Integrate Qwen3.8 model intake

This commit is contained in:
Georg Bauer
2026-09-03 19:56:23 +02:00
parent 4c83ac0360
commit 3773cfda2e
15 changed files with 1805 additions and 35 deletions

View File

@@ -2177,6 +2177,7 @@ impl SsdPlan {
ModelChoice::Glm52 | ModelChoice::Glm53Flash => {
unreachable!("GLM uses its dedicated executor")
}
ModelChoice::Qwen38FlashNext => unreachable!("Qwen uses its dedicated executor"),
};
for &(layer, expert) in hotlist {
if loaded == self.preload_experts {
@@ -4449,6 +4450,7 @@ impl Executor {
)
.map(Box::new)
.map(Self::Glm),
ModelFamily::Qwen => unreachable!("Qwen uses its dedicated executor"),
}
}
@@ -7679,7 +7681,9 @@ fn compression_ratio(shape: super::Shape, layer: u32) -> u32 {
}
crate::model::ModelChoice::DeepSeekV4Flash0731
| crate::model::ModelChoice::DeepSeekV4Pro => 128,
crate::model::ModelChoice::Glm52 | crate::model::ModelChoice::Glm53Flash => 0,
crate::model::ModelChoice::Glm52
| crate::model::ModelChoice::Glm53Flash
| crate::model::ModelChoice::Qwen38FlashNext => 0,
}
}

646
src/engine/qwen.rs Normal file
View File

@@ -0,0 +1,646 @@
use super::tokenizer::Tokenizer;
use serde::de::{MapAccess, Visitor};
use serde::{Deserialize, Deserializer};
use serde_json::Value;
use std::collections::{BTreeMap, BTreeSet};
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 GDN_STATE_BYTES: u64 = 113_246_208;
const GDN_CONV_BYTES: u64 = 2_211_840;
#[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(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 kv = KV_BYTES_PER_TOKEN
.checked_mul(u64::from(context))
.ok_or_else(|| "Qwen KV memory size overflows".to_owned())?;
let kv_and_recurrent = kv
.checked_add(GDN_STATE_BYTES + GDN_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)
.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, 6_557_908_992);
assert_eq!(plan.prefill_transient, 27_262_976);
assert_eq!(plan.admission, 80_000_430_099);
let without_mtp = memory_plan(262_144, false, 512).unwrap();
assert_eq!(without_mtp.optional_mtp, MTP_BYTES);
assert_eq!(without_mtp.admission, 78_327_854_567);
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()
}
}

View File

