878 lines
31 KiB
Rust
878 lines
31 KiB
Rust
use super::tokenizer::Tokenizer;
|
|
use super::{ChatTurn, ModelSummary};
|
|
use crate::model::ModelChoice;
|
|
use crate::settings::ReasoningMode;
|
|
use memmap2::{Mmap, MmapOptions};
|
|
use serde::de::{MapAccess, Visitor};
|
|
use serde::{Deserialize, Deserializer};
|
|
use serde_json::Value;
|
|
use sha2::{Digest, Sha256};
|
|
use std::collections::{BTreeMap, BTreeSet, HashMap};
|
|
use std::fmt;
|
|
use std::fs::{self, File};
|
|
use std::io::Read;
|
|
use std::path::{Path, PathBuf};
|
|
|
|
const MAX_SAFETENSORS_HEADER: u64 = 16 * 1024 * 1024;
|
|
const MANIFEST: &[u8] = include_bytes!("../../assets/models/qwen38-flash-next-bare-speed.json");
|
|
const INVENTORY: &str =
|
|
include_str!("../../assets/models/qwen38-flash-next-bare-speed-tensors.tsv");
|
|
const CORE_BYTES: u64 = 71_742_682_599;
|
|
const PLE_BYTES: u64 = 32_000_154_008;
|
|
const MTP_BYTES: u64 = 1_672_575_532;
|
|
const KV_BYTES_PER_TOKEN: u64 = 24_576;
|
|
const QSA_RAW_BYTES_PER_TOKEN: u64 = 3_072;
|
|
const QSA_POOLED_BYTES_PER_BLOCK: u64 = 3_072;
|
|
const QSA_POOL_RATIO: u64 = 4;
|
|
const QSA_FIXED_SCRATCH_BYTES: u64 = 6_656;
|
|
const QSA_TOPK_SCRATCH_BYTES_PER_BLOCK: u64 = 8;
|
|
const MTP_TOKEN_RESERVE: u64 = 3;
|
|
const GDN_STATE_BYTES: u64 = 113_246_208;
|
|
const GDN_CONV_BYTES: u64 = 2_211_840;
|
|
const PLE_CONV_BYTES: u64 = 184_320;
|
|
|
|
#[derive(Deserialize)]
|
|
struct Manifest {
|
|
config: BTreeMap<String, Value>,
|
|
runtime: BTreeMap<String, Value>,
|
|
}
|
|
|
|
#[derive(Clone, Debug, Eq, PartialEq)]
|
|
struct ExpectedTensor {
|
|
dtype: String,
|
|
shape: Vec<u64>,
|
|
quant_bits: Option<u32>,
|
|
group_size: Option<u64>,
|
|
quant_mode: Option<String>,
|
|
start: u64,
|
|
end: u64,
|
|
}
|
|
|
|
#[derive(Debug, Deserialize)]
|
|
struct Tensor {
|
|
dtype: String,
|
|
shape: Vec<u64>,
|
|
#[serde(rename = "data_offsets")]
|
|
offsets: [u64; 2],
|
|
}
|
|
|
|
#[derive(Debug, Eq, PartialEq)]
|
|
pub(super) struct MemoryPlan {
|
|
pub(super) resident_core: u64,
|
|
pub(super) mapped_ple: u64,
|
|
pub(super) optional_mtp: u64,
|
|
pub(super) kv_and_recurrent: u64,
|
|
pub(super) prefill_transient: u64,
|
|
pub(super) admission: u64,
|
|
}
|
|
|
|
pub(super) struct LoadedArtifacts {
|
|
#[allow(dead_code)]
|
|
pub(super) tokenizer: Tokenizer,
|
|
#[allow(dead_code)]
|
|
pub(super) memory: MemoryPlan,
|
|
#[allow(dead_code)]
|
|
pub(super) bindings: ArtifactBindings,
|
|
#[allow(dead_code)]
|
|
pub(super) tensor_count: usize,
|
|
}
|
|
|
|
pub(super) struct TensorBinding {
|
|
#[allow(dead_code)]
|
|
pub(super) file: PathBuf,
|
|
#[allow(dead_code)]
|
|
pub(super) name: String,
|
|
#[allow(dead_code)]
|
|
pub(super) dtype: String,
|
|
#[allow(dead_code)]
|
|
pub(super) shape: Vec<u64>,
|
|
#[allow(dead_code)]
|
|
pub(super) quant_bits: Option<u32>,
|
|
#[allow(dead_code)]
|
|
pub(super) group_size: Option<u64>,
|
|
#[allow(dead_code)]
|
|
pub(super) range: std::ops::Range<u64>,
|
|
}
|
|
|
|
pub(super) struct ArtifactBindings {
|
|
pub(super) core: Vec<TensorBinding>,
|
|
pub(super) ple: Vec<TensorBinding>,
|
|
pub(super) mtp: Vec<TensorBinding>,
|
|
}
|
|
|
|
pub(super) struct QwenMap {
|
|
path: PathBuf,
|
|
map: Mmap,
|
|
}
|
|
|
|
#[derive(Clone)]
|
|
pub(super) struct QwenTensor {
|
|
pub(super) map: usize,
|
|
pub(super) name: String,
|
|
pub(super) dtype: String,
|
|
pub(super) shape: Vec<u64>,
|
|
pub(super) quant_bits: Option<u32>,
|
|
pub(super) group_size: Option<u64>,
|
|
pub(super) range: std::ops::Range<u64>,
|
|
}
|
|
|
|
pub(super) struct QwenModel {
|
|
tokenizer: Tokenizer,
|
|
memory: MemoryPlan,
|
|
maps: Vec<QwenMap>,
|
|
tensors: HashMap<String, QwenTensor>,
|
|
identity: [u8; 32],
|
|
}
|
|
|
|
impl QwenModel {
|
|
pub(super) fn open(root: &Path, context: u32) -> Result<Self, String> {
|
|
let loaded = load(root, context, false)?;
|
|
let mut bindings = loaded.bindings.core;
|
|
bindings.extend(loaded.bindings.ple);
|
|
let mut paths = bindings
|
|
.iter()
|
|
.map(|binding| binding.file.clone())
|
|
.collect::<Vec<_>>();
|
|
paths.sort();
|
|
paths.dedup();
|
|
let mut maps = Vec::with_capacity(paths.len());
|
|
let mut map_indices = HashMap::with_capacity(paths.len());
|
|
for path in paths {
|
|
let file = File::open(&path).map_err(|error| format!("{}: {error}", path.display()))?;
|
|
// SAFETY: verified managed artifacts remain read-only while the model owns each mapping.
|
|
let map = unsafe { MmapOptions::new().map(&file) }
|
|
.map_err(|error| format!("cannot map {}: {error}", path.display()))?;
|
|
map_indices.insert(path.clone(), maps.len());
|
|
maps.push(QwenMap { path, map });
|
|
}
|
|
let tensors = bindings
|
|
.into_iter()
|
|
.map(|binding| {
|
|
let tensor = QwenTensor {
|
|
map: map_indices[&binding.file],
|
|
name: binding.name.clone(),
|
|
dtype: binding.dtype,
|
|
shape: binding.shape,
|
|
quant_bits: binding.quant_bits,
|
|
group_size: binding.group_size,
|
|
range: binding.range,
|
|
};
|
|
(binding.name, tensor)
|
|
})
|
|
.collect::<HashMap<_, _>>();
|
|
let mut hash = Sha256::new();
|
|
hash.update(b"DS4Server Qwen3.8 checkpoint identity v1");
|
|
hash.update(MANIFEST);
|
|
let identity = hash.finalize().into();
|
|
Ok(Self {
|
|
tokenizer: loaded.tokenizer,
|
|
memory: loaded.memory,
|
|
maps,
|
|
tensors,
|
|
identity,
|
|
})
|
|
}
|
|
|
|
pub(super) fn tensor(&self, name: &str) -> Result<&QwenTensor, String> {
|
|
self.tensors
|
|
.get(name)
|
|
.ok_or_else(|| format!("Qwen core tensor is missing: {name}"))
|
|
}
|
|
|
|
pub(super) fn map(&self, index: usize) -> (&[u8], &Path) {
|
|
(&self.maps[index].map, &self.maps[index].path)
|
|
}
|
|
|
|
pub(super) fn tensor_bytes<'a>(&'a self, tensor: &QwenTensor) -> Result<&'a [u8], String> {
|
|
let map = &self.maps[tensor.map].map;
|
|
let start = usize::try_from(tensor.range.start)
|
|
.map_err(|_| format!("{} starts beyond this platform", tensor.name))?;
|
|
let end = usize::try_from(tensor.range.end)
|
|
.map_err(|_| format!("{} ends beyond this platform", tensor.name))?;
|
|
map.get(start..end)
|
|
.ok_or_else(|| format!("{} is outside its mapped artifact", tensor.name))
|
|
}
|
|
|
|
pub(super) fn checkpoint_identity(&self) -> [u8; 32] {
|
|
self.identity
|
|
}
|
|
|
|
pub(super) fn summary(&self) -> ModelSummary {
|
|
ModelSummary {
|
|
model: ModelChoice::Qwen38FlashNext,
|
|
mapped_bytes: self.maps.iter().map(|item| item.map.len() as u64).sum(),
|
|
tensor_count: self.tensors.len(),
|
|
vocabulary_size: self.tokenizer.vocab_size(),
|
|
support_loaded: false,
|
|
vision_loaded: false,
|
|
}
|
|
}
|
|
|
|
pub(super) fn render_conversation(
|
|
&self,
|
|
system: &str,
|
|
messages: &[ChatTurn],
|
|
reasoning: ReasoningMode,
|
|
) -> Vec<i32> {
|
|
self.tokenizer
|
|
.encode_conversation(system, messages, reasoning)
|
|
}
|
|
|
|
pub(super) fn render_history(
|
|
&self,
|
|
system: &str,
|
|
messages: &[ChatTurn],
|
|
reasoning: ReasoningMode,
|
|
) -> Vec<i32> {
|
|
self.tokenizer.encode_history(system, messages, reasoning)
|
|
}
|
|
|
|
pub(super) fn render_continuation(
|
|
&self,
|
|
prompt: &str,
|
|
reasoning: ReasoningMode,
|
|
skip_previous_eos: bool,
|
|
) -> Vec<i32> {
|
|
self.tokenizer
|
|
.encode_continuation(prompt, reasoning, skip_previous_eos)
|
|
}
|
|
|
|
pub(super) fn token_bytes(&self, token: i32) -> Option<Vec<u8>> {
|
|
self.tokenizer.token_bytes(token)
|
|
}
|
|
|
|
pub(super) fn is_stop_token_for_reasoning(&self, token: i32, reasoning: ReasoningMode) -> bool {
|
|
self.tokenizer.is_stop(token)
|
|
|| (reasoning == ReasoningMode::Direct
|
|
&& (self.tokenizer.is_think_start(token) || self.tokenizer.is_think_end(token)))
|
|
}
|
|
|
|
pub(super) fn is_think_start_token(&self, token: i32) -> bool {
|
|
self.tokenizer.is_think_start(token)
|
|
}
|
|
|
|
pub(super) fn is_think_end_token(&self, token: i32) -> bool {
|
|
self.tokenizer.is_think_end(token)
|
|
}
|
|
|
|
pub(super) fn memory(&self) -> &MemoryPlan {
|
|
&self.memory
|
|
}
|
|
|
|
#[cfg(test)]
|
|
pub(super) fn mapped_residency(&self) -> Result<(u64, u64), String> {
|
|
// SAFETY: sysconf is read-only and has no pointer preconditions.
|
|
let page = unsafe { libc::sysconf(libc::_SC_PAGESIZE) };
|
|
if page <= 0 {
|
|
return Err("macOS did not report its virtual-memory page size".into());
|
|
}
|
|
let page = page as usize;
|
|
let mut core = 0_u64;
|
|
let mut ple = 0_u64;
|
|
for item in &self.maps {
|
|
let mut pages = vec![0_i8; item.map.len().div_ceil(page)];
|
|
// SAFETY: each read-only mmap and residency vector remain valid for this call.
|
|
if unsafe {
|
|
libc::mincore(
|
|
item.map.as_ptr().cast_mut().cast(),
|
|
item.map.len(),
|
|
pages.as_mut_ptr(),
|
|
)
|
|
} != 0
|
|
{
|
|
return Err(format!(
|
|
"cannot inspect residency for {}: {}",
|
|
item.path.display(),
|
|
std::io::Error::last_os_error()
|
|
));
|
|
}
|
|
let bytes = (pages.iter().filter(|value| **value & 1 != 0).count() * page)
|
|
.min(item.map.len()) as u64;
|
|
if item
|
|
.path
|
|
.file_name()
|
|
.is_some_and(|name| name == "ngram-table.safetensors")
|
|
{
|
|
ple += bytes;
|
|
} else {
|
|
core += bytes;
|
|
}
|
|
}
|
|
Ok((core, ple))
|
|
}
|
|
}
|
|
|
|
pub(crate) fn validate_artifacts(root: &Path) -> Result<(), String> {
|
|
load(root, 262_144, true).map(|_| ())
|
|
}
|
|
|
|
pub(super) fn load(root: &Path, context: u32, enable_mtp: bool) -> Result<LoadedArtifacts, String> {
|
|
let manifest: Manifest = serde_json::from_slice(MANIFEST)
|
|
.map_err(|error| format!("embedded Qwen manifest is invalid: {error}"))?;
|
|
validate_json(root, "config.json", &manifest.config)?;
|
|
validate_json(root, "mtplx_runtime.json", &manifest.runtime)?;
|
|
let expected = parse_inventory()?;
|
|
let bindings = validate_tensor_files(root, &expected)?;
|
|
validate_index(root, &expected)?;
|
|
let tokenizer = Tokenizer::load_qwen(&root.join("tokenizer.json"))?;
|
|
if tokenizer.vocab_size() != 248_320 {
|
|
return Err(format!(
|
|
"Qwen tokenizer has {} entries, expected 248320",
|
|
tokenizer.vocab_size()
|
|
));
|
|
}
|
|
tokenizer.validate_qwen_contract()?;
|
|
let memory = memory_plan(context, enable_mtp, 512)?;
|
|
Ok(LoadedArtifacts {
|
|
tokenizer,
|
|
memory,
|
|
bindings,
|
|
tensor_count: expected.len(),
|
|
})
|
|
}
|
|
|
|
fn validate_json(
|
|
root: &Path,
|
|
file_name: &str,
|
|
expected: &BTreeMap<String, Value>,
|
|
) -> Result<(), String> {
|
|
let path = root.join(file_name);
|
|
let value: Value = serde_json::from_slice(
|
|
&fs::read(&path).map_err(|error| format!("{}: {error}", path.display()))?,
|
|
)
|
|
.map_err(|error| format!("{}: {error}", path.display()))?;
|
|
validate_json_value(file_name, &value, expected)
|
|
}
|
|
|
|
fn validate_json_value(
|
|
file_name: &str,
|
|
value: &Value,
|
|
expected: &BTreeMap<String, Value>,
|
|
) -> Result<(), String> {
|
|
for (pointer, expected) in expected {
|
|
let actual = value
|
|
.pointer(pointer)
|
|
.ok_or_else(|| format!("{file_name} is missing {pointer}"))?;
|
|
if actual != expected {
|
|
return Err(format!(
|
|
"{file_name} {pointer} is {actual}, expected {expected}"
|
|
));
|
|
}
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
fn parse_inventory() -> Result<BTreeMap<(String, String), ExpectedTensor>, String> {
|
|
let mut tensors = BTreeMap::new();
|
|
for (line_number, line) in INVENTORY.lines().enumerate().skip(1) {
|
|
let fields = line.split('\t').collect::<Vec<_>>();
|
|
if fields.len() != 9 {
|
|
return Err(format!(
|
|
"embedded tensor inventory line {} is invalid",
|
|
line_number + 1
|
|
));
|
|
}
|
|
let parse = |field: &str, name: &str| {
|
|
field.parse::<u64>().map_err(|error| {
|
|
format!(
|
|
"inventory line {} has invalid {name}: {error}",
|
|
line_number + 1
|
|
)
|
|
})
|
|
};
|
|
let tensor = ExpectedTensor {
|
|
dtype: fields[2].to_owned(),
|
|
shape: fields[3]
|
|
.split('x')
|
|
.map(|dimension| parse(dimension, "shape"))
|
|
.collect::<Result<_, _>>()?,
|
|
quant_bits: (!fields[4].is_empty())
|
|
.then(|| fields[4].parse::<u32>())
|
|
.transpose()
|
|
.map_err(|error| format!("inventory has invalid quantization: {error}"))?,
|
|
group_size: (!fields[5].is_empty())
|
|
.then(|| parse(fields[5], "group size"))
|
|
.transpose()?,
|
|
quant_mode: (!fields[6].is_empty()).then(|| fields[6].to_owned()),
|
|
start: parse(fields[7], "start")?,
|
|
end: parse(fields[8], "end")?,
|
|
};
|
|
validate_precision(fields[1], &tensor)?;
|
|
if tensors
|
|
.insert((fields[0].to_owned(), fields[1].to_owned()), tensor)
|
|
.is_some()
|
|
{
|
|
return Err(format!("duplicate inventory tensor {}", fields[1]));
|
|
}
|
|
}
|
|
if tensors.len() != 2_527 {
|
|
return Err(format!(
|
|
"embedded tensor inventory has {} entries, expected 2527",
|
|
tensors.len()
|
|
));
|
|
}
|
|
Ok(tensors)
|
|
}
|
|
|
|
fn validate_precision(name: &str, tensor: &ExpectedTensor) -> Result<(), String> {
|
|
match (
|
|
tensor.quant_bits,
|
|
tensor.group_size,
|
|
tensor.quant_mode.as_deref(),
|
|
) {
|
|
(None, None, None) if matches!(tensor.dtype.as_str(), "BF16" | "I64") => Ok(()),
|
|
(Some(bits @ (2 | 4 | 8)), Some(group @ (32 | 64)), Some("affine"))
|
|
if tensor.dtype == "U32" || name.ends_with(".scales") || name.ends_with(".biases") =>
|
|
{
|
|
if group == 32 && !name.starts_with("ngram.") {
|
|
return Err(format!("{name} unexpectedly uses 32-value groups"));
|
|
}
|
|
if bits == 2 && !name.starts_with("mtp.") {
|
|
return Err(format!("{name} unexpectedly uses 2-bit weights"));
|
|
}
|
|
Ok(())
|
|
}
|
|
_ => Err(format!("{name} has an unsupported precision contract")),
|
|
}
|
|
}
|
|
|
|
fn validate_tensor_files(
|
|
root: &Path,
|
|
expected: &BTreeMap<(String, String), ExpectedTensor>,
|
|
) -> Result<ArtifactBindings, String> {
|
|
let files = expected
|
|
.keys()
|
|
.map(|(file, _)| file.as_str())
|
|
.collect::<BTreeSet<_>>();
|
|
let mut seen = BTreeSet::new();
|
|
let mut bindings = ArtifactBindings {
|
|
core: Vec::new(),
|
|
ple: Vec::new(),
|
|
mtp: Vec::new(),
|
|
};
|
|
for file_name in files {
|
|
let path = root.join(file_name);
|
|
let (data_start, tensors) = read_header(&path)?;
|
|
for (name, tensor) in tensors {
|
|
let key = (file_name.to_owned(), name.clone());
|
|
let contract = expected
|
|
.get(&key)
|
|
.ok_or_else(|| format!("{} contains unexpected tensor {name}", path.display()))?;
|
|
let start = data_start
|
|
.checked_add(tensor.offsets[0])
|
|
.ok_or_else(|| format!("{name} start offset overflows"))?;
|
|
let end = data_start
|
|
.checked_add(tensor.offsets[1])
|
|
.ok_or_else(|| format!("{name} end offset overflows"))?;
|
|
if tensor.dtype != contract.dtype
|
|
|| tensor.shape != contract.shape
|
|
|| start != contract.start
|
|
|| end != contract.end
|
|
{
|
|
return Err(format!("{name} does not match the frozen tensor layout"));
|
|
}
|
|
let binding = TensorBinding {
|
|
file: path.clone(),
|
|
name: name.clone(),
|
|
dtype: contract.dtype.clone(),
|
|
shape: contract.shape.clone(),
|
|
quant_bits: contract.quant_bits,
|
|
group_size: contract.group_size,
|
|
range: start..end,
|
|
};
|
|
match file_name {
|
|
"ngram-table.safetensors" => bindings.ple.push(binding),
|
|
"mtp.safetensors" => bindings.mtp.push(binding),
|
|
_ => bindings.core.push(binding),
|
|
}
|
|
seen.insert(key);
|
|
}
|
|
}
|
|
if seen.len() != expected.len() {
|
|
let missing = expected
|
|
.keys()
|
|
.find(|key| !seen.contains(*key))
|
|
.map(|(_, name)| name.as_str())
|
|
.unwrap_or("unknown tensor");
|
|
return Err(format!("artifact set is missing {missing}"));
|
|
}
|
|
let ple = seen
|
|
.iter()
|
|
.filter(|(file, _)| file == "ngram-table.safetensors")
|
|
.count();
|
|
let mtp = seen
|
|
.iter()
|
|
.filter(|(file, _)| file == "mtp.safetensors")
|
|
.count();
|
|
if ple != 3 || mtp != 58 {
|
|
return Err(format!(
|
|
"artifact set has {ple} PLE and {mtp} MTP tensors, expected 3 and 58"
|
|
));
|
|
}
|
|
Ok(bindings)
|
|
}
|
|
|
|
fn read_header(path: &Path) -> Result<(u64, BTreeMap<String, Tensor>), String> {
|
|
let mut file = File::open(path).map_err(|error| format!("{}: {error}", path.display()))?;
|
|
let size = file
|
|
.metadata()
|
|
.map_err(|error| format!("{}: {error}", path.display()))?
|
|
.len();
|
|
let mut length = [0_u8; 8];
|
|
file.read_exact(&mut length)
|
|
.map_err(|error| format!("{}: {error}", path.display()))?;
|
|
let length = u64::from_le_bytes(length);
|
|
if length == 0 || length > MAX_SAFETENSORS_HEADER {
|
|
return Err(format!(
|
|
"{} has invalid header size {length}",
|
|
path.display()
|
|
));
|
|
}
|
|
let data_start = 8_u64
|
|
.checked_add(length)
|
|
.ok_or_else(|| format!("{} header overflows", path.display()))?;
|
|
let mut bytes = vec![0_u8; length as usize];
|
|
file.read_exact(&mut bytes)
|
|
.map_err(|error| format!("{}: {error}", path.display()))?;
|
|
let mut values = serde_json::from_slice::<UniqueObject>(&bytes)
|
|
.map_err(|error| format!("{}: {error}", path.display()))?
|
|
.0;
|
|
values.remove("__metadata__");
|
|
let mut tensors = BTreeMap::new();
|
|
for (name, value) in values {
|
|
let tensor: Tensor = serde_json::from_value(value)
|
|
.map_err(|error| format!("{} tensor {name}: {error}", path.display()))?;
|
|
let end = data_start
|
|
.checked_add(tensor.offsets[1])
|
|
.ok_or_else(|| format!("{name} offset overflows"))?;
|
|
if tensor.shape.is_empty() || tensor.offsets[0] > tensor.offsets[1] || end > size {
|
|
return Err(format!(
|
|
"{} tensor {name} has invalid layout",
|
|
path.display()
|
|
));
|
|
}
|
|
if tensors.insert(name.clone(), tensor).is_some() {
|
|
return Err(format!(
|
|
"{} contains duplicate tensor {name}",
|
|
path.display()
|
|
));
|
|
}
|
|
}
|
|
Ok((data_start, tensors))
|
|
}
|
|
|
|
struct UniqueObject(BTreeMap<String, Value>);
|
|
|
|
impl<'de> Deserialize<'de> for UniqueObject {
|
|
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
|
where
|
|
D: Deserializer<'de>,
|
|
{
|
|
struct UniqueObjectVisitor;
|
|
|
|
impl<'de> Visitor<'de> for UniqueObjectVisitor {
|
|
type Value = UniqueObject;
|
|
|
|
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
|
formatter.write_str("a safetensors header with unique tensor names")
|
|
}
|
|
|
|
fn visit_map<M>(self, mut map: M) -> Result<Self::Value, M::Error>
|
|
where
|
|
M: MapAccess<'de>,
|
|
{
|
|
let mut values = BTreeMap::new();
|
|
while let Some((name, value)) = map.next_entry::<String, Value>()? {
|
|
if values.insert(name.clone(), value).is_some() {
|
|
return Err(serde::de::Error::custom(format!(
|
|
"duplicate tensor name {name}"
|
|
)));
|
|
}
|
|
}
|
|
Ok(UniqueObject(values))
|
|
}
|
|
}
|
|
|
|
deserializer.deserialize_map(UniqueObjectVisitor)
|
|
}
|
|
}
|
|
|
|
fn validate_index(
|
|
root: &Path,
|
|
expected: &BTreeMap<(String, String), ExpectedTensor>,
|
|
) -> Result<(), String> {
|
|
let path = root.join("model.safetensors.index.json");
|
|
let value: Value = serde_json::from_slice(
|
|
&fs::read(&path).map_err(|error| format!("{}: {error}", path.display()))?,
|
|
)
|
|
.map_err(|error| format!("{}: {error}", path.display()))?;
|
|
let map = value
|
|
.get("weight_map")
|
|
.and_then(Value::as_object)
|
|
.ok_or_else(|| "Qwen model index has no weight_map".to_owned())?;
|
|
let core = expected
|
|
.keys()
|
|
.filter(|(file, _)| file.starts_with("model-") && file.ends_with(".safetensors"))
|
|
.map(|(file, name)| (name.as_str(), file.as_str()))
|
|
.collect::<BTreeMap<_, _>>();
|
|
for (name, file) in &core {
|
|
if map.get(*name).and_then(Value::as_str) != Some(file) {
|
|
return Err(format!("model index does not bind {name} to {file}"));
|
|
}
|
|
}
|
|
if core.len() != 2_466 {
|
|
return Err(format!(
|
|
"Qwen core has {} tensors, expected 2466",
|
|
core.len()
|
|
));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
pub(super) fn memory_plan(
|
|
context: u32,
|
|
enable_mtp: bool,
|
|
prefill_chunk: u32,
|
|
) -> Result<MemoryPlan, String> {
|
|
if context == 0 || context > 262_144 {
|
|
return Err("Qwen context must be between 1 and 262144 tokens".into());
|
|
}
|
|
if prefill_chunk == 0 {
|
|
return Err("Qwen prefill chunk must be positive".into());
|
|
}
|
|
let token_capacity = u64::from(context)
|
|
.checked_add(MTP_TOKEN_RESERVE)
|
|
.ok_or_else(|| "Qwen attention capacity overflows".to_owned())?;
|
|
let block_capacity = token_capacity.div_ceil(QSA_POOL_RATIO);
|
|
let topk_scratch = if context > 2_048 {
|
|
u64::from(context) / QSA_POOL_RATIO * QSA_TOPK_SCRATCH_BYTES_PER_BLOCK
|
|
} else {
|
|
0
|
|
};
|
|
let kv = (KV_BYTES_PER_TOKEN + QSA_RAW_BYTES_PER_TOKEN)
|
|
.checked_mul(token_capacity)
|
|
.and_then(|bytes| {
|
|
QSA_POOLED_BYTES_PER_BLOCK
|
|
.checked_mul(block_capacity)
|
|
.and_then(|pooled| bytes.checked_add(pooled))
|
|
})
|
|
.ok_or_else(|| "Qwen KV memory size overflows".to_owned())?;
|
|
let kv_and_recurrent = kv
|
|
.checked_add(GDN_STATE_BYTES + GDN_CONV_BYTES + PLE_CONV_BYTES)
|
|
.ok_or_else(|| "Qwen recurrent memory size overflows".to_owned())?;
|
|
let prefill_transient = u64::from(prefill_chunk)
|
|
.checked_mul((4 * 2_560 + 2_048 + 2_048 + 6_144 + 6_144) * 2)
|
|
.and_then(|bytes| bytes.checked_add(block_capacity * 4))
|
|
.and_then(|bytes| bytes.checked_add(QSA_FIXED_SCRATCH_BYTES))
|
|
.and_then(|bytes| bytes.checked_add(topk_scratch))
|
|
.ok_or_else(|| "Qwen prefill transient size overflows".to_owned())?;
|
|
let admitted_mtp = if enable_mtp { MTP_BYTES } else { 0 };
|
|
let admission = CORE_BYTES
|
|
.checked_add(admitted_mtp)
|
|
.and_then(|bytes| bytes.checked_add(kv_and_recurrent))
|
|
.and_then(|bytes| bytes.checked_add(prefill_transient))
|
|
.ok_or_else(|| "Qwen admission size overflows".to_owned())?;
|
|
Ok(MemoryPlan {
|
|
resident_core: CORE_BYTES,
|
|
mapped_ple: PLE_BYTES,
|
|
optional_mtp: MTP_BYTES,
|
|
kv_and_recurrent,
|
|
prefill_transient,
|
|
admission,
|
|
})
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use crate::engine::ChatTurn;
|
|
use crate::settings::ReasoningMode;
|
|
use std::time::{SystemTime, UNIX_EPOCH};
|
|
|
|
#[test]
|
|
fn frozen_inventory_and_memory_categories_are_exact() {
|
|
let inventory = parse_inventory().unwrap();
|
|
assert_eq!(inventory.len(), 2_527);
|
|
assert_eq!(
|
|
inventory
|
|
.keys()
|
|
.filter(|(file, _)| file == "ngram-table.safetensors")
|
|
.count(),
|
|
3
|
|
);
|
|
let plan = memory_plan(262_144, true, 512).unwrap();
|
|
assert_eq!(plan.resident_core, CORE_BYTES);
|
|
assert_eq!(plan.mapped_ple, PLE_BYTES);
|
|
assert_eq!(plan.optional_mtp, MTP_BYTES);
|
|
assert_eq!(plan.kv_and_recurrent, 7_564_812_288);
|
|
assert_eq!(plan.prefill_transient, 28_056_068);
|
|
assert_eq!(plan.admission, 81_008_126_487);
|
|
let without_mtp = memory_plan(262_144, false, 512).unwrap();
|
|
assert_eq!(without_mtp.optional_mtp, MTP_BYTES);
|
|
assert_eq!(without_mtp.admission, 79_335_550_955);
|
|
assert!(memory_plan(0, false, 512).is_err());
|
|
assert!(memory_plan(262_145, false, 512).is_err());
|
|
assert!(memory_plan(1, false, 0).is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn metadata_and_safetensors_boundaries_fail_closed() {
|
|
let root = std::env::temp_dir().join(format!(
|
|
"ds4-qwen-loader-{}-{}",
|
|
std::process::id(),
|
|
SystemTime::now()
|
|
.duration_since(UNIX_EPOCH)
|
|
.unwrap()
|
|
.as_nanos()
|
|
));
|
|
fs::create_dir_all(&root).unwrap();
|
|
let expected = BTreeMap::from([("/model_type".into(), Value::String("qwen4_exp".into()))]);
|
|
fs::write(root.join("config.json"), br#"{"model_type":"qwen4_exp"}"#).unwrap();
|
|
assert!(validate_json(&root, "config.json", &expected).is_ok());
|
|
fs::write(root.join("config.json"), br#"{"model_type":"qwen3_next"}"#).unwrap();
|
|
assert!(validate_json(&root, "config.json", &expected).is_err());
|
|
fs::write(root.join("config.json"), b"{}").unwrap();
|
|
assert!(validate_json(&root, "config.json", &expected).is_err());
|
|
|
|
let header = serde_json::to_vec(&serde_json::json!({
|
|
"tensor": {"dtype":"BF16", "shape":[2], "data_offsets":[0,4]}
|
|
}))
|
|
.unwrap();
|
|
let tensor_path = root.join("fixture.safetensors");
|
|
let mut bytes = (header.len() as u64).to_le_bytes().to_vec();
|
|
bytes.extend(header);
|
|
bytes.extend([0_u8; 4]);
|
|
fs::write(&tensor_path, bytes).unwrap();
|
|
let (_, tensors) = read_header(&tensor_path).unwrap();
|
|
assert_eq!(tensors.len(), 1);
|
|
let mut truncated = fs::read(&tensor_path).unwrap();
|
|
truncated.pop();
|
|
fs::write(&tensor_path, truncated).unwrap();
|
|
assert!(read_header(&tensor_path).is_err());
|
|
let duplicate = br#"{"tensor":{"dtype":"BF16","shape":[1],"data_offsets":[0,2]},"tensor":{"dtype":"BF16","shape":[1],"data_offsets":[2,4]}}"#;
|
|
let mut bytes = (duplicate.len() as u64).to_le_bytes().to_vec();
|
|
bytes.extend(duplicate);
|
|
bytes.extend([0_u8; 4]);
|
|
fs::write(&tensor_path, bytes).unwrap();
|
|
assert!(
|
|
read_header(&tensor_path)
|
|
.unwrap_err()
|
|
.contains("duplicate tensor")
|
|
);
|
|
fs::remove_dir_all(root).unwrap();
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "requires DS4_QWEN38_ARTIFACTS to point at the pinned 105 GB source"]
|
|
fn pinned_artifact_set_loads_and_renders_goldens() {
|
|
let root = std::env::var_os("DS4_QWEN38_ARTIFACTS")
|
|
.map(std::path::PathBuf::from)
|
|
.expect("DS4_QWEN38_ARTIFACTS is set");
|
|
let loaded = load(&root, 131_072, false).unwrap();
|
|
assert_eq!(loaded.bindings.core.len(), 2_466);
|
|
assert_eq!(loaded.bindings.ple.len(), 3);
|
|
assert_eq!(loaded.bindings.mtp.len(), 58);
|
|
let manifest: Manifest = serde_json::from_slice(MANIFEST).unwrap();
|
|
for (file_name, contract) in [
|
|
("config.json", &manifest.config),
|
|
("mtplx_runtime.json", &manifest.runtime),
|
|
] {
|
|
let value: Value =
|
|
serde_json::from_slice(&fs::read(root.join(file_name)).unwrap()).unwrap();
|
|
assert!(validate_json_value(file_name, &value, contract).is_ok());
|
|
for pointer in contract.keys() {
|
|
let mut invalid = value.clone();
|
|
*invalid.pointer_mut(pointer).unwrap() = Value::Null;
|
|
assert!(
|
|
validate_json_value(file_name, &invalid, contract).is_err(),
|
|
"{file_name} accepted invalid {pointer}"
|
|
);
|
|
}
|
|
}
|
|
let user = ChatTurn {
|
|
user: true,
|
|
tool: false,
|
|
system: false,
|
|
skip_previous_eos: false,
|
|
reasoning: None,
|
|
reasoning_complete: true,
|
|
content: "Hi".into(),
|
|
};
|
|
let direct = loaded.tokenizer.encode_conversation(
|
|
"",
|
|
std::slice::from_ref(&user),
|
|
ReasoningMode::Direct,
|
|
);
|
|
assert_eq!(
|
|
direct,
|
|
[
|
|
248_045, 846, 198, 12_675, 248_046, 198, 248_045, 74_455, 198, 13_314, 741, 29,
|
|
271, 510, 26_003, 29, 271,
|
|
]
|
|
);
|
|
assert_eq!(
|
|
decode(&loaded.tokenizer, &direct),
|
|
"<|im_start|>user\nHi<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"
|
|
);
|
|
|
|
let thinking = loaded.tokenizer.encode_conversation(
|
|
"",
|
|
std::slice::from_ref(&user),
|
|
ReasoningMode::XHigh,
|
|
);
|
|
assert_eq!(
|
|
decode(&loaded.tokenizer, &thinking),
|
|
"<|im_start|>system\nReasoning effort is set to xhigh. Please think carefully through the task, validate key assumptions, consider plausible alternatives, and prioritize correctness, consistency, and clarity in the final answer.<|im_end|>\n<|im_start|>user\nHi<|im_end|>\n<|im_start|>assistant\n<think>\n"
|
|
);
|
|
|
|
let messages = [
|
|
user,
|
|
ChatTurn {
|
|
user: false,
|
|
tool: false,
|
|
system: false,
|
|
skip_previous_eos: false,
|
|
reasoning: Some("check".into()),
|
|
reasoning_complete: true,
|
|
content: "<tool_call>\n<function=read>\n<parameter=path>\na.rs\n</parameter>\n</function>\n</tool_call>".into(),
|
|
},
|
|
ChatTurn {
|
|
user: false,
|
|
tool: true,
|
|
system: false,
|
|
skip_previous_eos: false,
|
|
reasoning: None,
|
|
reasoning_complete: true,
|
|
content: "ok".into(),
|
|
},
|
|
ChatTurn {
|
|
user: false,
|
|
tool: true,
|
|
system: false,
|
|
skip_previous_eos: false,
|
|
reasoning: None,
|
|
reasoning_complete: true,
|
|
content: "done".into(),
|
|
},
|
|
];
|
|
let tools =
|
|
loaded
|
|
.tokenizer
|
|
.encode_conversation("system", &messages, ReasoningMode::Medium);
|
|
assert_eq!(
|
|
decode(&loaded.tokenizer, &tools),
|
|
"<|im_start|>system\nsystem<|im_end|>\n<|im_start|>user\nHi<|im_end|>\n<|im_start|>assistant\n<think>\ncheck\n</think>\n\n<tool_call>\n<function=read>\n<parameter=path>\na.rs\n</parameter>\n</function>\n</tool_call><|im_end|>\n<|im_start|>user\n<tool_response>\nok\n</tool_response>\n<tool_response>\ndone\n</tool_response><|im_end|>\n<|im_start|>assistant\n<think>\n"
|
|
);
|
|
}
|
|
|
|
fn decode(tokenizer: &Tokenizer, tokens: &[i32]) -> String {
|
|
String::from_utf8(
|
|
tokens
|
|
.iter()
|
|
.flat_map(|token| tokenizer.token_bytes(*token).unwrap())
|
|
.collect(),
|
|
)
|
|
.unwrap()
|
|
}
|
|
}
|