Integrate Qwen3.8 model intake
This commit is contained in:
@@ -69,6 +69,29 @@ The output directory must not already exist. Materialization uses hard links
|
||||
when possible and otherwise copies bytes, then re-runs the same full verifier
|
||||
against the output.
|
||||
|
||||
## Application intake
|
||||
|
||||
The native Model Manager exposes the complete 29-file set as one resumable
|
||||
artifact. It downloads from the immutable source revision above, verifies each
|
||||
file before promotion, retains interrupted `.part` files, and removes the core,
|
||||
PLE, MTP, tokenizer, metadata, checksum markers, and partials together.
|
||||
|
||||
Model selection is persisted as `qwen3.8-flash-next`; intake recognizes the
|
||||
artifact from the pinned `qwen4_exp` metadata rather than its directory name.
|
||||
Before Metal allocation, the Rust loader checks every configuration and runtime
|
||||
contract in the manifest and all 2,527 safetensors records against the frozen
|
||||
inventory. The tokenizer uses the pinned byte-level BPE, ChatML special tokens,
|
||||
official low/medium/xhigh/direct formatting, grouped tool responses, and the
|
||||
official temperature 1.0, top-p 0.95, and top-k 20 defaults.
|
||||
|
||||
At the native 262,144-token ceiling, admission reports the 71,742,682,599-byte
|
||||
resident core, the separately mapped 32,000,154,008-byte PLE table, the
|
||||
1,672,575,532-byte optional MTP artifact, 6,557,908,992 bytes of KV/recurrent
|
||||
state, and 27,262,976 bytes of prefill transient storage. PLE mapping is not
|
||||
counted as resident allocation, and MTP is counted only when enabled. Execution
|
||||
still fails with an explicit backend-unavailable error until the Rust Metal
|
||||
backend is delivered by issue #95; it never falls through to a GGUF executor.
|
||||
|
||||
## Verification receipt
|
||||
|
||||
On 2026-09-03 the pinned source passed full authentication and inventory,
|
||||
|
||||
108
src/agent.rs
108
src/agent.rs
@@ -3612,7 +3612,9 @@ pub(crate) fn parse_tool_calls(
|
||||
model: ModelChoice,
|
||||
text: &str,
|
||||
) -> Result<(String, Vec<ToolCall>), String> {
|
||||
let (content, calls) = if model.is_glm() {
|
||||
let (content, calls) = if model.is_qwen38() {
|
||||
parse_qwen_calls(text)?
|
||||
} else if model.is_glm() {
|
||||
parse_glm_calls(text)?
|
||||
} else {
|
||||
crate::dsml::parse_tool_calls(text)?
|
||||
@@ -3628,7 +3630,11 @@ pub(crate) fn parse_tool_calls(
|
||||
}
|
||||
|
||||
pub(crate) fn tool_protocol_model(text: &str) -> Option<ModelChoice> {
|
||||
[ModelChoice::DeepSeekV4Flash0731, ModelChoice::Glm52]
|
||||
[
|
||||
ModelChoice::DeepSeekV4Flash0731,
|
||||
ModelChoice::Glm52,
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
]
|
||||
.into_iter()
|
||||
.find(|model| parse_tool_calls(*model, text).is_ok_and(|(_, calls)| !calls.is_empty()))
|
||||
}
|
||||
@@ -3649,7 +3655,11 @@ fn system_prompt_with_tools(
|
||||
ralph_child: bool,
|
||||
) -> String {
|
||||
let schemas = tool_schemas(dev_brain, ralph_child);
|
||||
let tools = if model.is_glm() {
|
||||
let tools = if model.is_qwen38() {
|
||||
format!(
|
||||
"# Tools\n\nYou have access to the following functions:\n\n<tools>\n{schemas}\n</tools>\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>\n\nYou are a coding agent running in a local workspace. Use tools for local file and system work. Avoid printing large file contents or code blocks as answers; edit files with tools, then summarize briefly. Preserve the current system configuration unless explicitly asked otherwise."
|
||||
)
|
||||
} else if model.is_glm() {
|
||||
format!(
|
||||
"You are a coding agent running in a local workspace. Use tools for local file and system work. Avoid printing large file contents or code blocks as answers; edit files with tools, then summarize briefly.\n\n# Tools\n\n<tools>\n{schemas}\n</tools>\n\nFor a function call, output exactly: <tool_call>function-name<arg_key>key</arg_key><arg_value>value</arg_value></tool_call>\nTool calls are not allowed inside <think></think>. Pass numbers and booleans as JSON primitives, not quoted strings. When a tool fails validation or execution, use its code, field, expected, and received feedback to correct the next call. Preserve the current system configuration unless the user explicitly asks otherwise."
|
||||
)
|
||||
@@ -3943,7 +3953,7 @@ fn parse_glm_calls(text: &str) -> Result<(String, Vec<(String, Value)>), String>
|
||||
let body = &rest[..end];
|
||||
let name_end = body.find("<arg_key>").unwrap_or(body.len());
|
||||
let name = body[..name_end].trim();
|
||||
if name.is_empty() {
|
||||
if name.is_empty() || name.contains(['<', '>', '\n', '\r']) {
|
||||
return Err("GLM tool call without function name".into());
|
||||
}
|
||||
let mut arguments = Map::new();
|
||||
@@ -3959,6 +3969,66 @@ fn parse_glm_calls(text: &str) -> Result<(String, Vec<(String, Value)>), String>
|
||||
Ok((text[..visible_len].trim_end().to_owned(), calls))
|
||||
}
|
||||
|
||||
fn parse_qwen_calls(text: &str) -> Result<(String, Vec<(String, Value)>), String> {
|
||||
let scan = text
|
||||
.rfind("</think>")
|
||||
.map_or(text, |position| &text[position + "</think>".len()..]);
|
||||
let Some(first) = scan.find("<tool_call>") else {
|
||||
return Ok((text.to_owned(), Vec::new()));
|
||||
};
|
||||
let visible_len = text.len() - scan.len() + first;
|
||||
let mut rest = &scan[first..];
|
||||
let mut calls = Vec::new();
|
||||
while rest.starts_with("<tool_call>") {
|
||||
rest = rest["<tool_call>".len()..].trim_start();
|
||||
let function = rest
|
||||
.strip_prefix("<function=")
|
||||
.ok_or_else(|| "Qwen tool call is missing its function block".to_owned())?;
|
||||
let name_end = function
|
||||
.find('>')
|
||||
.ok_or_else(|| "Qwen tool call has an incomplete function name".to_owned())?;
|
||||
let name = function[..name_end].trim();
|
||||
if name.is_empty() {
|
||||
return Err("Qwen tool call has an empty function name".into());
|
||||
}
|
||||
let close_function = "</function>";
|
||||
let body = &function[name_end + 1..];
|
||||
let body_end = body
|
||||
.find(close_function)
|
||||
.ok_or_else(|| "Qwen tool call has an incomplete function block".to_owned())?;
|
||||
let mut parameters = body[..body_end].trim();
|
||||
let mut arguments = Map::new();
|
||||
while !parameters.is_empty() {
|
||||
let parameter = parameters
|
||||
.strip_prefix("<parameter=")
|
||||
.ok_or_else(|| "Qwen tool call has malformed parameters".to_owned())?;
|
||||
let key_end = parameter
|
||||
.find('>')
|
||||
.ok_or_else(|| "Qwen tool call has an incomplete parameter name".to_owned())?;
|
||||
let key = parameter[..key_end].trim();
|
||||
if key.is_empty() || arguments.contains_key(key) {
|
||||
return Err(format!("Qwen tool call has an invalid parameter {key}"));
|
||||
}
|
||||
let value = ¶meter[key_end + 1..];
|
||||
let value_end = value
|
||||
.find("</parameter>")
|
||||
.ok_or_else(|| format!("Qwen tool parameter {key} is incomplete"))?;
|
||||
arguments.insert(
|
||||
key.to_owned(),
|
||||
glm_argument(name, key, value[..value_end].trim()),
|
||||
);
|
||||
parameters = value[value_end + "</parameter>".len()..].trim();
|
||||
}
|
||||
calls.push((name.to_owned(), Value::Object(arguments)));
|
||||
rest = body[body_end + close_function.len()..].trim_start();
|
||||
rest = rest
|
||||
.strip_prefix("</tool_call>")
|
||||
.ok_or_else(|| "Qwen tool call is missing its closing tag".to_owned())?
|
||||
.trim_start();
|
||||
}
|
||||
Ok((text[..visible_len].trim_end().to_owned(), calls))
|
||||
}
|
||||
|
||||
fn glm_argument(tool: &str, key: &str, value: &str) -> Value {
|
||||
match tool_spec(tool)
|
||||
.and_then(|tool| {
|
||||
@@ -4239,6 +4309,18 @@ mod tests {
|
||||
assert_eq!(calls[0].name, "read");
|
||||
assert_eq!(calls[0].arguments["path"], "src/main.rs");
|
||||
|
||||
let qwen = "done<tool_call>\n<function=read>\n<parameter=path>\nsrc/main.rs\n</parameter>\n</function>\n</tool_call>";
|
||||
let (visible, calls) = parse_tool_calls(ModelChoice::Qwen38FlashNext, qwen).unwrap();
|
||||
assert_eq!(visible, "done");
|
||||
assert_eq!(calls[0].name, "read");
|
||||
assert_eq!(calls[0].arguments["path"], "src/main.rs");
|
||||
assert_eq!(
|
||||
tool_protocol_model(qwen),
|
||||
Some(ModelChoice::Qwen38FlashNext)
|
||||
);
|
||||
let qwen_prompt = system_prompt(ModelChoice::Qwen38FlashNext, "", false);
|
||||
assert!(qwen_prompt.contains("<function=example_function_name>"));
|
||||
|
||||
let dsml = "done<|DSML|tool_calls><|DSML|invoke name=\"read\"><|DSML|parameter name=\"path\" string=\"true\">src/main.rs</|DSML|parameter></|DSML|invoke></|DSML|tool_calls>";
|
||||
let (visible, calls) = parse_tool_calls(ModelChoice::DeepSeekV4Flash0731, dsml).unwrap();
|
||||
assert_eq!(visible, "done");
|
||||
@@ -4673,6 +4755,18 @@ mod tests {
|
||||
assert_eq!(calls[0].arguments["timeout_sec"], Value::from(3));
|
||||
assert!(validate_tool_call(&calls[0]).is_ok());
|
||||
|
||||
let qwen = "<tool_call>\n<function=bash>\n<parameter=command>\npwd\n</parameter>\n<parameter=timeout_sec>\n3\n</parameter>\n</function>\n</tool_call>";
|
||||
let (_, calls) = parse_tool_calls(ModelChoice::Qwen38FlashNext, qwen).unwrap();
|
||||
assert_eq!(calls[0].arguments["timeout_sec"], Value::from(3));
|
||||
assert!(validate_tool_call(&calls[0]).is_ok());
|
||||
assert!(
|
||||
parse_tool_calls(
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
"<tool_call><function=read><parameter=path>x"
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
|
||||
let dsml = "<|DSML|tool_calls><|DSML|invoke name=\"read\"><|DSML|parameter name=\"path\" string=\"true\">README.md</|DSML|parameter><|DSML|parameter name=\"whole\" string=\"false\">true</|DSML|parameter></|DSML|invoke></|DSML|tool_calls>";
|
||||
let (_, calls) = parse_tool_calls(ModelChoice::DeepSeekV4Flash0731, dsml).unwrap();
|
||||
assert_eq!(calls[0].arguments["whole"], true);
|
||||
@@ -4702,7 +4796,11 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn ralph_report_schemas_are_child_only_and_parse_for_both_models() {
|
||||
for model in [ModelChoice::DeepSeekV4Flash0731, ModelChoice::Glm52] {
|
||||
for model in [
|
||||
ModelChoice::DeepSeekV4Flash0731,
|
||||
ModelChoice::Glm52,
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
] {
|
||||
let parent = system_prompt(model, "", false);
|
||||
let child = ralph_system_prompt(model, "", false);
|
||||
assert!(parent.contains("\"name\":\"ralph\""));
|
||||
|
||||
@@ -72,7 +72,7 @@ impl Default for Config {
|
||||
.reasoning_modes()
|
||||
.iter()
|
||||
.copied()
|
||||
.map(|mode| (mode, GenerationPreferences::default()))
|
||||
.map(|mode| (mode, GenerationPreferences::defaults_for(model)))
|
||||
.collect(),
|
||||
)
|
||||
})
|
||||
@@ -427,7 +427,9 @@ impl Config {
|
||||
.or_insert_with(|| ModelPreferences::defaults_for(model));
|
||||
let profiles = self.generation_profiles.entry(model).or_default();
|
||||
for &mode in model.reasoning_modes() {
|
||||
profiles.entry(mode).or_default();
|
||||
profiles
|
||||
.entry(mode)
|
||||
.or_insert_with(|| GenerationPreferences::defaults_for(model));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -579,6 +581,31 @@ mod tests {
|
||||
assert!(invalid.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen_selection_and_native_defaults_survive_restart() {
|
||||
let directory =
|
||||
std::env::temp_dir().join(format!("ds4-config-qwen-{}", std::process::id()));
|
||||
let path = directory.join("config.yaml");
|
||||
let mut config = Config {
|
||||
model: ModelChoice::Qwen38FlashNext,
|
||||
..Config::default()
|
||||
};
|
||||
config
|
||||
.model_profiles
|
||||
.get_mut(&ModelChoice::Qwen38FlashNext)
|
||||
.unwrap()
|
||||
.reasoning_mode = ReasoningMode::Medium;
|
||||
config.save(&path).unwrap();
|
||||
let loaded = Config::load(&path).unwrap();
|
||||
assert_eq!(loaded.model, ModelChoice::Qwen38FlashNext);
|
||||
assert_eq!(
|
||||
loaded.reasoning_mode(ModelChoice::Qwen38FlashNext),
|
||||
ReasoningMode::Medium
|
||||
);
|
||||
assert_eq!(loaded.active_generation().context_tokens, 131_072);
|
||||
fs::remove_dir_all(directory).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_0731_dspark_opt_out_survives_the_default() {
|
||||
let directory =
|
||||
|
||||
@@ -3,6 +3,7 @@ mod gguf;
|
||||
mod kvstore;
|
||||
#[cfg(target_os = "macos")]
|
||||
mod metal;
|
||||
mod qwen;
|
||||
mod tokenizer;
|
||||
mod validation;
|
||||
|
||||
@@ -82,6 +83,7 @@ unsafe extern "C" {
|
||||
fn munmap(address: *mut std::ffi::c_void, length: usize) -> i32;
|
||||
}
|
||||
|
||||
pub(crate) use qwen::validate_artifacts as validate_qwen_artifacts;
|
||||
pub(crate) use validation::{validate_model_artifact, validate_vision_artifact};
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
@@ -98,6 +100,7 @@ const DSPARK_DENSE: &[u32] = &[F16, F32, Q8_0];
|
||||
enum ModelFamily {
|
||||
DeepSeek,
|
||||
Glm,
|
||||
Qwen,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
@@ -316,6 +319,20 @@ impl Model {
|
||||
settings.speculative.dspark,
|
||||
&settings.artifacts,
|
||||
)?;
|
||||
if settings.model.is_qwen38() {
|
||||
let context = u32::try_from(settings.context_tokens)
|
||||
.map_err(|_| "Qwen context must be a positive whole number")?;
|
||||
let loaded = qwen::load(&settings.artifacts.model, context, false)?;
|
||||
return Err(format!(
|
||||
"Qwen3.8 artifacts are valid, but its Rust Metal execution backend is not available until issue #95. Memory plan: {} resident core bytes, {} mapped PLE bytes, {} optional MTP bytes, {} KV/recurrent bytes, {} prefill transient bytes, {} admitted bytes.",
|
||||
loaded.memory.resident_core,
|
||||
loaded.memory.mapped_ple,
|
||||
loaded.memory.optional_mtp,
|
||||
loaded.memory.kv_and_recurrent,
|
||||
loaded.memory.prefill_transient,
|
||||
loaded.memory.admission,
|
||||
));
|
||||
}
|
||||
let mut model = Self::open_main(&settings.artifacts.model, settings.model)?;
|
||||
if settings.execution.warm_weights {
|
||||
model.main.warm()?;
|
||||
@@ -1717,6 +1734,8 @@ fn conversation_key(system: &str, reasoning: ReasoningMode, messages: &[ChatTurn
|
||||
ReasoningMode::Low => 1,
|
||||
ReasoningMode::High => 2,
|
||||
ReasoningMode::Max => 3,
|
||||
ReasoningMode::Medium => 4,
|
||||
ReasoningMode::XHigh => 5,
|
||||
});
|
||||
for message in messages {
|
||||
output.push(u8::from(message.user));
|
||||
|
||||
@@ -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
646
src/engine/qwen.rs
Normal 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()
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -848,6 +848,7 @@ fn model_code(model: ModelChoice) -> u8 {
|
||||
ModelChoice::Glm52 => 3,
|
||||
ModelChoice::DeepSeekV4Flash0731 => 4,
|
||||
ModelChoice::Glm53Flash => 5,
|
||||
ModelChoice::Qwen38FlashNext => 6,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -858,6 +859,7 @@ fn model_name(value: u8) -> &'static str {
|
||||
3 => "GLM 5.2",
|
||||
4 => "DeepSeek V4 Flash 0731",
|
||||
5 => "GLM 5.3 Flash Q2",
|
||||
6 => "Qwen3.8 Flash Next Bare Speed",
|
||||
_ => "No model loaded",
|
||||
}
|
||||
}
|
||||
|
||||
59
src/model.rs
59
src/model.rs
@@ -1,3 +1,4 @@
|
||||
mod qwen;
|
||||
mod transfer;
|
||||
|
||||
pub(crate) use transfer::{
|
||||
@@ -9,19 +10,21 @@ use std::fmt;
|
||||
use std::fs;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
pub(crate) const MODEL_CHOICES: [ModelChoice; 4] = [
|
||||
pub(crate) const MODEL_CHOICES: [ModelChoice; 5] = [
|
||||
ModelChoice::DeepSeekV4Flash0731,
|
||||
ModelChoice::DeepSeekV4Pro,
|
||||
ModelChoice::Glm52,
|
||||
ModelChoice::Glm53Flash,
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
];
|
||||
pub(crate) const MANAGED_ARTIFACTS: [ManagedArtifactId; 6] = [
|
||||
pub(crate) const MANAGED_ARTIFACTS: [ManagedArtifactId; 7] = [
|
||||
ManagedArtifactId::DeepSeekV4Flash0731,
|
||||
ManagedArtifactId::DeepSeekV4Flash0731Dspark,
|
||||
ManagedArtifactId::DeepSeekV4Pro,
|
||||
ManagedArtifactId::Glm52,
|
||||
ManagedArtifactId::Glm53Flash,
|
||||
ManagedArtifactId::Glm53FlashVision,
|
||||
ManagedArtifactId::Qwen38FlashNext,
|
||||
];
|
||||
|
||||
const DEEPSEEK_REPOSITORY: &str = "antirez/deepseek-v4-gguf";
|
||||
@@ -94,6 +97,8 @@ pub(crate) enum ModelChoice {
|
||||
Glm52,
|
||||
#[serde(rename = "glm-5.3-flash")]
|
||||
Glm53Flash,
|
||||
#[serde(rename = "qwen3.8-flash-next")]
|
||||
Qwen38FlashNext,
|
||||
}
|
||||
|
||||
impl ModelChoice {
|
||||
@@ -103,6 +108,7 @@ impl ModelChoice {
|
||||
Self::DeepSeekV4Pro => "deepseek-v4-pro",
|
||||
Self::Glm52 => "glm-5.2",
|
||||
Self::Glm53Flash => "glm-5.3-flash",
|
||||
Self::Qwen38FlashNext => "qwen3.8-flash-next",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -118,13 +124,21 @@ impl ModelChoice {
|
||||
matches!(self, Self::Glm52 | Self::Glm53Flash)
|
||||
}
|
||||
|
||||
pub(crate) fn is_qwen38(self) -> bool {
|
||||
self == Self::Qwen38FlashNext
|
||||
}
|
||||
|
||||
pub(crate) fn supports_glm_mtp(self) -> bool {
|
||||
self.is_glm()
|
||||
}
|
||||
|
||||
pub(crate) fn main_artifact_size(self) -> u64 {
|
||||
if self.is_qwen38() {
|
||||
qwen::core_bytes()
|
||||
} else {
|
||||
self.main_artifact().size
|
||||
}
|
||||
}
|
||||
|
||||
fn main_artifact(self) -> &'static Artifact {
|
||||
match self {
|
||||
@@ -132,13 +146,14 @@ impl ModelChoice {
|
||||
Self::DeepSeekV4Pro => &PRO,
|
||||
Self::Glm52 => &GLM,
|
||||
Self::Glm53Flash => &GLM53_FLASH,
|
||||
Self::Qwen38FlashNext => unreachable!("Qwen uses a pinned artifact set"),
|
||||
}
|
||||
}
|
||||
|
||||
fn dspark_artifact(self) -> Option<&'static Artifact> {
|
||||
match self {
|
||||
Self::DeepSeekV4Flash0731 => Some(&FLASH_0731_DSPARK),
|
||||
Self::DeepSeekV4Pro | Self::Glm52 | Self::Glm53Flash => None,
|
||||
Self::DeepSeekV4Pro | Self::Glm52 | Self::Glm53Flash | Self::Qwen38FlashNext => None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -170,6 +185,13 @@ pub(crate) fn engine_artifacts(
|
||||
dspark_enabled: bool,
|
||||
models_path: &Path,
|
||||
) -> EngineArtifacts {
|
||||
if model.is_qwen38() {
|
||||
return EngineArtifacts {
|
||||
model: qwen::root(models_path),
|
||||
support: None,
|
||||
vision: None,
|
||||
};
|
||||
}
|
||||
EngineArtifacts {
|
||||
model: model.main_artifact().path(model, models_path),
|
||||
support: if dspark_enabled {
|
||||
@@ -195,6 +217,12 @@ pub(crate) fn validate_engine_artifacts(
|
||||
if dspark_enabled && !model.supports_dspark() {
|
||||
return Err(format!("DSpark is not compatible with {model}"));
|
||||
}
|
||||
if model.is_qwen38() {
|
||||
if artifacts.support.is_some() || artifacts.vision.is_some() {
|
||||
return Err("Qwen3.8 does not accept GGUF support or vision artifacts".into());
|
||||
}
|
||||
return qwen::validate_installed(&artifacts.model);
|
||||
}
|
||||
model
|
||||
.main_artifact()
|
||||
.validate_installed_path(&artifacts.model)?;
|
||||
@@ -229,6 +257,7 @@ pub(crate) enum ManagedArtifactId {
|
||||
Glm52,
|
||||
Glm53Flash,
|
||||
Glm53FlashVision,
|
||||
Qwen38FlashNext,
|
||||
}
|
||||
|
||||
impl ManagedArtifactId {
|
||||
@@ -240,6 +269,7 @@ impl ManagedArtifactId {
|
||||
Self::DeepSeekV4Pro => ModelChoice::DeepSeekV4Pro,
|
||||
Self::Glm52 => ModelChoice::Glm52,
|
||||
Self::Glm53Flash | Self::Glm53FlashVision => ModelChoice::Glm53Flash,
|
||||
Self::Qwen38FlashNext => ModelChoice::Qwen38FlashNext,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -251,15 +281,20 @@ impl ManagedArtifactId {
|
||||
Self::Glm52 => &GLM,
|
||||
Self::Glm53Flash => &GLM53_FLASH,
|
||||
Self::Glm53FlashVision => &GLM53_FLASH_VISION,
|
||||
Self::Qwen38FlashNext => unreachable!("Qwen uses a pinned artifact set"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for ManagedArtifactId {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
if *self == Self::Qwen38FlashNext {
|
||||
formatter.write_str(qwen::LABEL)
|
||||
} else {
|
||||
formatter.write_str(self.artifact().label)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub(crate) enum ManagedArtifactState {
|
||||
@@ -346,6 +381,7 @@ impl fmt::Display for ModelChoice {
|
||||
Self::DeepSeekV4Pro => "DeepSeek V4 Pro 0813",
|
||||
Self::Glm52 => "GLM 5.2",
|
||||
Self::Glm53Flash => "GLM 5.3 Flash Q2",
|
||||
Self::Qwen38FlashNext => "Qwen3.8 Flash Next Bare Speed",
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -444,6 +480,9 @@ pub(crate) fn managed_artifacts(models_path: &Path) -> Vec<ManagedArtifact> {
|
||||
MANAGED_ARTIFACTS
|
||||
.into_iter()
|
||||
.map(|id| {
|
||||
if id == ManagedArtifactId::Qwen38FlashNext {
|
||||
return qwen::managed_artifact(models_path);
|
||||
}
|
||||
let model = id.model();
|
||||
let artifact = id.artifact();
|
||||
let stored = artifact.stored_bytes(model, models_path);
|
||||
@@ -474,7 +513,13 @@ pub(crate) fn managed_artifacts(models_path: &Path) -> Vec<ManagedArtifact> {
|
||||
pub(crate) fn installed_models(models_path: &Path) -> Vec<ModelChoice> {
|
||||
MODEL_CHOICES
|
||||
.into_iter()
|
||||
.filter(|model| model.main_artifact().is_installed(*model, models_path))
|
||||
.filter(|model| {
|
||||
if model.is_qwen38() {
|
||||
qwen::is_installed(&qwen::root(models_path))
|
||||
} else {
|
||||
model.main_artifact().is_installed(*model, models_path)
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
@@ -482,6 +527,9 @@ pub(crate) fn artifact_download_progress(
|
||||
id: ManagedArtifactId,
|
||||
models_path: &Path,
|
||||
) -> DownloadProgress {
|
||||
if id == ManagedArtifactId::Qwen38FlashNext {
|
||||
return qwen::download_progress(models_path);
|
||||
}
|
||||
let model = id.model();
|
||||
let artifact = id.artifact();
|
||||
let downloaded = artifact.downloaded_bytes(model, models_path);
|
||||
@@ -512,6 +560,9 @@ pub(crate) fn artifact_verification_progress(
|
||||
id: ManagedArtifactId,
|
||||
verified: u64,
|
||||
) -> DownloadProgress {
|
||||
if id == ManagedArtifactId::Qwen38FlashNext {
|
||||
return qwen::verification_progress(verified);
|
||||
}
|
||||
let artifact = id.artifact();
|
||||
DownloadProgress {
|
||||
downloaded: artifact.size,
|
||||
|
||||
338
src/model/qwen.rs
Normal file
338
src/model/qwen.rs
Normal file
@@ -0,0 +1,338 @@
|
||||
use super::{
|
||||
DownloadPhase, DownloadProgress, ManagedArtifact, ManagedArtifactId, ManagedArtifactState,
|
||||
VerificationProgress,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::fs::{self, File};
|
||||
use std::io::Read;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::OnceLock;
|
||||
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
|
||||
|
||||
pub(super) const LABEL: &str = "Qwen3.8 Flash Next Bare Speed artifact set";
|
||||
pub(super) const REPOSITORY: &str = "Youssofal/Qwen3.8-Flash-Next-MTPLX-Bare-Speed";
|
||||
pub(super) const REVISION: &str = "74559cdf34fbfc0b593de72d17e93f37fd4f9ea7";
|
||||
const MANIFEST_SHA256: &str = "eeec490fd3d0c0b1c9093be0a7fe7b9be4b49389ed5276de7530c2d3c11290ca";
|
||||
const MANIFEST_BYTES: &[u8] =
|
||||
include_bytes!("../../assets/models/qwen38-flash-next-bare-speed.json");
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Manifest {
|
||||
format: u32,
|
||||
source: Source,
|
||||
files: Vec<Artifact>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Source {
|
||||
repository: String,
|
||||
revision: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(super) struct Artifact {
|
||||
pub(super) path: String,
|
||||
pub(super) role: String,
|
||||
pub(super) size: u64,
|
||||
pub(super) sha256: String,
|
||||
}
|
||||
|
||||
pub(super) fn root(models_path: &Path) -> PathBuf {
|
||||
models_path.join("qwen3.8-flash-next")
|
||||
}
|
||||
|
||||
fn manifest() -> &'static Manifest {
|
||||
static MANIFEST: OnceLock<Manifest> = OnceLock::new();
|
||||
MANIFEST.get_or_init(|| {
|
||||
let digest = hex(Sha256::digest(MANIFEST_BYTES));
|
||||
assert_eq!(digest, MANIFEST_SHA256, "embedded Qwen manifest changed");
|
||||
let manifest: Manifest = serde_json::from_slice(MANIFEST_BYTES)
|
||||
.expect("embedded Qwen manifest must be valid JSON");
|
||||
assert_eq!(manifest.format, 1, "unsupported embedded Qwen manifest");
|
||||
assert_eq!(manifest.source.repository, REPOSITORY);
|
||||
assert_eq!(manifest.source.revision, REVISION);
|
||||
assert_eq!(manifest.files.len(), 29);
|
||||
manifest
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn artifacts() -> &'static [Artifact] {
|
||||
&manifest().files
|
||||
}
|
||||
|
||||
pub(super) fn total_bytes() -> u64 {
|
||||
artifacts().iter().map(|artifact| artifact.size).sum()
|
||||
}
|
||||
|
||||
pub(super) fn core_bytes() -> u64 {
|
||||
artifacts()
|
||||
.iter()
|
||||
.filter(|artifact| artifact.role == "core")
|
||||
.map(|artifact| artifact.size)
|
||||
.sum()
|
||||
}
|
||||
|
||||
pub(super) fn url(artifact: &Artifact) -> String {
|
||||
format!(
|
||||
"https://huggingface.co/{REPOSITORY}/resolve/{REVISION}/{}",
|
||||
artifact.path
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn path(root: &Path, artifact: &Artifact) -> PathBuf {
|
||||
root.join(&artifact.path)
|
||||
}
|
||||
|
||||
pub(super) fn partial_path(root: &Path, artifact: &Artifact) -> PathBuf {
|
||||
root.join(format!("{}.part", artifact.path))
|
||||
}
|
||||
|
||||
pub(super) fn verification_path(root: &Path, artifact: &Artifact) -> PathBuf {
|
||||
root.join(format!("{}.sha256", artifact.path))
|
||||
}
|
||||
|
||||
pub(super) fn artifact_is_installed(root: &Path, artifact: &Artifact) -> bool {
|
||||
path(root, artifact)
|
||||
.metadata()
|
||||
.is_ok_and(|metadata| metadata.len() == artifact.size)
|
||||
&& fs::read_to_string(verification_path(root, artifact))
|
||||
.is_ok_and(|digest| digest.trim() == artifact.sha256)
|
||||
}
|
||||
|
||||
pub(super) fn is_installed(root: &Path) -> bool {
|
||||
artifacts()
|
||||
.iter()
|
||||
.all(|artifact| artifact_is_installed(root, artifact))
|
||||
}
|
||||
|
||||
pub(super) fn validate_installed(root: &Path) -> Result<(), String> {
|
||||
for artifact in artifacts() {
|
||||
let artifact_path = path(root, artifact);
|
||||
let size = artifact_path
|
||||
.metadata()
|
||||
.map_err(|error| format!("Could not inspect {}: {error}", artifact_path.display()))?
|
||||
.len();
|
||||
if size != artifact.size {
|
||||
return Err(format!(
|
||||
"{} has {size} bytes, expected {}",
|
||||
artifact_path.display(),
|
||||
artifact.size
|
||||
));
|
||||
}
|
||||
let marker = fs::read_to_string(verification_path(root, artifact)).map_err(|_| {
|
||||
format!(
|
||||
"{} has not passed checksum verification",
|
||||
artifact_path.display()
|
||||
)
|
||||
})?;
|
||||
if marker.trim() != artifact.sha256 {
|
||||
return Err(format!(
|
||||
"{} has the wrong artifact identity",
|
||||
artifact_path.display()
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn stored_bytes(root: &Path) -> u64 {
|
||||
artifacts()
|
||||
.iter()
|
||||
.flat_map(|artifact| [path(root, artifact), partial_path(root, artifact)])
|
||||
.filter_map(|path| path.metadata().ok())
|
||||
.map(|metadata| metadata.len())
|
||||
.sum()
|
||||
}
|
||||
|
||||
fn downloaded_bytes(root: &Path) -> u64 {
|
||||
artifacts()
|
||||
.iter()
|
||||
.map(|artifact| {
|
||||
if artifact_is_installed(root, artifact) {
|
||||
artifact.size
|
||||
} else {
|
||||
path(root, artifact)
|
||||
.metadata()
|
||||
.or_else(|_| partial_path(root, artifact).metadata())
|
||||
.map_or(0, |metadata| metadata.len().min(artifact.size))
|
||||
}
|
||||
})
|
||||
.sum()
|
||||
}
|
||||
|
||||
fn complete_payloads(root: &Path) -> bool {
|
||||
artifacts().iter().all(|artifact| {
|
||||
path(root, artifact)
|
||||
.metadata()
|
||||
.or_else(|_| partial_path(root, artifact).metadata())
|
||||
.is_ok_and(|metadata| metadata.len() == artifact.size)
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn managed_artifact(models_path: &Path) -> ManagedArtifact {
|
||||
let root = root(models_path);
|
||||
let stored = stored_bytes(&root);
|
||||
let state = if is_installed(&root) {
|
||||
ManagedArtifactState::Ready
|
||||
} else if complete_payloads(&root) {
|
||||
ManagedArtifactState::NeedsVerification
|
||||
} else if stored > 0 {
|
||||
ManagedArtifactState::Partial
|
||||
} else {
|
||||
ManagedArtifactState::Missing
|
||||
};
|
||||
ManagedArtifact {
|
||||
id: ManagedArtifactId::Qwen38FlashNext,
|
||||
stored,
|
||||
expected: total_bytes(),
|
||||
state,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn download_progress(models_path: &Path) -> DownloadProgress {
|
||||
let root = root(models_path);
|
||||
let downloaded = downloaded_bytes(&root);
|
||||
let installed = is_installed(&root);
|
||||
let verification = (!installed && complete_payloads(&root)).then_some(VerificationProgress {
|
||||
verified: 0,
|
||||
total: total_bytes(),
|
||||
});
|
||||
let phase = if installed {
|
||||
DownloadPhase::Complete
|
||||
} else if verification.is_some() {
|
||||
DownloadPhase::Verifying(LABEL)
|
||||
} else if downloaded > 0 {
|
||||
DownloadPhase::Downloading(LABEL)
|
||||
} else {
|
||||
DownloadPhase::Pending(LABEL)
|
||||
};
|
||||
DownloadProgress {
|
||||
downloaded,
|
||||
total: total_bytes(),
|
||||
phase,
|
||||
verification,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn verification_progress(verified: u64) -> DownloadProgress {
|
||||
let total = total_bytes();
|
||||
DownloadProgress {
|
||||
downloaded: total,
|
||||
total,
|
||||
phase: DownloadPhase::Verifying(LABEL),
|
||||
verification: Some(VerificationProgress {
|
||||
verified: verified.min(total),
|
||||
total,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn verify(
|
||||
path: &Path,
|
||||
artifact: &Artifact,
|
||||
cancel: &AtomicBool,
|
||||
verified_bytes: &AtomicU64,
|
||||
) -> Result<bool, String> {
|
||||
let size = path
|
||||
.metadata()
|
||||
.map_err(|error| format!("{}: {error}", path.display()))?
|
||||
.len();
|
||||
if size != artifact.size {
|
||||
return Err(format!(
|
||||
"{} has size {size}, expected {}",
|
||||
path.display(),
|
||||
artifact.size
|
||||
));
|
||||
}
|
||||
let mut file = File::open(path).map_err(|error| format!("{}: {error}", path.display()))?;
|
||||
let mut hash = Sha256::new();
|
||||
let mut buffer = vec![0_u8; 1024 * 1024];
|
||||
loop {
|
||||
if cancel.load(Ordering::Relaxed) {
|
||||
return Ok(false);
|
||||
}
|
||||
let count = file
|
||||
.read(&mut buffer)
|
||||
.map_err(|error| format!("{}: {error}", path.display()))?;
|
||||
if count == 0 {
|
||||
break;
|
||||
}
|
||||
hash.update(&buffer[..count]);
|
||||
verified_bytes.fetch_add(count as u64, Ordering::Relaxed);
|
||||
}
|
||||
if hex(hash.finalize()) != artifact.sha256 {
|
||||
return Err(format!(
|
||||
"Checksum verification failed for {}",
|
||||
path.display()
|
||||
));
|
||||
}
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
fn hex(bytes: impl AsRef<[u8]>) -> String {
|
||||
let bytes = bytes.as_ref();
|
||||
let mut result = String::with_capacity(bytes.len() * 2);
|
||||
for byte in bytes {
|
||||
use std::fmt::Write as _;
|
||||
write!(result, "{byte:02x}").expect("writing to String cannot fail");
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
#[test]
|
||||
fn embedded_manifest_is_the_frozen_complete_set() {
|
||||
assert_eq!(artifacts().len(), 29);
|
||||
assert_eq!(total_bytes(), 105_438_842_430);
|
||||
assert_eq!(core_bytes(), 71_742_682_599);
|
||||
assert_eq!(
|
||||
artifacts()
|
||||
.iter()
|
||||
.filter(|artifact| artifact.role == "ple")
|
||||
.map(|artifact| artifact.size)
|
||||
.sum::<u64>(),
|
||||
32_000_154_008
|
||||
);
|
||||
assert_eq!(
|
||||
artifacts()
|
||||
.iter()
|
||||
.filter(|artifact| artifact.role == "mtp")
|
||||
.map(|artifact| artifact.size)
|
||||
.sum::<u64>(),
|
||||
1_672_575_532
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn verification_rejects_hash_mismatch_and_honors_cancellation() {
|
||||
let root = std::env::temp_dir().join(format!(
|
||||
"ds4-qwen-verify-{}-{}",
|
||||
std::process::id(),
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos()
|
||||
));
|
||||
fs::create_dir_all(&root).unwrap();
|
||||
let path = root.join("fixture");
|
||||
fs::write(&path, b"qwen").unwrap();
|
||||
let artifact = Artifact {
|
||||
path: "fixture".into(),
|
||||
role: "core".into(),
|
||||
size: 4,
|
||||
sha256: hex(Sha256::digest(b"qwen")),
|
||||
};
|
||||
let progress = AtomicU64::new(0);
|
||||
assert!(verify(&path, &artifact, &AtomicBool::new(false), &progress).unwrap());
|
||||
assert_eq!(progress.load(Ordering::Relaxed), 4);
|
||||
fs::write(&path, b"fail").unwrap();
|
||||
assert!(verify(&path, &artifact, &AtomicBool::new(false), &progress).is_err());
|
||||
assert!(!verify(&path, &artifact, &AtomicBool::new(true), &progress).unwrap());
|
||||
fs::remove_dir_all(root).unwrap();
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,9 @@ pub(crate) fn download_managed_artifact(
|
||||
cancel: &AtomicBool,
|
||||
verified_bytes: &AtomicU64,
|
||||
) -> Result<DownloadOutcome, String> {
|
||||
if id == ManagedArtifactId::Qwen38FlashNext {
|
||||
return download_qwen(models_path, cancel, verified_bytes);
|
||||
}
|
||||
download_artifact_with_cancel(
|
||||
id.model(),
|
||||
id.artifact(),
|
||||
@@ -29,6 +32,9 @@ pub(crate) fn validate_managed_artifact(
|
||||
cancel: &AtomicBool,
|
||||
verified_bytes: &AtomicU64,
|
||||
) -> Result<DownloadOutcome, String> {
|
||||
if id == ManagedArtifactId::Qwen38FlashNext {
|
||||
return validate_qwen(models_path, cancel, verified_bytes);
|
||||
}
|
||||
let model = id.model();
|
||||
let artifact = id.artifact();
|
||||
let destination = artifact.path(model, models_path);
|
||||
@@ -67,6 +73,9 @@ pub(crate) fn delete_managed_artifact(
|
||||
id: ManagedArtifactId,
|
||||
models_path: &Path,
|
||||
) -> Result<(), String> {
|
||||
if id == ManagedArtifactId::Qwen38FlashNext {
|
||||
return delete_qwen(models_path);
|
||||
}
|
||||
let model = id.model();
|
||||
let artifact = id.artifact();
|
||||
for path in [
|
||||
@@ -83,6 +92,144 @@ pub(crate) fn delete_managed_artifact(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn download_qwen(
|
||||
models_path: &Path,
|
||||
cancel: &AtomicBool,
|
||||
verified_bytes: &AtomicU64,
|
||||
) -> Result<DownloadOutcome, String> {
|
||||
let root = qwen::root(models_path);
|
||||
fs::create_dir_all(&root).map_err(|error| format!("{}: {error}", root.display()))?;
|
||||
verified_bytes.store(0, Ordering::Relaxed);
|
||||
for artifact in qwen::artifacts() {
|
||||
if cancel.load(Ordering::Relaxed) {
|
||||
return Ok(DownloadOutcome::Stopped);
|
||||
}
|
||||
if qwen::artifact_is_installed(&root, artifact) {
|
||||
verified_bytes.fetch_add(artifact.size, Ordering::Relaxed);
|
||||
continue;
|
||||
}
|
||||
let destination = qwen::path(&root, artifact);
|
||||
let partial = qwen::partial_path(&root, artifact);
|
||||
let source = if destination.exists() {
|
||||
destination.clone()
|
||||
} else {
|
||||
if partial.metadata().map_or(0, |metadata| metadata.len()) > artifact.size {
|
||||
File::create(&partial)
|
||||
.map_err(|error| format!("{}: {error}", partial.display()))?;
|
||||
}
|
||||
if partial.metadata().map_or(0, |metadata| metadata.len()) != artifact.size
|
||||
&& download_url_to_partial(&qwen::url(artifact), &partial, cancel)?
|
||||
== DownloadOutcome::Stopped
|
||||
{
|
||||
return Ok(DownloadOutcome::Stopped);
|
||||
}
|
||||
partial.clone()
|
||||
};
|
||||
match qwen::verify(&source, artifact, cancel, verified_bytes) {
|
||||
Ok(false) => return Ok(DownloadOutcome::Stopped),
|
||||
Ok(true) => {}
|
||||
Err(error) => {
|
||||
let _ = fs::remove_file(qwen::verification_path(&root, artifact));
|
||||
if source == partial {
|
||||
fs::remove_file(&partial).map_err(|remove_error| {
|
||||
format!("{error}; could not remove partial file: {remove_error}")
|
||||
})?;
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
if source == partial {
|
||||
fs::rename(&partial, &destination)
|
||||
.map_err(|error| format!("{}: {error}", destination.display()))?;
|
||||
}
|
||||
fs::write(qwen::verification_path(&root, artifact), &artifact.sha256)
|
||||
.map_err(|error| error.to_string())?;
|
||||
}
|
||||
if let Err(error) = crate::engine::validate_qwen_artifacts(&root) {
|
||||
clear_qwen_markers(&root)?;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(DownloadOutcome::Complete)
|
||||
}
|
||||
|
||||
fn validate_qwen(
|
||||
models_path: &Path,
|
||||
cancel: &AtomicBool,
|
||||
verified_bytes: &AtomicU64,
|
||||
) -> Result<DownloadOutcome, String> {
|
||||
let root = qwen::root(models_path);
|
||||
verified_bytes.store(0, Ordering::Relaxed);
|
||||
for artifact in qwen::artifacts() {
|
||||
if !qwen::path(&root, artifact).exists() && !qwen::partial_path(&root, artifact).exists() {
|
||||
return Err(format!("{} is not downloaded", artifact.path));
|
||||
}
|
||||
}
|
||||
for artifact in qwen::artifacts() {
|
||||
let destination = qwen::path(&root, artifact);
|
||||
let partial = qwen::partial_path(&root, artifact);
|
||||
let (source, promote) = if destination.exists() {
|
||||
(destination.clone(), false)
|
||||
} else if partial.exists() {
|
||||
(partial.clone(), true)
|
||||
} else {
|
||||
return Err(format!("{} is not downloaded", artifact.path));
|
||||
};
|
||||
match qwen::verify(&source, artifact, cancel, verified_bytes) {
|
||||
Ok(false) => return Ok(DownloadOutcome::Stopped),
|
||||
Ok(true) => {}
|
||||
Err(error) => {
|
||||
let _ = fs::remove_file(qwen::verification_path(&root, artifact));
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
if promote {
|
||||
fs::rename(&partial, &destination)
|
||||
.map_err(|error| format!("{}: {error}", destination.display()))?;
|
||||
}
|
||||
fs::write(qwen::verification_path(&root, artifact), &artifact.sha256)
|
||||
.map_err(|error| error.to_string())?;
|
||||
}
|
||||
if let Err(error) = crate::engine::validate_qwen_artifacts(&root) {
|
||||
clear_qwen_markers(&root)?;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(DownloadOutcome::Complete)
|
||||
}
|
||||
|
||||
fn clear_qwen_markers(root: &Path) -> Result<(), String> {
|
||||
for artifact in qwen::artifacts() {
|
||||
let path = qwen::verification_path(root, artifact);
|
||||
match fs::remove_file(&path) {
|
||||
Ok(()) => {}
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
|
||||
Err(error) => return Err(format!("{}: {error}", path.display())),
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn delete_qwen(models_path: &Path) -> Result<(), String> {
|
||||
let root = qwen::root(models_path);
|
||||
for artifact in qwen::artifacts() {
|
||||
for path in [
|
||||
qwen::path(&root, artifact),
|
||||
qwen::partial_path(&root, artifact),
|
||||
qwen::verification_path(&root, artifact),
|
||||
] {
|
||||
match fs::remove_file(&path) {
|
||||
Ok(()) => {}
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
|
||||
Err(error) => return Err(format!("{}: {error}", path.display())),
|
||||
}
|
||||
}
|
||||
}
|
||||
match fs::remove_dir(&root) {
|
||||
Ok(()) => Ok(()),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
|
||||
Err(error) => Err(format!("{}: {error}", root.display())),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn download_artifact(
|
||||
model: ModelChoice,
|
||||
@@ -332,8 +479,12 @@ mod tests {
|
||||
Some(ModelChoice::DeepSeekV4Flash0731)
|
||||
);
|
||||
assert!(ModelChoice::from_id("unknown").is_none());
|
||||
assert_eq!(MODEL_CHOICES.len(), 4);
|
||||
assert_eq!(MANAGED_ARTIFACTS.len(), 6);
|
||||
assert_eq!(MODEL_CHOICES.len(), 5);
|
||||
assert_eq!(MANAGED_ARTIFACTS.len(), 7);
|
||||
assert_eq!(
|
||||
ModelChoice::from_id("qwen3.8-flash-next"),
|
||||
Some(ModelChoice::Qwen38FlashNext)
|
||||
);
|
||||
assert_eq!(ModelChoice::Glm52.main_artifact().size, 211_075_856_448);
|
||||
assert_eq!(
|
||||
ModelChoice::DeepSeekV4Flash0731.main_artifact().size,
|
||||
@@ -342,6 +493,10 @@ mod tests {
|
||||
assert_eq!(ModelChoice::DeepSeekV4Flash0731.artifacts(true).count(), 2);
|
||||
assert_eq!(ModelChoice::Glm52.artifacts(true).count(), 1);
|
||||
assert_eq!(ModelChoice::Glm53Flash.artifacts(true).count(), 2);
|
||||
assert_eq!(
|
||||
engine_artifacts(ModelChoice::Qwen38FlashNext, false, Path::new("/models")).model,
|
||||
Path::new("/models/qwen3.8-flash-next")
|
||||
);
|
||||
|
||||
let id = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
@@ -417,6 +572,52 @@ mod tests {
|
||||
fs::remove_dir_all(models_path).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires DS4_QWEN38_ARTIFACTS to point at the pinned 105 GB source"]
|
||||
fn qwen_set_verifies_reloads_and_removes_as_one_managed_artifact() {
|
||||
let source = std::env::var_os("DS4_QWEN38_ARTIFACTS")
|
||||
.map(std::path::PathBuf::from)
|
||||
.expect("DS4_QWEN38_ARTIFACTS is set");
|
||||
let models = std::env::temp_dir().join(format!(
|
||||
"ds4-qwen-managed-{}-{}",
|
||||
std::process::id(),
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos()
|
||||
));
|
||||
let root = qwen::root(&models);
|
||||
fs::create_dir_all(&root).unwrap();
|
||||
for artifact in qwen::artifacts() {
|
||||
fs::hard_link(source.join(&artifact.path), qwen::path(&root, artifact)).unwrap();
|
||||
}
|
||||
fs::remove_file(qwen::path(
|
||||
&root,
|
||||
qwen::artifacts()
|
||||
.iter()
|
||||
.find(|artifact| artifact.role == "mtp")
|
||||
.unwrap(),
|
||||
))
|
||||
.unwrap();
|
||||
assert!(
|
||||
validate_qwen(&models, &AtomicBool::new(false), &AtomicU64::new(0))
|
||||
.unwrap_err()
|
||||
.contains("mtp.safetensors is not downloaded")
|
||||
);
|
||||
fs::hard_link(source.join("mtp.safetensors"), root.join("mtp.safetensors")).unwrap();
|
||||
let verified = AtomicU64::new(0);
|
||||
assert_eq!(
|
||||
validate_qwen(&models, &AtomicBool::new(false), &verified).unwrap(),
|
||||
DownloadOutcome::Complete
|
||||
);
|
||||
assert_eq!(verified.load(Ordering::Relaxed), qwen::total_bytes());
|
||||
assert!(qwen::is_installed(&root));
|
||||
assert!(crate::engine::validate_qwen_artifacts(&root).is_ok());
|
||||
delete_qwen(&models).unwrap();
|
||||
assert!(!root.exists());
|
||||
fs::remove_dir_all(models).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn verification_reports_bytes_read() {
|
||||
let path = std::env::temp_dir().join(format!(
|
||||
|
||||
@@ -563,6 +563,9 @@ fn model_alias(id: &str) -> Option<ModelChoice> {
|
||||
| "glm-5.3-flash-reasoner"
|
||||
| "zai/glm-5.3-flash"
|
||||
| "zai/glm-5.3-flash-reasoner" => Some(ModelChoice::Glm53Flash),
|
||||
"qwen3.8-flash-next" | "qwen/qwen3.8-flash-next" | "mtplx-flash-next-bare-speed" => {
|
||||
Some(ModelChoice::Qwen38FlashNext)
|
||||
}
|
||||
_ => ModelChoice::from_id(id),
|
||||
}
|
||||
}
|
||||
@@ -1027,6 +1030,10 @@ mod tests {
|
||||
Some(ModelChoice::DeepSeekV4Flash0731)
|
||||
);
|
||||
assert_eq!(model_alias("glm-5.3-flash-chat"), None);
|
||||
assert_eq!(
|
||||
model_alias("mtplx-flash-next-bare-speed"),
|
||||
Some(ModelChoice::Qwen38FlashNext)
|
||||
);
|
||||
let model = model_json(
|
||||
"deepseek-reasoner",
|
||||
ModelChoice::DeepSeekV4Flash0731,
|
||||
@@ -1047,6 +1054,15 @@ mod tests {
|
||||
model_json("glm-5.3-flash", ModelChoice::Glm53Flash, 32_768, 50_000)["supported_reasoning_efforts"],
|
||||
json!(["low", "high", "max"])
|
||||
);
|
||||
assert_eq!(
|
||||
model_json(
|
||||
"qwen3.8-flash-next",
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
131_072,
|
||||
50_000
|
||||
)["supported_reasoning_efforts"],
|
||||
json!(["low", "medium", "xhigh", "none"])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -603,8 +603,10 @@ fn request_reasoning(
|
||||
let explicit_effort = match request.reasoning_effort.as_deref() {
|
||||
Some("none") => Some(ReasoningMode::Direct),
|
||||
Some("low") => Some(ReasoningMode::Low),
|
||||
Some("medium") => Some(ReasoningMode::Medium),
|
||||
Some("high") => Some(ReasoningMode::High),
|
||||
Some("max") => Some(ReasoningMode::Max),
|
||||
Some("xhigh") => Some(ReasoningMode::XHigh),
|
||||
None => None,
|
||||
Some(value) => return Err((400, format!("unsupported reasoning_effort: {value}"))),
|
||||
};
|
||||
@@ -679,6 +681,10 @@ mod tests {
|
||||
request_reasoning(&default, ModelChoice::Glm53Flash, "").unwrap(),
|
||||
ReasoningMode::Max
|
||||
);
|
||||
assert_eq!(
|
||||
request_reasoning(&default, ModelChoice::Qwen38FlashNext, "").unwrap(),
|
||||
ReasoningMode::XHigh
|
||||
);
|
||||
|
||||
for (effort, mode) in [
|
||||
("none", ReasoningMode::Direct),
|
||||
@@ -698,6 +704,32 @@ mod tests {
|
||||
}
|
||||
|
||||
assert!(request_reasoning(&request(Some("medium"), None), ModelChoice::Glm52, "").is_err());
|
||||
assert_eq!(
|
||||
request_reasoning(
|
||||
&request(Some("medium"), None),
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
""
|
||||
)
|
||||
.unwrap(),
|
||||
ReasoningMode::Medium
|
||||
);
|
||||
assert_eq!(
|
||||
request_reasoning(
|
||||
&request(Some("xhigh"), None),
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
""
|
||||
)
|
||||
.unwrap(),
|
||||
ReasoningMode::XHigh
|
||||
);
|
||||
assert!(
|
||||
request_reasoning(
|
||||
&request(Some("high"), None),
|
||||
ModelChoice::Qwen38FlashNext,
|
||||
""
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
assert!(
|
||||
request_reasoning(&request(None, Some(false)), ModelChoice::Glm53Flash, "").is_err()
|
||||
);
|
||||
|
||||
@@ -16,6 +16,12 @@ const GLM_52_REASONING_MODES: [ReasoningMode; 3] = [
|
||||
];
|
||||
const GLM_53_REASONING_MODES: [ReasoningMode; 3] =
|
||||
[ReasoningMode::Low, ReasoningMode::High, ReasoningMode::Max];
|
||||
const QWEN_REASONING_MODES: [ReasoningMode; 4] = [
|
||||
ReasoningMode::Low,
|
||||
ReasoningMode::Medium,
|
||||
ReasoningMode::XHigh,
|
||||
ReasoningMode::Direct,
|
||||
];
|
||||
const MAX_CPU_THREADS: u32 = 32;
|
||||
pub(crate) const GIB: u64 = 1024 * 1024 * 1024;
|
||||
/// DS4 disk KV cache defaults, from `ds4_kvstore.h` and `--kv-disk-space-mb`.
|
||||
@@ -493,8 +499,10 @@ pub(crate) enum ReasoningMode {
|
||||
Direct,
|
||||
#[default]
|
||||
Low,
|
||||
Medium,
|
||||
High,
|
||||
Max,
|
||||
XHigh,
|
||||
}
|
||||
|
||||
impl fmt::Display for ReasoningMode {
|
||||
@@ -502,8 +510,10 @@ impl fmt::Display for ReasoningMode {
|
||||
formatter.write_str(match self {
|
||||
Self::Direct => "Direct (no thinking)",
|
||||
Self::Low => "Think Low",
|
||||
Self::Medium => "Think Medium",
|
||||
Self::High => "Think High",
|
||||
Self::Max => "Think Max",
|
||||
Self::XHigh => "Think XHigh",
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -513,8 +523,10 @@ impl ReasoningMode {
|
||||
match self {
|
||||
Self::Direct => "none",
|
||||
Self::Low => "low",
|
||||
Self::Medium => "medium",
|
||||
Self::High => "high",
|
||||
Self::Max => "max",
|
||||
Self::XHigh => "xhigh",
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -525,6 +537,7 @@ impl ModelChoice {
|
||||
Self::DeepSeekV4Flash0731 | Self::DeepSeekV4Pro => &DEEPSEEK_REASONING_MODES,
|
||||
Self::Glm52 => &GLM_52_REASONING_MODES,
|
||||
Self::Glm53Flash => &GLM_53_REASONING_MODES,
|
||||
Self::Qwen38FlashNext => &QWEN_REASONING_MODES,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -532,6 +545,7 @@ impl ModelChoice {
|
||||
match self {
|
||||
Self::DeepSeekV4Flash0731 | Self::DeepSeekV4Pro => ReasoningMode::Low,
|
||||
Self::Glm52 | Self::Glm53Flash => ReasoningMode::Max,
|
||||
Self::Qwen38FlashNext => ReasoningMode::XHigh,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -581,6 +595,13 @@ impl Default for GenerationPreferences {
|
||||
}
|
||||
|
||||
impl GenerationPreferences {
|
||||
pub(crate) fn defaults_for(model: ModelChoice) -> Self {
|
||||
Self {
|
||||
context_tokens: if model.is_qwen38() { 131_072 } else { 32_768 },
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn validate(&self) -> Result<(), String> {
|
||||
if self.context_tokens <= 0 {
|
||||
return Err("Context tokens must be a positive whole number.".into());
|
||||
@@ -603,15 +624,16 @@ impl GenerationPreferences {
|
||||
kv_cache: KvCacheSettings,
|
||||
) -> TurnSettings {
|
||||
let glm = model.is_glm();
|
||||
let qwen = model.is_qwen38();
|
||||
TurnSettings {
|
||||
kv_cache,
|
||||
context_tokens: self.context_tokens,
|
||||
max_generated_tokens: self.max_generated_tokens,
|
||||
system_prompt: self.system_prompt.clone(),
|
||||
temperature: self.temperature.unwrap_or(1.0),
|
||||
top_p: self.top_p.unwrap_or(if glm { 0.95 } else { 1.0 }),
|
||||
min_p: self.min_p.unwrap_or(if glm { 0.0 } else { 0.05 }),
|
||||
top_k: 0,
|
||||
top_p: self.top_p.unwrap_or(if glm || qwen { 0.95 } else { 1.0 }),
|
||||
min_p: self.min_p.unwrap_or(if glm || qwen { 0.0 } else { 0.05 }),
|
||||
top_k: if qwen { 20 } else { 0 },
|
||||
stops: Vec::new(),
|
||||
seed: self.seed,
|
||||
reasoning_mode: self.reasoning_mode,
|
||||
@@ -675,6 +697,9 @@ pub(crate) fn effective_settings(
|
||||
models_path: &Path,
|
||||
) -> Result<EffectiveSettings, String> {
|
||||
generation.validate()?;
|
||||
if model.is_qwen38() && generation.context_tokens > 262_144 {
|
||||
return Err("Qwen context cannot exceed its native 262144-token limit.".into());
|
||||
}
|
||||
model.validate_reasoning_mode(generation.reasoning_mode)?;
|
||||
Ok(EffectiveSettings {
|
||||
engine: runtime.engine_settings(model, generation.context_tokens, models_path)?,
|
||||
@@ -701,6 +726,27 @@ mod tests {
|
||||
|
||||
let glm = defaults.turn_settings(ModelChoice::Glm52, cache);
|
||||
assert_eq!((glm.temperature, glm.top_p, glm.min_p), (1.0, 0.95, 0.0));
|
||||
let qwen = GenerationPreferences::defaults_for(ModelChoice::Qwen38FlashNext);
|
||||
let qwen_turn = qwen.turn_settings(ModelChoice::Qwen38FlashNext, cache);
|
||||
assert_eq!(qwen.context_tokens, 131_072);
|
||||
assert_eq!(
|
||||
ModelChoice::Qwen38FlashNext.reasoning_modes(),
|
||||
&[
|
||||
ReasoningMode::Low,
|
||||
ReasoningMode::Medium,
|
||||
ReasoningMode::XHigh,
|
||||
ReasoningMode::Direct
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
(
|
||||
qwen_turn.temperature,
|
||||
qwen_turn.top_p,
|
||||
qwen_turn.min_p,
|
||||
qwen_turn.top_k
|
||||
),
|
||||
(1.0, 0.95, 0.0, 20)
|
||||
);
|
||||
assert_eq!(defaults.reasoning_mode, ReasoningMode::Low);
|
||||
|
||||
let explicit = GenerationPreferences {
|
||||
|
||||
Reference in New Issue
Block a user