@@ -4,7 +4,10 @@ use super::{
VISION_TOKEN_END, VISION_TOKEN_START,
};
use crate::settings::ReasoningMode;
use serde::Deserialize;
use std::collections::HashMap;
use std::fs;
use std::path::Path;
const HIGH_REASONING_PREFIX: &str = "Reasoning Effort: Absolute maximum with no shortcuts permitted.\n\
You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\n\
@@ -20,7 +23,14 @@ fn reasoning_prefix(family: ModelFamily, reasoning: ReasoningMode) -> Option<&'s
(ModelFamily::Glm, ReasoningMode::Low) => Some("Reasoning Effort: Low"),
(ModelFamily::Glm, ReasoningMode::High) => Some("Reasoning Effort: High"),
(ModelFamily::Glm, ReasoningMode::Max) => Some("Reasoning Effort: Max"),
(_, ReasoningMode::Direct | ReasoningMode::Low) => None,
(ModelFamily::Qwen, _) => None,
(
_,
ReasoningMode::Direct
| ReasoningMode::Low
| ReasoningMode::Medium
| ReasoningMode::XHigh,
) => None,
}
}
@@ -38,7 +48,30 @@ pub(super) struct Tokenizer {
sop: i32,
think_start: i32,
think_end: i32,
rendered_specials: Vec<(&'static [u8], i32)>,
alternate_eos: i32,
rendered_specials: Vec<(Vec<u8>, i32)>,
}
#[derive(Deserialize)]
struct QwenTokenizerFile {
added_tokens: Vec<QwenAddedToken>,
model: QwenBpe,
}
#[derive(Deserialize)]
struct QwenAddedToken {
id: usize,
content: String,
special: bool,
}
#[derive(Deserialize)]
struct QwenBpe {
#[serde(rename = "type")]
kind: String,
vocab: HashMap<String, usize>,
merges: Vec<String>,
byte_fallback: bool,
}
impl Tokenizer {
@@ -95,6 +128,9 @@ impl Tokenizer {
lookup(b"<think>"),
lookup(b"</think>"),
),
ModelFamily::Qwen => {
return Err("Qwen tokenizers must be loaded from tokenizer.json".into());
}
};
if [bos, user, assistant, think_start, think_end]
.into_iter()
@@ -128,6 +164,7 @@ impl Tokenizer {
]
.into_iter()
.filter(|(_, token)| *token >= 0)
.map(|(marker, token)| (marker.to_vec(), token))
.collect();
Ok(Self {
@@ -144,6 +181,98 @@ impl Tokenizer {
sop,
think_start,
think_end,
alternate_eos: -1,
rendered_specials,
})
}
pub(super) fn load_qwen(path: &Path) -> Result<Self, String> {
let file: QwenTokenizerFile = serde_json::from_slice(
&fs::read(path).map_err(|error| format!("{}: {error}", path.display()))?,
)
.map_err(|error| format!("{}: {error}", path.display()))?;
if file.model.kind != "BPE" || file.model.byte_fallback {
return Err("Qwen tokenizer must use BPE without byte fallback".into());
}
let maximum = file
.model
.vocab
.values()
.copied()
.chain(file.added_tokens.iter().map(|token| token.id))
.max()
.ok_or_else(|| "Qwen tokenizer vocabulary is empty".to_owned())?;
if maximum != 248_076 {
return Err(format!(
"Qwen tokenizer ends at token id {maximum}, expected 248076"
));
}
// The model head has 248320 rows; the final 243 ids are deliberately
// reserved and have no tokenizer spelling in the pinned source.
let mut tokens = vec![Vec::new(); 248_320];
for (token, id) in file.model.vocab {
if !tokens[id].is_empty() {
return Err(format!("Qwen tokenizer has duplicate token id {id}"));
}
tokens[id] = token.into_bytes();
}
let mut rendered_specials = Vec::new();
for token in file.added_tokens {
if !tokens[token.id].is_empty() {
return Err(format!(
"Qwen tokenizer has duplicate token id {}",
token.id
));
}
tokens[token.id] = token.content.as_bytes().to_vec();
if token.special {
rendered_specials.push((token.content.into_bytes(), token.id as i32));
}
}
if tokens[..=maximum].iter().any(Vec::is_empty) {
return Err("Qwen tokenizer token ids are not contiguous".into());
}
let token_to_id = tokens
.iter()
.enumerate()
.filter(|(_, token)| !token.is_empty())
.map(|(id, token)| (token.clone(), id as i32))
.collect::<HashMap<_, _>>();
let required = |token: &[u8]| {
token_to_id.get(token).copied().ok_or_else(|| {
format!(
"required Qwen tokenizer token is missing: {}",
String::from_utf8_lossy(token)
)
})
};
let merge_rank = file
.model
.merges
.into_iter()
.enumerate()
.map(|(rank, merge)| (merge.into_bytes(), rank))
.collect();
let alternate_eos = required(b"<|endoftext|>")?;
let im_start = required(b"<|im_start|>")?;
let im_end = required(b"<|im_end|>")?;
let think_start = required(b"<think>")?;
let think_end = required(b"</think>")?;
Ok(Self {
family: ModelFamily::Qwen,
tokens,
token_to_id,
merge_rank,
bos: alternate_eos,
eos: im_end,
system: im_start,
user: im_start,
assistant: im_start,
observation: im_start,
sop: im_end,
think_start,
think_end,
alternate_eos,
rendered_specials,
})
}
@@ -152,6 +281,40 @@ impl Tokenizer {
self.tokens.len()
}
pub(super) fn validate_qwen_contract(&self) -> Result<(), String> {
if self.family != ModelFamily::Qwen {
return Err("tokenizer is not Qwen".into());
}
for (token, expected) in [
(b"<|endoftext|>".as_slice(), 248_044),
(b"<|im_start|>".as_slice(), 248_045),
(b"<|im_end|>".as_slice(), 248_046),
(b"<tool_call>".as_slice(), 248_058),
(b"</tool_call>".as_slice(), 248_059),
(b"<tool_response>".as_slice(), 248_066),
(b"</tool_response>".as_slice(), 248_067),
(b"<think>".as_slice(), 248_068),
(b"</think>".as_slice(), 248_069),
] {
if self.token_to_id.get(token).copied() != Some(expected) {
return Err(format!(
"Qwen special token {} does not have id {expected}",
String::from_utf8_lossy(token)
));
}
}
if ![248_044, 248_045, 248_046]
.into_iter()
.all(|token| self.is_ple_reset(token))
|| ![248_044, 248_046]
.into_iter()
.all(|token| self.is_stop(token))
{
return Err("Qwen EOS or PLE reset markers are invalid".into());
}
Ok(())
}
pub(super) fn token_bytes(&self, token: i32) -> Option<Vec<u8>> {
let token = self.tokens.get(usize::try_from(token).ok()?)?;
if token.windows(3).any(|window| window == [0xef, 0xbd, 0x9c]) {
@@ -163,10 +326,10 @@ impl Tokenizer {
pub(super) fn tokenize(&self, text: &str) -> Vec<i32> {
let mut output = Vec::new();
if self.family == ModelFamily::Glm {
self.tokenize_glm(text, &mut output);
} else {
self.tokenize_joyai(text, &mut output);
match self.family {
ModelFamily::DeepSeek => self.tokenize_joyai(text, &mut output),
ModelFamily::Glm => self.tokenize_gpt4(text, &mut output, 3),
ModelFamily::Qwen => self.tokenize_gpt4(text, &mut output, 1),
}
output
}
@@ -211,10 +374,10 @@ impl Tokenizer {
}
fn tokenize_plain(&self, text: &str, output: &mut Vec<i32>) {
if self.family == ModelFamily::Glm {
self.tokenize_glm(text, output);
} else {
self.tokenize_joyai(text, output);
match self.family {
ModelFamily::DeepSeek => self.tokenize_joyai(text, output),
ModelFamily::Glm => self.tokenize_gpt4(text, output, 3),
ModelFamily::Qwen => self.tokenize_gpt4(text, output, 1),
}
}
@@ -264,6 +427,14 @@ impl Tokenizer {
reasoning: ReasoningMode,
continue_assistant: bool,
) -> Vec<i32> {
if self.family == ModelFamily::Qwen {
return self.encode_qwen_messages(
system_prompt,
messages,
reasoning,
continue_assistant,
);
}
let mut output = vec![self.bos];
if self.family == ModelFamily::Glm && self.sop >= 0 {
output.push(self.sop);
@@ -353,12 +524,95 @@ impl Tokenizer {
output
}
fn encode_qwen_messages(
&self,
system_prompt: &str,
messages: &[ChatTurn],
reasoning: ReasoningMode,
continue_assistant: bool,
) -> Vec<i32> {
let instruction = match reasoning {
ReasoningMode::Direct | ReasoningMode::Medium => "",
ReasoningMode::Low => {
"Reasoning effort is set to low. Keep your thinking brief and focused, moving directly to the conclusion without unnecessary elaboration."
}
ReasoningMode::XHigh => {
"Reasoning 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."
}
ReasoningMode::High | ReasoningMode::Max => {
unreachable!("legacy reasoning modes are not valid for Qwen")
}
};
let mut rendered = String::new();
let system_prompt = system_prompt.trim();
if !instruction.is_empty() || !system_prompt.is_empty() {
rendered.push_str("<|im_start|>system\n");
if !instruction.is_empty() {
rendered.push_str(instruction);
if !system_prompt.is_empty() {
rendered.push_str("\n\n");
}
}
rendered.push_str(system_prompt);
rendered.push_str("<|im_end|>\n");
}
for (index, message) in messages.iter().enumerate() {
let content = message.content.trim();
if message.system {
rendered.push_str("<|im_start|>system\n");
rendered.push_str(content);
rendered.push_str("<|im_end|>\n");
} else if message.tool {
if index == 0 || !messages[index - 1].tool {
rendered.push_str("<|im_start|>user");
}
rendered.push_str("\n<tool_response>\n");
rendered.push_str(content);
rendered.push_str("\n</tool_response>");
if index + 1 == messages.len() || !messages[index + 1].tool {
rendered.push_str("<|im_end|>\n");
}
} else if message.user {
rendered.push_str("<|im_start|>user\n");
rendered.push_str(content);
rendered.push_str("<|im_end|>\n");
} else {
rendered.push_str("<|im_start|>assistant\n<think>\n");
if let Some(thinking) = &message.reasoning {
rendered.push_str(thinking.trim());
}
rendered.push_str("\n</think>\n\n");
rendered.push_str(content);
rendered.push_str("<|im_end|>\n");
}
}
if continue_assistant {
rendered.push_str("<|im_start|>assistant\n<think>\n");
if reasoning == ReasoningMode::Direct {
rendered.push_str("\n</think>\n\n");
}
}
self.tokenize_rendered(&rendered)
}
pub(super) fn encode_continuation(
&self,
prompt: &str,
reasoning: ReasoningMode,
skip_previous_eos: bool,
) -> Vec<i32> {
if self.family == ModelFamily::Qwen {
let messages = [ChatTurn {
user: true,
tool: false,
system: false,
skip_previous_eos,
reasoning: None,
reasoning_complete: true,
content: prompt.to_owned(),
}];
return self.encode_qwen_messages("", &messages, reasoning, true);
}
let mut output = Vec::new();
if self.family == ModelFamily::DeepSeek && !skip_previous_eos {
output.push(self.eos);
@@ -382,6 +636,7 @@ impl Tokenizer {
pub(super) fn is_stop(&self, token: i32) -> bool {
token == self.eos
|| token == self.alternate_eos
|| (self.family == ModelFamily::Glm
&& [self.system, self.user, self.assistant, self.observation].contains(&token))
}
@@ -394,6 +649,11 @@ impl Tokenizer {
token == self.think_end
}
pub(super) fn is_ple_reset(&self, token: i32) -> bool {
self.family == ModelFamily::Qwen
&& [self.system, self.sop, self.alternate_eos].contains(&token)
}
fn emit_piece(&self, raw: &[u8], output: &mut Vec<i32>) {
if raw.is_empty() {
return;
@@ -512,7 +772,7 @@ impl Tokenizer {
}
}
fn tokenize_glm(&self, text: &str, output: &mut Vec<i32>) {
fn tokenize_gpt4(&self, text: &str, output: &mut Vec<i32>, max_digits: usize) {
let mut position = 0;
while position < text.len() {
let start = position;
@@ -548,7 +808,10 @@ impl Tokenizer {
}
} else if current.number {
let mut digits = 0;
while position < text.len() && char_info(text, position).number && digits < 3 {
while position < text.len()
&& char_info(text, position).number
&& digits < max_digits
{
position = char_info(text, position).next;
digits += 1;
}

View File

@@ -9,7 +9,10 @@ pub(crate) fn validate_model_artifact(
let model = Gguf::open(path)?;
let shape = match expected {
ModelChoice::DeepSeekV4Flash0731 => FLASH_0731,
ModelChoice::DeepSeekV4Pro | ModelChoice::Glm52 | ModelChoice::Glm53Flash => {
ModelChoice::DeepSeekV4Pro
| ModelChoice::Glm52
| ModelChoice::Glm53Flash
| ModelChoice::Qwen38FlashNext => {
return Err(format!("{expected} does not use an external support GGUF"));
}
};
@@ -407,6 +410,7 @@ fn validate_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> {
validate_glm53_tensors(model, shape)
}
ModelFamily::Glm => validate_glm_tensors(model, shape),
ModelFamily::Qwen => unreachable!("Qwen does not use GGUF validation"),
}
}
@@ -1064,7 +1068,7 @@ fn compression_ratio(shape: &Shape, layer: u32) -> u32 {
4
}
ModelChoice::DeepSeekV4Flash0731 | ModelChoice::DeepSeekV4Pro => 128,
ModelChoice::Glm52 | ModelChoice::Glm53Flash => 0,
ModelChoice::Glm52 | ModelChoice::Glm53Flash | ModelChoice::Qwen38FlashNext => 0,
}
}