From bf82df77cbbb3e0bce4d52f4e51002ca727caee8 Mon Sep 17 00:00:00 2001 From: Georg Bauer Date: Thu, 3 Sep 2026 22:38:52 +0200 Subject: [PATCH] Add native Qwen MTP speculation --- .../qwen38-flash-next-bare-speed-tensors.tsv | 18 +- .../models/qwen38-flash-next-bare-speed.json | 4 +- src/app/preferences.rs | 14 +- src/app/view/preferences.rs | 16 +- src/app/view/stats.rs | 17 + src/engine.rs | 108 +- src/engine/metal.rs | 35 +- src/engine/metal/qwen.rs | 1071 ++++++++++++++++- src/engine/qwen.rs | 52 +- src/metrics.rs | 23 +- src/model.rs | 4 +- src/model/qwen.rs | 2 +- src/settings.rs | 8 +- tools/qwen38-artifacts.rs | 8 +- 14 files changed, 1302 insertions(+), 78 deletions(-) diff --git a/assets/models/qwen38-flash-next-bare-speed-tensors.tsv b/assets/models/qwen38-flash-next-bare-speed-tensors.tsv index db342d9..fa85d57 100644 --- a/assets/models/qwen38-flash-next-bare-speed-tensors.tsv +++ b/assets/models/qwen38-flash-next-bare-speed-tensors.tsv @@ -2492,15 +2492,15 @@ mtp.safetensors mtp.layers.0.mlp.shared_expert.up_proj.weight U32 640x320 4 64 a mtp.safetensors mtp.layers.0.mlp.shared_expert_gate.biases BF16 1x40 4 64 affine 140045196 140045276 mtp.safetensors mtp.layers.0.mlp.shared_expert_gate.scales BF16 1x40 4 64 affine 140045276 140045356 mtp.safetensors mtp.layers.0.mlp.shared_expert_gate.weight U32 1x320 4 64 affine 620915756 620917036 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.down_proj.biases BF16 512x2560x20 2 64 affine 140966956 193395756 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.down_proj.scales BF16 512x2560x20 2 64 affine 641468716 693897516 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.down_proj.weight U32 512x2560x80 2 64 affine 1253145132 1672575532 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.gate_proj.biases BF16 512x640x80 2 64 affine 693897772 746326572 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.gate_proj.scales BF16 512x640x80 2 64 affine 20277900 72706700 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.gate_proj.weight U32 512x640x320 2 64 affine 760314412 1179744812 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.up_proj.biases BF16 512x640x80 2 64 affine 87616396 140045196 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.up_proj.scales BF16 512x640x80 2 64 affine 1192852012 1245280812 -mtp.safetensors mtp.layers.0.mlp.switch_mlp.up_proj.weight U32 512x640x320 2 64 affine 194808876 614239276 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.down_proj.biases BF16 512x2560x20 4 32 affine 140966956 193395756 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.down_proj.scales BF16 512x2560x20 4 32 affine 641468716 693897516 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.down_proj.weight U32 512x2560x80 4 32 affine 1253145132 1672575532 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.gate_proj.biases BF16 512x640x80 4 32 affine 693897772 746326572 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.gate_proj.scales BF16 512x640x80 4 32 affine 20277900 72706700 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.gate_proj.weight U32 512x640x320 4 32 affine 760314412 1179744812 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.up_proj.biases BF16 512x640x80 4 32 affine 87616396 140045196 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.up_proj.scales BF16 512x640x80 4 32 affine 1192852012 1245280812 +mtp.safetensors mtp.layers.0.mlp.switch_mlp.up_proj.weight U32 512x640x320 4 32 affine 194808876 614239276 mtp.safetensors mtp.layers.0.mlp_hyper_connection.block_inject_weight.weight BF16 4x10240 74345100 74427020 mtp.safetensors mtp.layers.0.mlp_hyper_connection.hc_norm.weight BF16 10240 752921132 752941612 mtp.safetensors mtp.layers.0.mlp_hyper_connection.input_mix_weight_down.weight BF16 320x10240 81062540 87616140 diff --git a/assets/models/qwen38-flash-next-bare-speed.json b/assets/models/qwen38-flash-next-bare-speed.json index 3f584b8..ab411e7 100644 --- a/assets/models/qwen38-flash-next-bare-speed.json +++ b/assets/models/qwen38-flash-next-bare-speed.json @@ -11,7 +11,7 @@ "converter": "qwen38-artifacts-v1-identity" }, "tensor_inventory": "qwen38-flash-next-bare-speed-tensors.tsv", - "tensor_inventory_sha256": "369cdc7e53c09eaa5eff5cdcc49f0d0eab0ddc175f34965ba2ed1be2673d5858", + "tensor_inventory_sha256": "b5731e6febcf865d276a0e7b144da02375f2d7e3129594f271ee9c6c351f4c8f", "config": { "/architectures/0": "Qwen4ExpForConditionalGeneration", "/model_type": "qwen4_exp", @@ -73,7 +73,7 @@ {"class":"qsa","file":"model-00016-of-00017.safetensors","tensor":"language_model.model.layers.11.self_attn.indexer.index_qk_proj.weight","row":0,"values":64,"bits":8,"group_size":64,"sha256":"c803ea77148621a5d6dfa4060a52c75f57185234a0bec4771671f69336bb8346"}, {"class":"gdn","file":"model-00016-of-00017.safetensors","tensor":"language_model.model.layers.0.linear_attn.A_log","row":0,"values":48,"sha256":"88c53a2a04bda1d96ee1ade6fa7dfa9c49d3245b1cc1e83673f02b7a36e07a85"}, {"class":"ple","file":"ngram-table.safetensors","tensor":"ngram.weight","row":0,"values":64,"bits":4,"group_size":32,"sha256":"2243dd9766046bb80d98e3baf5e59ba246958cfe4f340ebce8c1fff57a2810d9"}, - {"class":"mtp","file":"mtp.safetensors","tensor":"mtp.layers.0.mlp.switch_mlp.gate_proj.weight","row":0,"values":64,"bits":2,"group_size":64,"sha256":"6293cce5e1b25a4f35cb8af46e6e13b8378ec4ca84ec7f68859066dda48049cc"} + {"class":"mtp","file":"mtp.safetensors","tensor":"mtp.layers.0.mlp.switch_mlp.gate_proj.weight","row":0,"values":64,"bits":4,"group_size":32,"sha256":"393d82675ae7e25275243fd4907d3894146ac1c5662b601aae998c489d10af5c"} ], "files": [ {"path":"model-00001-of-00017.safetensors","role":"core","size":4666167150,"sha256":"a27232c9434b9d8961f198cf36f44e228b3f23a5f8d874ba4ca8af32a1b23ffc"}, diff --git a/src/app/preferences.rs b/src/app/preferences.rs index 6d9e34b..58b84e2 100644 --- a/src/app/preferences.rs +++ b/src/app/preferences.rs @@ -865,16 +865,22 @@ impl App { self.preference_error = None; } Message::PreferenceGlmMtpChanged(value) => { - self.preference_draft.glm_mtp = - self.preference_draft.acceleration_model.supports_glm_mtp() && value; + self.preference_draft.glm_mtp = self + .preference_draft + .acceleration_model + .supports_integrated_mtp() + && value; if !self.preference_draft.glm_mtp { self.preference_draft.glm_mtp_timing = false; } self.preference_error = None; } Message::PreferenceGlmMtpTimingChanged(value) => { - self.preference_draft.glm_mtp_timing = - self.preference_draft.acceleration_model.supports_glm_mtp() && value; + self.preference_draft.glm_mtp_timing = self + .preference_draft + .acceleration_model + .supports_integrated_mtp() + && value; if self.preference_draft.glm_mtp_timing { self.preference_draft.glm_mtp = true; } diff --git a/src/app/view/preferences.rs b/src/app/view/preferences.rs index 6ca4978..6dad994 100644 --- a/src/app/view/preferences.rs +++ b/src/app/view/preferences.rs @@ -17,12 +17,12 @@ impl App { let glm_mtp_toggle: Option Message> = self .preference_draft .acceleration_model - .supports_glm_mtp() + .supports_integrated_mtp() .then_some(Message::PreferenceGlmMtpChanged); let glm_mtp_timing_toggle: Option Message> = self .preference_draft .acceleration_model - .supports_glm_mtp() + .supports_integrated_mtp() .then_some(Message::PreferenceGlmMtpTimingChanged); let keep_vision_loaded_toggle: Option Message> = (self.preference_draft.acceleration_model == ModelChoice::Glm53Flash) @@ -608,13 +608,13 @@ impl App { text("SPECULATIVE DECODING").size(11).color(muted_text()), hint( toggle(self.preference_draft.glm_mtp) - .label("Enable integrated GLM MTP") + .label("Enable integrated MTP") .on_toggle_maybe(glm_mtp_toggle), - "Uses the prediction head built into GLM for speculative decoding, so no separate draft model is loaded.", + "Uses the selected model's managed prediction head for speculative decoding. Qwen loads its pinned MTP sidecar; GLM uses its embedded head.", ), hint( toggle(self.preference_draft.glm_mtp_timing) - .label("Log GLM MTP timing counters") + .label("Log MTP timing counters") .on_toggle_maybe(glm_mtp_timing_toggle), "Records per-stage timings of the speculative path to the log, to show where the acceleration actually goes. A diagnostic aid that costs a little throughput.", ), @@ -644,8 +644,8 @@ impl App { ), text(if self.preference_draft.acceleration_model.supports_dspark() { "DeepSeek V4 Flash 0731 uses its managed DSpark support artifact." - } else if self.preference_draft.acceleration_model.supports_glm_mtp() { - "GLM MTP is integrated; DSpark is unavailable for this model." + } else if self.preference_draft.acceleration_model.supports_integrated_mtp() { + "Integrated MTP is available; DSpark is unavailable for this model." } else { "No speculative-decoding support is available for this model." }) @@ -656,7 +656,7 @@ impl App { |engine| { let settings = engine.speculative; format!( - "Engine: GLM MTP {} • timing {} • DSpark {} • confidence {}{} • target-only {} • exact sampling {}", + "Engine: integrated MTP {} • timing {} • DSpark {} • confidence {}{} • target-only {} • exact sampling {}", if settings.glm_mtp { "on" } else { "off" }, if settings.glm_mtp_timing { "on" } else { "off" }, if settings.dspark { "on" } else { "off" }, diff --git a/src/app/view/stats.rs b/src/app/view/stats.rs index 4082513..f0e79d6 100644 --- a/src/app/view/stats.rs +++ b/src/app/view/stats.rs @@ -369,6 +369,7 @@ impl App { match stats.speculative_mode { 2 => "DSpark", 3 => "GLM MTP", + 4 => "Qwen MTP", _ => "Off", }, ), @@ -376,11 +377,27 @@ impl App { metric_row("Drafted", format_count(stats.drafted_tokens)), metric_row("Accepted", format_count(stats.accepted_draft_tokens)), metric_row("Acceptance", format!("{:.1}%", acceptance * 100.0)), + metric_row( + "Mean accepted depth", + if stats.speculative_cycles == 0 { + "0.00".into() + } else { + format!( + "{:.2}", + stats.accepted_depth_total as f64 / stats.speculative_cycles as f64 + ) + } + ), + metric_row( + "Maximum accepted depth", + format_count(stats.accepted_depth_max) + ), metric_row( "Target verifier passes", format_count(stats.verifier_passes) ), metric_row("Verifier wall time", format_milliseconds(stats.verifier_ms)), + metric_row("Repair wall time", format_milliseconds(stats.repair_ms)), metric_row( "Effective target-pass speedup", format!("{effective_speedup:.2}×") diff --git a/src/engine.rs b/src/engine.rs index 41fe7af..6927349 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -329,9 +329,13 @@ impl LoadedModel { )?; let context = u32::try_from(settings.context_tokens) .map_err(|_| "Qwen context must be a positive whole number")?; - qwen::QwenModel::open(&settings.artifacts.model, context) - .map(Box::new) - .map(Self::Qwen) + qwen::QwenModel::open_configured( + &settings.artifacts.model, + context, + settings.speculative.glm_mtp, + ) + .map(Box::new) + .map(Self::Qwen) } else { Model::open(settings).map(Box::new).map(Self::Gguf) } @@ -815,8 +819,11 @@ impl Generator { stats.speculative_cycles, stats.drafted_tokens, stats.accepted_draft_tokens, + stats.accepted_depth_total, + stats.accepted_depth_max, stats.verifier_passes, stats.verifier_ms, + stats.repair_ms, ); self.metrics.ssd_stats(SsdStats { enabled: stats.ssd_enabled, @@ -1988,7 +1995,7 @@ fn sample_probabilities( .sum(); let mut choice = rng.unit() * total; for (token, probability) in probabilities { - if Some(*token) == excluded { + if Some(*token) == excluded || !probability.is_finite() || *probability <= 0.0 { continue; } choice -= probability; @@ -1999,7 +2006,9 @@ fn sample_probabilities( probabilities .iter() .rev() - .find(|(token, _)| Some(*token) != excluded) + .find(|(token, probability)| { + Some(*token) != excluded && probability.is_finite() && *probability > 0.0 + }) .map_or(0, |(token, _)| *token as i32) } @@ -2028,6 +2037,67 @@ fn exact_delta_sample( ) } +#[cfg(any(target_os = "macos", test))] +fn sample_from_logits( + logits: &[f32], + temperature: f32, + top_p: f32, + min_p: f32, + top_k: i32, + rng: &mut Rng, +) -> i32 { + sample_probabilities( + &sampling_probabilities(logits, temperature, top_p, min_p, top_k), + rng, + None, + ) +} + +#[cfg(any(target_os = "macos", test))] +#[allow(clippy::too_many_arguments)] +fn exact_speculative_sample( + target_logits: &[f32], + draft_logits: &[f32], + draft: i32, + temperature: f32, + top_p: f32, + min_p: f32, + top_k: i32, + rng: &mut Rng, +) -> (i32, bool) { + let mut target = sampling_probabilities(target_logits, temperature, top_p, min_p, top_k); + let mut proposal = sampling_probabilities(draft_logits, temperature, top_p, min_p, top_k); + target.sort_unstable_by_key(|(token, _)| *token); + proposal.sort_unstable_by_key(|(token, _)| *token); + let probability = |values: &[(usize, f32)], token: usize| { + values + .binary_search_by_key(&token, |(candidate, _)| *candidate) + .ok() + .map_or(0.0, |index| values[index].1) + }; + let target_probability = probability(&target, draft as usize); + let draft_probability = probability(&proposal, draft as usize); + if draft_probability > 0.0 && rng.unit() <= (target_probability / draft_probability).min(1.0) { + return (draft, true); + } + + let target_best = target + .iter() + .max_by(|left, right| left.1.total_cmp(&right.1)) + .map_or(draft, |(token, _)| *token as i32); + let mut residual = Vec::with_capacity(target.len()); + for (token, target_probability) in target { + let remaining = target_probability - probability(&proposal, token); + if remaining > 0.0 { + residual.push((token, remaining)); + } + } + if residual.is_empty() { + return (target_best, false); + } + (sample_probabilities(&residual, rng, None), false) +} + #[cfg(any(target_os = "macos", test))] struct Rng(u64); @@ -2113,6 +2183,34 @@ mod sampling_tests { assert_eq!((replacement, was_draft), (1, false)); } + #[test] + fn probability_ratio_speculation_preserves_the_target_distribution() { + let target = [0.75_f32.ln(), 0.25_f32.ln()]; + let proposal = [0.25_f32.ln(), 0.75_f32.ln()]; + let mut rng = Rng::new(0x51a7_1c5e); + let mut counts = [0_u32; 2]; + for _ in 0..40_000 { + let draft = sample_from_logits(&proposal, 1.0, 1.0, 0.0, 0, &mut rng); + let (token, _) = + exact_speculative_sample(&target, &proposal, draft, 1.0, 1.0, 0.0, 0, &mut rng); + counts[token as usize] += 1; + } + let observed = counts[0] as f32 / counts.iter().sum::() as f32; + assert!((observed - 0.75).abs() < 0.015, "observed {observed}"); + + let run = |seed| { + let mut rng = Rng::new(seed); + (0..32) + .map(|_| { + let draft = sample_from_logits(&proposal, 1.0, 1.0, 0.0, 0, &mut rng); + exact_speculative_sample(&target, &proposal, draft, 1.0, 1.0, 0.0, 0, &mut rng) + }) + .collect::>() + }; + assert_eq!(run(77), run(77)); + assert!(run(77).iter().any(|(_, accepted)| !accepted)); + } + #[test] fn sampling_probabilities_match_ds4_filter_order() { let probabilities = sampling_probabilities(&[0.0, 2.0, 1.0], 1.0, 0.8, 0.2, 0); diff --git a/src/engine/metal.rs b/src/engine/metal.rs index 147835a..9e8649c 100644 --- a/src/engine/metal.rs +++ b/src/engine/metal.rs @@ -14,7 +14,10 @@ use qwen::QwenExecutor; use super::gguf::{BF16, F16, F32, Gguf, IQ2_XXS, MXFP4, Q2_K, Q4_K, Q8_0, Tensor as GgufTensor}; use super::validation::{DsparkConfig, SupportKind, dspark_config}; -use super::{LoadedModel, Model, ModelFamily, ModelRef, Rng, exact_delta_sample}; +use super::{ + LoadedModel, Model, ModelFamily, ModelRef, Rng, exact_delta_sample, exact_speculative_sample, + sample_from_logits, +}; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode, @@ -2949,8 +2952,11 @@ pub(super) struct ExecutionStats { pub(super) speculative_cycles: u64, pub(super) drafted_tokens: u64, pub(super) accepted_draft_tokens: u64, + pub(super) accepted_depth_total: u64, + pub(super) accepted_depth_max: u64, pub(super) verifier_passes: u64, pub(super) verifier_ms: u64, + pub(super) repair_ms: u64, pub(super) ssd_enabled: bool, pub(super) ssd_resident_bytes: u64, pub(super) ssd_cache_bytes: u64, @@ -4460,9 +4466,11 @@ impl Executor { .map(Self::Glm) } LoadedModel::Gguf(_) => unreachable!("Qwen never uses a GGUF model"), - LoadedModel::Qwen(model) => QwenExecutor::open(*model, context) - .map(Box::new) - .map(Self::Qwen), + LoadedModel::Qwen(model) => { + { QwenExecutor::open_configured(*model, context, speculative) } + .map(Box::new) + .map(Self::Qwen) + } } } @@ -4495,10 +4503,17 @@ impl Executor { executor.eval(token)?; Ok(vec![token]) } - Self::Qwen(executor) => { - executor.eval(token)?; - Ok(vec![token]) - } + Self::Qwen(executor) => executor.eval_speculative_sampled( + token, + max_tokens, + reasoning, + temperature, + top_p, + min_p, + top_k, + rng, + cancelled, + ), } } @@ -4548,9 +4563,7 @@ impl Executor { executor.eval_speculative_greedy(token, max_tokens, cancelled) } Self::Qwen(executor) => { - let _ = (max_tokens, reasoning, cancelled); - executor.eval(token)?; - Ok(vec![token]) + executor.eval_speculative_greedy(token, max_tokens, reasoning, cancelled) } } } diff --git a/src/engine/metal/qwen.rs b/src/engine/metal/qwen.rs index fd4ad67..0ee1570 100644 --- a/src/engine/metal/qwen.rs +++ b/src/engine/metal/qwen.rs @@ -26,6 +26,7 @@ const QSA_WIDTH: u32 = (QSA_HEADS + QSA_KV_HEADS) * QSA_DIM; const QSA_RATIO: u32 = 4; const QSA_TOP_K: u32 = 512; const MTP_TOKEN_RESERVE: u32 = 3; +const MTP_CAPTURE_ROWS: usize = MTP_TOKEN_RESERVE as usize + 1; const EXPERTS: u32 = 512; const EXPERTS_USED: usize = 10; const EXPERT_WIDTH: u32 = 640; @@ -38,7 +39,7 @@ const PLE_CONV_STATE: u32 = 9; const EOS_TOKEN: i32 = 248_044; const PLE_ROW_BYTES: usize = 100; const CHECKPOINT_MAGIC: &[u8; 8] = b"DS4QWN01"; -const CHECKPOINT_VERSION: u32 = 3; +const CHECKPOINT_VERSION: u32 = 4; const CHECKPOINT_CHUNK: usize = 8 * 1024 * 1024; #[derive(Clone, Copy)] @@ -171,6 +172,35 @@ struct PleState { conv: Buffer, } +struct RecurrentSnapshot { + gdn: Vec>, + ple_conv: Buffer, + ple_history: [i32; PLE_HISTORY], + hidden: Buffer, + logits: Vec, + position: u32, + tokens: usize, +} + +struct QwenMtp { + attention: LayerState, + captures: Vec, + timing: bool, + cycles: u64, + drafted: u64, + accepted: u64, + accepted_depth_total: u64, + accepted_depth_max: u64, + verifier_passes: u64, + verifier_ns: u64, + repair_ns: u64, +} + +struct QwenDraft { + token: i32, + logits: Vec, +} + pub(in crate::engine) struct QwenExecutor { model: QwenModel, states: Vec, @@ -182,6 +212,7 @@ pub(in crate::engine) struct QwenExecutor { position: u32, context: u32, checkpoint_tag: [u8; 32], + mtp: Option, _context: Context, } @@ -192,13 +223,27 @@ pub(in crate::engine) struct QwenResidentState { tokens: Vec, position: u32, checkpoint_tag: [u8; 32], + mtp: Option, } impl QwenExecutor { + #[cfg(test)] pub(super) fn open(model: QwenModel, context: u32) -> Result { + Self::open_configured(model, context, EngineSpeculativeSettings::default()) + } + + pub(super) fn open_configured( + model: QwenModel, + context: u32, + speculative: EngineSpeculativeSettings, + ) -> Result { let ple_contract = ple_contract(&model)?; let native = Context::open_qwen(model.memory().admission)?; let states = allocate_states(context)?; + let mtp = speculative + .glm_mtp + .then(|| allocate_mtp(context, speculative.glm_mtp_timing)) + .transpose()?; Ok(Self { model, states, @@ -210,11 +255,19 @@ impl QwenExecutor { position: 0, context, checkpoint_tag: [0; 32], + mtp, _context: native, }) } pub(super) fn eval(&mut self, token: i32) -> Result<(), String> { + if self.position > 0 && self.mtp.is_some() { + let _ = self.mtp_step(token, self.position - 1, false)?; + } + self.eval_target(token) + } + + fn eval_target(&mut self, token: i32) -> Result<(), String> { if token < 0 || token as u32 >= VOCAB { return Err(format!("token {token} is outside the Qwen vocabulary")); } @@ -421,9 +474,17 @@ impl QwenExecutor { self.hyper_read(&format!("{prefix}.attn_hyper_connection"))?; match &self.states[layer] { LayerState::Gdn { .. } => self.gdn(&prefix, layer)?, - LayerState::Attention { .. } => self.attention(&prefix, layer)?, + state @ LayerState::Attention { .. } => { + self.attention(&prefix, state, self.position)? + } } self.hyper_write()?; + commands.finish()?; + self.encode_moe(&prefix, &layer.to_string()) + } + + fn encode_moe(&mut self, prefix: &str, layer_name: &str) -> Result<(), String> { + let commands = Commands::begin()?; self.hyper_read(&format!("{prefix}.mlp_hyper_connection"))?; self.affine_mv_into( &self.affine(&format!("{prefix}.mlp.gate"), HIDDEN, EXPERTS, None)?, @@ -462,12 +523,12 @@ impl QwenExecutor { self.scratch.block.read_f32(&mut block)?; let block_non_finite = block.iter().filter(|value| !value.is_finite()).count(); return Err(format!( - "Qwen layer {layer} router produced invalid expert {expert} with weight {weight} ({non_finite} non-finite logits, {block_non_finite} non-finite inputs)" + "Qwen layer {layer_name} router produced invalid expert {expert} with weight {weight} ({non_finite} non-finite logits, {block_non_finite} non-finite inputs)" )); } - self.expert(&prefix, expert as u32, weight)?; + self.expert(prefix, expert as u32, weight)?; } - self.shared_expert(&prefix)?; + self.shared_expert(prefix)?; let mut inject_args = args(); inject_args.u[0] = HIDDEN; self.dispatch( @@ -703,12 +764,12 @@ impl QwenExecutor { ) } - fn attention(&self, prefix: &str, layer: usize) -> Result<(), String> { + fn attention(&self, prefix: &str, state: &LayerState, position: u32) -> Result<(), String> { let LayerState::Attention { kv, qsa_raw, qsa_pooled, - } = &self.states[layer] + } = state else { return Err("Qwen attention graph received GDN state".into()); }; @@ -727,7 +788,7 @@ impl QwenExecutor { let mut index = args(); index.u[0] = QSA_DIM; index.u[1] = QSA_HEADS; - index.u[2] = self.position; + index.u[2] = position; self.dispatch( c"kernel_qwen_qsa_store_raw", qsa_raw, @@ -744,14 +805,14 @@ impl QwenExecutor { &self.scratch.qsa_qk, self.weight(&format!("{prefix}.self_attn.indexer.q_layernorm.weight"))?, QSA_HEADS, - self.position, + position, )?; - if (self.position + 1).is_multiple_of(QSA_RATIO) { + if (position + 1).is_multiple_of(QSA_RATIO) { let mut pool = args(); pool.u[0] = QSA_DIM; - pool.u[1] = self.position / QSA_RATIO; + pool.u[1] = position / QSA_RATIO; pool.u[2] = QSA_RATIO; - pool.u[3] = self.position + 1 - QSA_RATIO; + pool.u[3] = position + 1 - QSA_RATIO; pool.f[0] = 1.0e-6; pool.f[1] = 10_000_000.0; self.dispatch( @@ -822,16 +883,18 @@ impl QwenExecutor { &self.scratch.attention, self.weight(&format!("{prefix}.self_attn.q_norm.weight"))?, ATTN_HEADS, + position, )?; self.head_norm_rope( &self.scratch.k, &self.scratch.k_rope, self.weight(&format!("{prefix}.self_attn.k_norm.weight"))?, ATTN_KV_HEADS, + position, )?; let mut store = args(); store.u[0] = ATTN_KV_WIDTH; - store.u[1] = self.position; + store.u[1] = position; self.dispatch( c"kernel_qwen_store_kv_bf16", kv, @@ -843,7 +906,7 @@ impl QwenExecutor { ATTN_KV_WIDTH, 1, )?; - self.attend(kv, qsa_pooled)?; + self.attend(kv, qsa_pooled, position)?; let mut gate = args(); gate.u[0] = ATTN_WIDTH; self.dispatch( @@ -899,8 +962,8 @@ impl QwenExecutor { ) } - fn attend(&self, kv: &Buffer, qsa_pooled: &Buffer) -> Result<(), String> { - let tokens = self.position + 1; + fn attend(&self, kv: &Buffer, qsa_pooled: &Buffer, position: u32) -> Result<(), String> { + let tokens = position + 1; let mut values = args(); values.u[0] = ATTN_HEADS; values.u[1] = ATTN_KV_HEADS; @@ -985,12 +1048,13 @@ impl QwenExecutor { output: &Buffer, weight: Weight<'_>, heads: u32, + position: u32, ) -> Result<(), String> { let mut values = args(); values.u[0] = ATTN_DIM; values.u[1] = 64; values.u[2] = heads; - values.u[3] = self.position; + values.u[3] = position; values.f[0] = 1.0e-6; values.f[1] = 10_000_000.0; self.dispatch( @@ -1132,7 +1196,7 @@ impl QwenExecutor { fn final_output(&mut self) -> Result<(), String> { let commands = Commands::begin()?; - self.final_mix()?; + self.mix_output("language_model.model.hyper_connection_mixer")?; self.affine_mv_into( &self.affine("language_model.lm_head", HIDDEN, VOCAB, None)?, &self.scratch.block, @@ -1144,8 +1208,7 @@ impl QwenExecutor { self.scratch.logits.read_f32(&mut self.logits) } - fn final_mix(&self) -> Result<(), String> { - let prefix = "language_model.model.hyper_connection_mixer"; + fn mix_output(&self, prefix: &str) -> Result<(), String> { let norm = self.weight(&format!("{prefix}.hc_norm.weight"))?; let down = self.weight(&format!("{prefix}.input_mix_weight_down.weight"))?; let up = self.weight(&format!("{prefix}.input_mix_weight_up.weight"))?; @@ -1379,6 +1442,582 @@ impl QwenExecutor { dispatch_qwen(kernel, out, a, b, c, weights, values, grid_x, grid_y) } + fn mtp_step( + &mut self, + token: i32, + position: u32, + draft_requested: bool, + ) -> Result>, String> { + let mut mtp = self.mtp.take().ok_or("Qwen MTP is not configured")?; + let started = Instant::now(); + let result = (|| { + if token < 0 || token as u32 >= VOCAB { + return Err(format!("token {token} is outside the Qwen vocabulary")); + } + if position >= self.context + MTP_TOKEN_RESERVE { + return Err("Qwen MTP position exceeds its bounded cache".into()); + } + + let embedding = + self.affine("language_model.model.embed_tokens", HIDDEN, VOCAB, None)?; + let mut embed = args(); + embed.u[0] = HIDDEN; + embed.u[2] = embedding.bits; + embed.u[3] = embedding.group; + embed.u[4] = token as u32; + let commands = Commands::begin()?; + self.dispatch( + c"kernel_qwen_affine_embedding", + &self.scratch.hidden, + None, + None, + None, + &[ + self.view(embedding.packed), + self.view(embedding.scales), + self.view(embedding.biases), + ], + &embed, + HIDDEN, + 1, + )?; + let mut norm = args(); + norm.u[0] = HIDDEN; + norm.u[1] = HIDDEN; + norm.f[0] = 1.0e-6; + self.dispatch( + c"kernel_qwen_zero_rms", + &self.scratch.block, + Some(&self.scratch.hidden), + None, + None, + &[self.view(self.weight("mtp.pre_fc_norm_embedding.weight")?)], + &norm, + 1, + 1, + )?; + self.bf16_mv( + self.weight("mtp.fc_embedding.weight")?, + &self.scratch.block, + &self.scratch.hidden, + HIDDEN, + HIDDEN, + )?; + norm.u[0] = HC_WIDTH; + norm.u[1] = HC_WIDTH; + self.dispatch( + c"kernel_qwen_zero_rms", + &self.scratch.hc_norm, + Some(&self.scratch.hc), + None, + None, + &[self.view(self.weight("mtp.pre_fc_norm_hidden.weight")?)], + &norm, + 1, + 1, + )?; + let hidden_weight = self.weight("mtp.fc_hidden.weight")?; + for stream in 0..HC { + let offset = u64::from(stream * HIDDEN) * 4; + let input = self.scratch.hc_norm.view(offset, u64::from(HIDDEN) * 4)?; + let output = self.scratch.hc_mix.view(offset, u64::from(HIDDEN) * 4)?; + self.bf16_mv(hidden_weight, &input, &output, HIDDEN, HIDDEN)?; + } + let mut repeat = args(); + repeat.u[0] = HIDDEN; + self.dispatch( + c"kernel_qwen_repeat4", + &self.scratch.hc_norm, + Some(&self.scratch.hidden), + None, + None, + &[], + &repeat, + HC_WIDTH, + 1, + )?; + let mut add = args(); + add.u[0] = HC_WIDTH; + self.dispatch( + c"kernel_qwen_add", + &self.scratch.hc, + Some(&self.scratch.hc_mix), + Some(&self.scratch.hc_norm), + None, + &[], + &add, + HC_WIDTH, + 1, + )?; + commands.finish()?; + + let prefix = "mtp.layers.0"; + let commands = Commands::begin()?; + self.hyper_read(&format!("{prefix}.attn_hyper_connection"))?; + self.attention(prefix, &mtp.attention, position)?; + self.hyper_write()?; + commands.finish()?; + self.encode_moe(prefix, "MTP")?; + + if !draft_requested { + return Ok(None); + } + let commands = Commands::begin()?; + self.mix_output("mtp.hyper_connection_mixer")?; + self.affine_mv_into( + &self.affine("language_model.lm_head", HIDDEN, VOCAB, None)?, + &self.scratch.block, + &self.scratch.logits, + HIDDEN, + VOCAB, + )?; + commands.finish()?; + self.scratch.logits.read_f32(&mut self.logits)?; + mtp.drafted += 1; + if mtp.timing { + eprintln!( + "ds4: Qwen MTP proposal at position {position} in {:.1} ms", + started.elapsed().as_secs_f64() * 1000.0 + ); + } + Ok(Some(self.logits.clone())) + })(); + self.mtp = Some(mtp); + result + } + + fn capture_recurrent(&mut self, slot: usize) -> Result<(), String> { + let commands = Commands::begin()?; + let mtp = self.mtp.as_mut().ok_or("Qwen MTP is not configured")?; + let snapshot = mtp + .captures + .get_mut(slot) + .ok_or("Qwen MTP capture slot is out of range")?; + for (state, saved) in self.states.iter().zip(&snapshot.gdn) { + match (state, saved) { + (LayerState::Gdn { conv, recurrent }, Some((saved_conv, saved_recurrent))) => { + saved_conv.copy_from( + 0, + conv, + 0, + u64::from(GDN_QKV) * 3 * 2, + "capturing Qwen MTP GDN convolution state", + )?; + saved_recurrent.copy_from( + 0, + recurrent, + 0, + u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4, + "capturing Qwen MTP GDN recurrent state", + )?; + } + (LayerState::Attention { .. }, None) => {} + _ => return Err("Qwen MTP capture layout is invalid".into()), + } + } + snapshot.ple_conv.copy_from( + 0, + &self.ple_state.conv, + 0, + u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2, + "capturing Qwen MTP PLE state", + )?; + snapshot.hidden.copy_from( + 0, + &self.scratch.hc, + 0, + u64::from(HC_WIDTH) * 4, + "capturing Qwen MTP target hidden state", + )?; + snapshot.ple_history = self.ple_state.history; + snapshot.logits.copy_from_slice(&self.logits); + snapshot.position = self.position; + snapshot.tokens = self.tokens.len(); + commands.finish() + } + + fn restore_recurrent(&mut self, slot: usize) -> Result<(), String> { + let commands = Commands::begin()?; + let mtp = self.mtp.as_ref().ok_or("Qwen MTP is not configured")?; + let snapshot = mtp + .captures + .get(slot) + .ok_or("Qwen MTP restore slot is out of range")?; + for (state, saved) in self.states.iter().zip(&snapshot.gdn) { + match (state, saved) { + (LayerState::Gdn { conv, recurrent }, Some((saved_conv, saved_recurrent))) => { + conv.copy_from( + 0, + saved_conv, + 0, + u64::from(GDN_QKV) * 3 * 2, + "restoring Qwen MTP GDN convolution state", + )?; + recurrent.copy_from( + 0, + saved_recurrent, + 0, + u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4, + "restoring Qwen MTP GDN recurrent state", + )?; + } + (LayerState::Attention { .. }, None) => {} + _ => return Err("Qwen MTP restore layout is invalid".into()), + } + } + self.ple_state.conv.copy_from( + 0, + &snapshot.ple_conv, + 0, + u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2, + "restoring Qwen MTP PLE state", + )?; + self.scratch.hc.copy_from( + 0, + &snapshot.hidden, + 0, + u64::from(HC_WIDTH) * 4, + "restoring Qwen MTP target hidden state", + )?; + self.ple_state.history = snapshot.ple_history; + self.logits.copy_from_slice(&snapshot.logits); + self.position = snapshot.position; + self.tokens.truncate(snapshot.tokens); + commands.finish() + } + + fn draft_limit(&self, max_tokens: u32, reasoning: ReasoningMode, token: i32) -> u32 { + if self.mtp.is_none() + || self.position == 0 + || max_tokens <= 1 + || self.model.is_stop_token_for_reasoning(token, reasoning) + { + return 0; + } + max_tokens + .saturating_sub(1) + .min(MTP_TOKEN_RESERVE) + .min(self.context.saturating_sub(self.position + 1)) + } + + fn draft_greedy( + &mut self, + token: i32, + count: u32, + reasoning: ReasoningMode, + ) -> Result, String> { + let mut drafts = Vec::with_capacity(count as usize); + let mut input = token; + let start = self.position - 1; + for step in 0..count { + let logits = self + .mtp_step(input, start + step, true)? + .ok_or("Qwen MTP did not return draft logits")?; + let token = argmax(&logits); + drafts.push(QwenDraft { token, logits }); + input = token; + if self.model.is_stop_token_for_reasoning(token, reasoning) { + break; + } + } + Ok(drafts) + } + + #[allow(clippy::too_many_arguments)] + fn draft_sampled( + &mut self, + token: i32, + count: u32, + reasoning: ReasoningMode, + temperature: f32, + top_p: f32, + min_p: f32, + top_k: i32, + rng: &mut Rng, + ) -> Result, String> { + let mut drafts = Vec::with_capacity(count as usize); + let mut input = token; + let start = self.position - 1; + for step in 0..count { + let logits = self + .mtp_step(input, start + step, true)? + .ok_or("Qwen MTP did not return draft logits")?; + let token = sample_from_logits(&logits, temperature, top_p, min_p, top_k, rng); + drafts.push(QwenDraft { token, logits }); + input = token; + if self.model.is_stop_token_for_reasoning(token, reasoning) { + break; + } + } + Ok(drafts) + } + + fn evaluate_drafts(&mut self, drafts: &[QwenDraft]) -> Result<(Vec>, bool), String> { + let mut captures_complete = self.capture_recurrent(0).is_ok(); + let mut target_logits = Vec::with_capacity(drafts.len()); + for (index, draft) in drafts.iter().enumerate() { + target_logits.push(self.logits.clone()); + self.eval_target(draft.token)?; + if captures_complete { + captures_complete = self.capture_recurrent(index + 1).is_ok(); + } + } + Ok((target_logits, captures_complete)) + } + + fn restore_or_replay( + &mut self, + slot: usize, + tokens: &[i32], + verified_end: u32, + mtp_end: u32, + captures_complete: bool, + ) -> Result<(), String> { + if captures_complete + && self.restore_recurrent(slot).is_ok() + && self.trim_rejected_attention(verified_end, mtp_end).is_ok() + { + return Ok(()); + } + self.replay_tokens(tokens) + } + + fn trim_rejected_attention(&self, verified_end: u32, mtp_end: u32) -> Result<(), String> { + for state in &self.states { + if matches!(state, LayerState::Attention { .. }) { + clear_attention_rows(state, self.position, verified_end)?; + } + } + if let Some(mtp) = &self.mtp { + clear_attention_rows(&mtp.attention, self.position.saturating_sub(1), mtp_end)?; + } + Ok(()) + } + + fn replay_tokens(&mut self, tokens: &[i32]) -> Result<(), String> { + let tag = self.checkpoint_tag; + let counters = self.mtp.as_ref().map(|mtp| { + ( + mtp.cycles, + mtp.drafted, + mtp.accepted, + mtp.accepted_depth_total, + mtp.accepted_depth_max, + mtp.verifier_passes, + mtp.verifier_ns, + mtp.repair_ns, + ) + }); + self.reset()?; + for &token in tokens { + self.eval(token)?; + } + self.checkpoint_tag = tag; + if let (Some(mtp), Some(counters)) = (&mut self.mtp, counters) { + ( + mtp.cycles, + mtp.drafted, + mtp.accepted, + mtp.accepted_depth_total, + mtp.accepted_depth_max, + mtp.verifier_passes, + mtp.verifier_ns, + mtp.repair_ns, + ) = counters; + } + Ok(()) + } + + fn commit_final_mtp_row( + &mut self, + draft: &QwenDraft, + final_slot: usize, + captures_complete: bool, + ) -> Result<(), String> { + let tokens = self.tokens.clone(); + if captures_complete + && self.restore_recurrent(final_slot - 1).is_ok() + && self + .mtp_step(draft.token, self.position.saturating_sub(1), false) + .is_ok() + && self.restore_recurrent(final_slot).is_ok() + { + return Ok(()); + } + self.replay_tokens(&tokens) + } + + fn note_mtp_verification(&mut self, accepted: usize, started: Instant) { + if let Some(mtp) = &mut self.mtp { + mtp.accepted = mtp.accepted.saturating_add(accepted as u64); + mtp.accepted_depth_total = mtp.accepted_depth_total.saturating_add(accepted as u64); + mtp.accepted_depth_max = mtp.accepted_depth_max.max(accepted as u64); + mtp.verifier_passes = mtp.verifier_passes.saturating_add(1); + mtp.verifier_ns = mtp + .verifier_ns + .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); + } + } + + pub(super) fn eval_speculative_greedy( + &mut self, + token: i32, + max_tokens: u32, + reasoning: ReasoningMode, + cancelled: &std::sync::atomic::AtomicBool, + ) -> Result, String> { + if self.mtp.is_none() { + self.eval_target(token)?; + return Ok(vec![token]); + } + if let Some(mtp) = &mut self.mtp { + mtp.cycles = mtp.cycles.saturating_add(1); + } + let count = self.draft_limit(max_tokens, reasoning, token); + if count == 0 || cancelled.load(std::sync::atomic::Ordering::Relaxed) { + if self.position > 0 && !self.model.is_stop_token_for_reasoning(token, reasoning) { + let _ = self.mtp_step(token, self.position - 1, false)?; + } + self.eval_target(token)?; + return Ok(vec![token]); + } + let drafts = self.draft_greedy(token, count, reasoning)?; + self.eval_target(token)?; + if drafts.is_empty() || cancelled.load(std::sync::atomic::Ordering::Relaxed) { + return Ok(vec![token]); + } + let started = Instant::now(); + let baseline_tokens = self.tokens.clone(); + let mtp_end = self.position.saturating_sub(2) + drafts.len() as u32; + let (target_logits, captures_complete) = self.evaluate_drafts(&drafts)?; + let verified_end = self.position; + let accepted = drafts + .iter() + .zip(&target_logits) + .take_while(|(draft, target)| argmax(target) == draft.token) + .count(); + if accepted != drafts.len() { + let repair = Instant::now(); + let mut tokens = baseline_tokens; + tokens.extend(drafts[..accepted].iter().map(|draft| draft.token)); + self.restore_or_replay(accepted, &tokens, verified_end, mtp_end, captures_complete)?; + if let Some(mtp) = &mut self.mtp { + mtp.repair_ns = mtp + .repair_ns + .saturating_add(u64::try_from(repair.elapsed().as_nanos()).unwrap_or(u64::MAX)); + } + } else if let Some(last) = drafts.last() + && !self + .model + .is_stop_token_for_reasoning(last.token, reasoning) + { + self.commit_final_mtp_row(last, drafts.len(), captures_complete)?; + } + self.note_mtp_verification(accepted, started); + let mut emitted = Vec::with_capacity(accepted + 1); + emitted.push(token); + emitted.extend(drafts[..accepted].iter().map(|draft| draft.token)); + Ok(emitted) + } + + #[allow(clippy::too_many_arguments)] + pub(super) fn eval_speculative_sampled( + &mut self, + token: i32, + max_tokens: u32, + reasoning: ReasoningMode, + temperature: f32, + top_p: f32, + min_p: f32, + top_k: i32, + rng: &mut Rng, + cancelled: &std::sync::atomic::AtomicBool, + ) -> Result, String> { + if self.mtp.is_none() { + self.eval_target(token)?; + return Ok(vec![token]); + } + if let Some(mtp) = &mut self.mtp { + mtp.cycles = mtp.cycles.saturating_add(1); + } + let count = self.draft_limit(max_tokens, reasoning, token); + if count == 0 || cancelled.load(std::sync::atomic::Ordering::Relaxed) { + if self.position > 0 && !self.model.is_stop_token_for_reasoning(token, reasoning) { + let _ = self.mtp_step(token, self.position - 1, false)?; + } + self.eval_target(token)?; + return Ok(vec![token]); + } + let drafts = self.draft_sampled( + token, + count, + reasoning, + temperature, + top_p, + min_p, + top_k, + rng, + )?; + self.eval_target(token)?; + if drafts.is_empty() || cancelled.load(std::sync::atomic::Ordering::Relaxed) { + return Ok(vec![token]); + } + let started = Instant::now(); + let baseline_tokens = self.tokens.clone(); + let mtp_end = self.position.saturating_sub(2) + drafts.len() as u32; + let (target_logits, captures_complete) = self.evaluate_drafts(&drafts)?; + let verified_end = self.position; + let mut accepted = 0; + let mut replacement = None; + for ((draft, target), index) in drafts.iter().zip(&target_logits).zip(0..) { + let (sampled, was_draft) = exact_speculative_sample( + target, + &draft.logits, + draft.token, + temperature, + top_p, + min_p, + top_k, + rng, + ); + if !was_draft { + replacement = Some((index, sampled)); + break; + } + accepted += 1; + } + let mut emitted = Vec::with_capacity(accepted + 2); + emitted.push(token); + emitted.extend(drafts[..accepted].iter().map(|draft| draft.token)); + if let Some((slot, replacement)) = replacement { + let repair = Instant::now(); + let mut tokens = baseline_tokens; + tokens.extend(drafts[..accepted].iter().map(|draft| draft.token)); + self.restore_or_replay(slot, &tokens, verified_end, mtp_end, captures_complete)?; + if !self + .model + .is_stop_token_for_reasoning(replacement, reasoning) + { + let _ = self.mtp_step(replacement, self.position - 1, false)?; + } + self.eval_target(replacement)?; + emitted.push(replacement); + if let Some(mtp) = &mut self.mtp { + mtp.repair_ns = mtp + .repair_ns + .saturating_add(u64::try_from(repair.elapsed().as_nanos()).unwrap_or(u64::MAX)); + } + } else if let Some(last) = drafts.last() + && !self + .model + .is_stop_token_for_reasoning(last.token, reasoning) + { + self.commit_final_mtp_row(last, drafts.len(), captures_complete)?; + } + self.note_mtp_verification(accepted, started); + Ok(emitted) + } + pub(super) fn prefill( &mut self, tokens: &[i32], @@ -1400,7 +2039,20 @@ impl QwenExecutor { &self.logits } pub(super) fn execution_stats(&self) -> ExecutionStats { - ExecutionStats::default() + self.mtp + .as_ref() + .map_or_else(ExecutionStats::default, |mtp| ExecutionStats { + speculative_mode: 4, + speculative_cycles: mtp.cycles, + drafted_tokens: mtp.drafted, + accepted_draft_tokens: mtp.accepted, + accepted_depth_total: mtp.accepted_depth_total, + accepted_depth_max: mtp.accepted_depth_max, + verifier_passes: mtp.verifier_passes, + verifier_ms: mtp.verifier_ns / 1_000_000, + repair_ms: mtp.repair_ns / 1_000_000, + ..ExecutionStats::default() + }) } pub(super) fn model(&self) -> &QwenModel { &self.model @@ -1424,6 +2076,9 @@ impl QwenExecutor { pub(super) fn reset(&mut self) -> Result<(), String> { self.states = allocate_states(self.context)?; self.ple_state = allocate_ple_state()?; + if let Some(mtp) = &self.mtp { + self.mtp = Some(allocate_mtp(self.context, mtp.timing)?); + } self.logits.fill(0.0); self.tokens.clear(); self.position = 0; @@ -1446,6 +2101,11 @@ impl QwenExecutor { tokens: Vec::new(), position: 0, checkpoint_tag: [0; 32], + mtp: self + .mtp + .as_ref() + .map(|mtp| allocate_mtp(self.context, mtp.timing)) + .transpose()?, }) } @@ -1460,6 +2120,7 @@ impl QwenExecutor { std::mem::swap(&mut self.tokens, &mut incoming.tokens); std::mem::swap(&mut self.position, &mut incoming.position); std::mem::swap(&mut self.checkpoint_tag, &mut incoming.checkpoint_tag); + std::mem::swap(&mut self.mtp, &mut incoming.mtp); *state = Some(incoming); Ok(()) } @@ -1483,6 +2144,7 @@ impl QwenExecutor { self.position, VOCAB, LAYERS as u32, + u32::from(self.mtp.is_some()), ] { write_u32(&mut file, value)?; } @@ -1540,6 +2202,20 @@ impl QwenExecutor { } } } + if let Some(mtp) = &self.mtp { + let LayerState::Attention { + kv, + qsa_raw, + qsa_pooled, + } = &mtp.attention + else { + return Err("Qwen MTP checkpoint state is invalid".into()); + }; + let [kv_bytes, raw_bytes, pooled_bytes] = attention_checkpoint_bytes(self.position); + write_buffer(&mut file, kv, 0, kv_bytes, &mut chunk, progress)?; + write_buffer(&mut file, qsa_raw, 0, raw_bytes, &mut chunk, progress)?; + write_buffer(&mut file, qsa_pooled, 0, pooled_bytes, &mut chunk, progress)?; + } file.sync_all().map_err(|error| error.to_string())?; fs::rename(temporary, path).map_err(|error| error.to_string())?; self.checkpoint_tag = tag; @@ -1574,6 +2250,10 @@ impl QwenExecutor { { return Err("Qwen checkpoint shape is invalid".into()); } + let checkpoint_mtp = read_u32(&mut file)?; + if checkpoint_mtp > 1 || (checkpoint_mtp != 0) != self.mtp.is_some() { + return Err("Qwen checkpoint MTP configuration does not match the executor".into()); + } let mut identity = [0; 32]; file.read_exact(&mut identity) .map_err(|error| error.to_string())?; @@ -1645,6 +2325,20 @@ impl QwenExecutor { } } } + if let Some(mtp) = &self.mtp { + let LayerState::Attention { + kv, + qsa_raw, + qsa_pooled, + } = &mtp.attention + else { + return Err("Qwen MTP checkpoint state is invalid".into()); + }; + let [kv_bytes, raw_bytes, pooled_bytes] = attention_checkpoint_bytes(position); + read_buffer(&mut file, kv, 0, kv_bytes, &mut chunk, progress)?; + read_buffer(&mut file, qsa_raw, 0, raw_bytes, &mut chunk, progress)?; + read_buffer(&mut file, qsa_pooled, 0, pooled_bytes, &mut chunk, progress)?; + } let mut trailing = [0]; if file .read(&mut trailing) @@ -1691,6 +2385,92 @@ fn allocate_states(context: u32) -> Result, String> { .collect() } +fn allocate_attention_state(context: u32) -> Result { + let token_capacity = context + .checked_add(MTP_TOKEN_RESERVE) + .ok_or_else(|| "Qwen MTP attention capacity overflows".to_owned())?; + Ok(LayerState::Attention { + kv: Buffer::bytes(u64::from(token_capacity) * ATTN_KV_WIDTH as u64 * 4)?, + qsa_raw: Buffer::bytes(u64::from(token_capacity) * QSA_DIM as u64 * 2)?, + qsa_pooled: Buffer::bytes(u64::from(qsa_block_capacity(context)?) * QSA_DIM as u64 * 2)?, + }) +} + +fn allocate_recurrent_snapshot() -> Result { + let gdn = (0..LAYERS) + .map(|layer| { + if layer % 4 == 3 { + Ok(None) + } else { + Ok(Some(( + Buffer::bytes(u64::from(GDN_QKV) * 3 * 2)?, + Buffer::bytes(u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4)?, + ))) + } + }) + .collect::>()?; + Ok(RecurrentSnapshot { + gdn, + ple_conv: Buffer::bytes(u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2)?, + ple_history: [0; PLE_HISTORY], + hidden: Buffer::floats(HC_WIDTH.into())?, + logits: vec![0.0; VOCAB as usize], + position: 0, + tokens: 0, + }) +} + +fn allocate_mtp(context: u32, timing: bool) -> Result { + Ok(QwenMtp { + attention: allocate_attention_state(context)?, + captures: (0..MTP_CAPTURE_ROWS) + .map(|_| allocate_recurrent_snapshot()) + .collect::>()?, + timing, + cycles: 0, + drafted: 0, + accepted: 0, + accepted_depth_total: 0, + accepted_depth_max: 0, + verifier_passes: 0, + verifier_ns: 0, + repair_ns: 0, + }) +} + +fn clear_attention_rows(state: &LayerState, start: u32, end: u32) -> Result<(), String> { + if end <= start { + return Ok(()); + } + let LayerState::Attention { + kv, + qsa_raw, + qsa_pooled, + } = state + else { + return Err("Qwen attention trim received GDN state".into()); + }; + fn clear(buffer: &Buffer, row_bytes: u64, start: u32, end: u32) -> Result<(), String> { + let bytes = u64::from(end - start) + .checked_mul(row_bytes) + .ok_or_else(|| "Qwen attention trim size overflows".to_owned())?; + let zeros = vec![ + 0; + usize::try_from(bytes) + .map_err(|_| "Qwen attention trim exceeds addressable memory")? + ]; + buffer.write(u64::from(start) * row_bytes, &zeros) + } + clear(kv, u64::from(ATTN_KV_WIDTH) * 4, start, end)?; + clear(qsa_raw, u64::from(QSA_DIM) * 2, start, end)?; + clear( + qsa_pooled, + u64::from(QSA_DIM) * 2, + start / QSA_RATIO, + end / QSA_RATIO, + ) +} + fn qsa_block_capacity(context: u32) -> Result { context .checked_add(MTP_TOKEN_RESERVE) @@ -2949,7 +3729,7 @@ mod tests { .write(0, &vec![0; (QSA_TOP_K * QSA_DIM * 2) as usize]) .unwrap(); executor - .attention("language_model.model.layers.3", 3) + .attention("language_model.model.layers.3", &executor.states[3], 3) .unwrap(); let mut hidden = vec![0.0; HIDDEN as usize]; executor.scratch.hidden.read_f32(&mut hidden).unwrap(); @@ -3090,7 +3870,7 @@ mod tests { .unwrap(); assert_eq!( fs::metadata(&checkpoint).unwrap().len(), - 116_635_748 + u64::from(depth) * 4 + live_attention_bytes + 116_635_752 + u64::from(depth) * 4 + live_attention_bytes ); assert!(executor.load_checkpoint(&checkpoint, &mut |_| {}).unwrap()); assert_eq!(executor.position, depth); @@ -3311,4 +4091,249 @@ mod tests { assert_eq!(executor.checkpoint_tag, [8; 32]); fs::remove_file(checkpoint).unwrap(); } + + #[test] + #[ignore = "requires the pinned 105 GB Qwen artifact set and Apple Metal"] + fn qwen_mtp_matches_target_and_restores_its_cache() { + fn digest(logits: &[f32]) -> [u8; 32] { + let mut hash = Sha256::new(); + for value in logits { + hash.update(value.to_bits().to_le_bytes()); + } + hash.finalize().into() + } + + fn target_state_digest(executor: &QwenExecutor) -> [u8; 32] { + fn buffer(hash: &mut Sha256, buffer: &Buffer, bytes: u64) { + let mut values = vec![0; usize::try_from(bytes).unwrap()]; + buffer.read(0, &mut values).unwrap(); + hash.update(values); + } + let mut hash = Sha256::new(); + hash.update(executor.position.to_le_bytes()); + for token in &executor.tokens { + hash.update(token.to_le_bytes()); + } + for value in &executor.logits { + hash.update(value.to_bits().to_le_bytes()); + } + for value in executor.ple_state.history { + hash.update(value.to_le_bytes()); + } + buffer( + &mut hash, + &executor.ple_state.conv, + u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2, + ); + buffer(&mut hash, &executor.scratch.hc, u64::from(HC_WIDTH) * 4); + for state in &executor.states { + match state { + LayerState::Gdn { conv, recurrent } => { + buffer(&mut hash, conv, u64::from(GDN_QKV) * 3 * 2); + buffer( + &mut hash, + recurrent, + u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4, + ); + } + LayerState::Attention { + kv, + qsa_raw, + qsa_pooled, + } => { + let [kv_bytes, raw_bytes, pooled_bytes] = + attention_checkpoint_bytes(executor.position); + buffer(&mut hash, kv, kv_bytes); + buffer(&mut hash, qsa_raw, raw_bytes); + buffer(&mut hash, qsa_pooled, pooled_bytes); + } + } + } + hash.finalize().into() + } + + fn mtp_attention_digest(executor: &QwenExecutor) -> [u8; 32] { + let LayerState::Attention { + kv, + qsa_raw, + qsa_pooled, + } = &executor.mtp.as_ref().unwrap().attention + else { + unreachable!() + }; + let capacity = executor.context + MTP_TOKEN_RESERVE; + let mut hash = Sha256::new(); + for (buffer, bytes) in [ + (kv, u64::from(capacity) * ATTN_KV_WIDTH as u64 * 4), + (qsa_raw, u64::from(capacity) * QSA_DIM as u64 * 2), + ( + qsa_pooled, + u64::from(qsa_block_capacity(executor.context).unwrap()) * QSA_DIM as u64 * 2, + ), + ] { + let mut values = vec![0; usize::try_from(bytes).unwrap()]; + buffer.read(0, &mut values).unwrap(); + hash.update(values); + } + hash.finalize().into() + } + + configure_sources().unwrap(); + let root = std::env::var_os("DS4SERVER_QWEN38_SOURCE") + .map(PathBuf::from) + .expect("set DS4SERVER_QWEN38_SOURCE to the pinned artifact directory"); + let mut baseline = QwenExecutor::open(QwenModel::open(&root, 8).unwrap(), 8).unwrap(); + baseline.eval(1).unwrap(); + baseline.eval(2).unwrap(); + let seeded = digest(baseline.logits()); + let mut token = argmax(baseline.logits()); + let mut target_tokens = Vec::new(); + let mut target_digests = Vec::new(); + let mut target_state_digests = Vec::new(); + for _ in 0..4 { + target_tokens.push(token); + baseline.eval(token).unwrap(); + target_digests.push(digest(baseline.logits())); + target_state_digests.push(target_state_digest(&baseline)); + token = argmax(baseline.logits()); + } + drop(baseline); + + let settings = EngineSpeculativeSettings { + glm_mtp: true, + ..EngineSpeculativeSettings::default() + }; + let model = QwenModel::open_configured(&root, 8, true).unwrap(); + let mut executor = QwenExecutor::open_configured(model, 8, settings).unwrap(); + executor.eval(1).unwrap(); + executor.eval(2).unwrap(); + assert_eq!(digest(executor.logits()), seeded); + let LayerState::Attention { kv, .. } = &executor.mtp.as_ref().unwrap().attention else { + unreachable!() + }; + let mut mtp_frontier = vec![0; ATTN_KV_WIDTH as usize * 4]; + kv.read(0, &mut mtp_frontier).unwrap(); + assert!(mtp_frontier.iter().any(|byte| *byte != 0)); + + let pinned = executor + .draft_greedy(target_tokens[0], 3, ReasoningMode::Direct) + .unwrap(); + assert_eq!( + pinned + .iter() + .map(|draft| (draft.token, digest(&draft.logits))) + .collect::>(), + vec![ + ( + 1788, + [ + 101, 126, 239, 116, 151, 56, 134, 52, 71, 6, 11, 42, 11, 2, 178, 196, 80, + 234, 100, 88, 71, 26, 168, 15, 170, 92, 98, 85, 194, 250, 179, 145, + ], + ), + ( + 1151, + [ + 102, 84, 243, 59, 200, 12, 148, 234, 248, 175, 4, 153, 221, 246, 199, 149, + 56, 191, 105, 233, 72, 18, 86, 132, 62, 215, 116, 22, 151, 136, 56, 63, + ], + ), + ( + 8598, + [ + 102, 99, 10, 59, 130, 251, 42, 204, 133, 69, 46, 151, 25, 67, 85, 230, 132, + 176, 220, 200, 168, 160, 46, 136, 204, 105, 122, 0, 58, 239, 198, 72, + ], + ), + ] + ); + let mut mtp_hidden = vec![0.0; HIDDEN as usize]; + executor.scratch.block.read_f32(&mut mtp_hidden).unwrap(); + assert_eq!( + digest(&mtp_hidden), + [ + 37, 149, 12, 95, 102, 7, 237, 89, 2, 170, 48, 69, 151, 191, 178, 241, 145, 208, + 147, 40, 132, 7, 25, 233, 103, 20, 249, 90, 86, 102, 255, 240, + ] + ); + + executor.eval_target(target_tokens[0]).unwrap(); + let target_drafts = target_tokens[1..] + .iter() + .map(|&token| QwenDraft { + token, + logits: Vec::new(), + }) + .collect::>(); + let (_, complete) = executor.evaluate_drafts(&target_drafts).unwrap(); + assert!(complete); + let verified_end = executor.position; + for slot in (0..=target_drafts.len()).rev() { + executor.restore_recurrent(slot).unwrap(); + executor.trim_rejected_attention(verified_end, 4).unwrap(); + assert_eq!( + target_state_digest(&executor), + target_state_digests[slot], + "target state after accepting {slot} draft tokens", + ); + } + let fallback_tokens = vec![1, 2, target_tokens[0], target_tokens[1]]; + executor + .restore_or_replay(1, &fallback_tokens, verified_end, 4, false) + .unwrap(); + assert_eq!( + target_state_digest(&executor), + target_state_digests[1], + "ordinary-forward fallback state", + ); + executor.reset().unwrap(); + executor.eval(1).unwrap(); + executor.eval(2).unwrap(); + + let checkpoint = + std::env::temp_dir().join(format!("ds4-qwen98-mtp-checkpoint-{}", std::process::id())); + executor + .save_checkpoint(&checkpoint, [98; 32], &mut |_| {}) + .unwrap(); + executor.reset().unwrap(); + assert!(executor.load_checkpoint(&checkpoint, &mut |_| {}).unwrap()); + let LayerState::Attention { kv, .. } = &executor.mtp.as_ref().unwrap().attention else { + unreachable!() + }; + let mut restored = vec![0; mtp_frontier.len()]; + kv.read(0, &mut restored).unwrap(); + assert_eq!(restored, mtp_frontier); + assert_eq!(executor.checkpoint_tag(), [98; 32]); + + let first = target_tokens[0]; + let emitted = executor + .eval_speculative_greedy( + first, + 4, + ReasoningMode::Direct, + &std::sync::atomic::AtomicBool::new(false), + ) + .unwrap(); + assert_eq!(emitted, target_tokens[..emitted.len()]); + assert_eq!(digest(executor.logits()), target_digests[emitted.len() - 1]); + let stats = executor.execution_stats(); + assert_eq!(stats.speculative_cycles, 1); + assert_eq!(stats.drafted_tokens, 3); + assert_eq!(stats.accepted_draft_tokens, (emitted.len() - 1) as u64); + assert_eq!(stats.verifier_passes, 1); + let mtp_before_eos = mtp_attention_digest(&executor); + assert_eq!( + executor + .eval_speculative_greedy( + EOS_TOKEN, + 4, + ReasoningMode::Direct, + &std::sync::atomic::AtomicBool::new(false), + ) + .unwrap(), + [EOS_TOKEN], + ); + assert_eq!(mtp_attention_digest(&executor), mtp_before_eos); + fs::remove_file(checkpoint).unwrap(); + } } diff --git a/src/engine/qwen.rs b/src/engine/qwen.rs index 8c783e5..0a9832e 100644 --- a/src/engine/qwen.rs +++ b/src/engine/qwen.rs @@ -30,6 +30,8 @@ const MTP_TOKEN_RESERVE: u64 = 3; const GDN_STATE_BYTES: u64 = 113_246_208; const GDN_CONV_BYTES: u64 = 2_211_840; const PLE_CONV_BYTES: u64 = 184_320; +const MTP_CAPTURE_HIDDEN_BYTES: u64 = 10_240 * 4; +const MTP_CAPTURE_LOGITS_BYTES: u64 = 248_320 * 4; #[derive(Deserialize)] struct Manifest { @@ -125,10 +127,22 @@ pub(super) struct QwenModel { } impl QwenModel { + #[cfg(test)] pub(super) fn open(root: &Path, context: u32) -> Result { - let loaded = load(root, context, false)?; + Self::open_configured(root, context, false) + } + + pub(super) fn open_configured( + root: &Path, + context: u32, + enable_mtp: bool, + ) -> Result { + let loaded = load(root, context, enable_mtp)?; let mut bindings = loaded.bindings.core; bindings.extend(loaded.bindings.ple); + if enable_mtp { + bindings.extend(loaded.bindings.mtp); + } let mut paths = bindings .iter() .map(|binding| binding.file.clone()) @@ -176,7 +190,7 @@ impl QwenModel { pub(super) fn tensor(&self, name: &str) -> Result<&QwenTensor, String> { self.tensors .get(name) - .ok_or_else(|| format!("Qwen core tensor is missing: {name}")) + .ok_or_else(|| format!("Qwen tensor is missing: {name}")) } pub(super) fn map(&self, index: usize) -> (&[u8], &Path) { @@ -424,10 +438,13 @@ fn validate_precision(name: &str, tensor: &ExpectedTensor) -> Result<(), String> (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.") { + if group == 32 + && !name.starts_with("ngram.") + && !name.starts_with("mtp.layers.0.mlp.switch_mlp.") + { return Err(format!("{name} unexpectedly uses 32-value groups")); } - if bits == 2 && !name.starts_with("mtp.") { + if bits == 2 { return Err(format!("{name} unexpectedly uses 2-bit weights")); } Ok(()) @@ -666,7 +683,30 @@ pub(super) fn memory_plan( .and_then(|bytes| bytes.checked_add(QSA_FIXED_SCRATCH_BYTES)) .and_then(|bytes| bytes.checked_add(topk_scratch)) .ok_or_else(|| "Qwen prefill transient size overflows".to_owned())?; - let admitted_mtp = if enable_mtp { MTP_BYTES } else { 0 }; + let admitted_mtp = if enable_mtp { + let capture_rows = MTP_TOKEN_RESERVE + 1; + let recurrent_capture = (GDN_STATE_BYTES + + GDN_CONV_BYTES + + PLE_CONV_BYTES + + MTP_CAPTURE_HIDDEN_BYTES + + MTP_CAPTURE_LOGITS_BYTES) + .checked_mul(capture_rows) + .ok_or_else(|| "Qwen MTP verifier capture size overflows".to_owned())?; + let mtp_attention = (KV_BYTES_PER_TOKEN + QSA_RAW_BYTES_PER_TOKEN) + .checked_mul(token_capacity) + .and_then(|bytes| { + QSA_POOLED_BYTES_PER_BLOCK + .checked_mul(block_capacity) + .and_then(|pooled| bytes.checked_add(pooled)) + }) + .ok_or_else(|| "Qwen MTP attention size overflows".to_owned())?; + MTP_BYTES + .checked_add(recurrent_capture) + .and_then(|bytes| bytes.checked_add(mtp_attention)) + .ok_or_else(|| "Qwen MTP admission size overflows".to_owned())? + } else { + 0 + }; let admission = CORE_BYTES .checked_add(admitted_mtp) .and_then(|bytes| bytes.checked_add(kv_and_recurrent)) @@ -706,7 +746,7 @@ mod tests { assert_eq!(plan.optional_mtp, MTP_BYTES); assert_eq!(plan.kv_and_recurrent, 7_564_812_288); assert_eq!(plan.prefill_transient, 28_056_068); - assert_eq!(plan.admission, 81_008_126_487); + assert_eq!(plan.admission, 88_924_002_839); let without_mtp = memory_plan(262_144, false, 512).unwrap(); assert_eq!(without_mtp.optional_mtp, MTP_BYTES); assert_eq!(without_mtp.admission, 79_335_550_955); diff --git a/src/metrics.rs b/src/metrics.rs index cf8e2cb..b944d14 100644 --- a/src/metrics.rs +++ b/src/metrics.rs @@ -113,8 +113,11 @@ pub(crate) struct MetricsSnapshot { pub(crate) speculative_cycles: u64, pub(crate) drafted_tokens: u64, pub(crate) accepted_draft_tokens: u64, + pub(crate) accepted_depth_total: u64, + pub(crate) accepted_depth_max: u64, pub(crate) verifier_passes: u64, pub(crate) verifier_ms: u64, + pub(crate) repair_ms: u64, pub(crate) ssd_enabled: bool, pub(crate) ssd_resident_bytes: u64, pub(crate) ssd_cache_bytes: u64, @@ -206,8 +209,11 @@ pub(crate) struct Metrics { speculative_cycles: AtomicU64, drafted_tokens: AtomicU64, accepted_draft_tokens: AtomicU64, + accepted_depth_total: AtomicU64, + accepted_depth_max: AtomicU64, verifier_passes: AtomicU64, verifier_ms: AtomicU64, + repair_ms: AtomicU64, ssd_enabled: AtomicBool, ssd_resident_bytes: AtomicU64, ssd_cache_bytes: AtomicU64, @@ -304,8 +310,11 @@ impl Metrics { speculative_cycles: AtomicU64::new(0), drafted_tokens: AtomicU64::new(0), accepted_draft_tokens: AtomicU64::new(0), + accepted_depth_total: AtomicU64::new(0), + accepted_depth_max: AtomicU64::new(0), verifier_passes: AtomicU64::new(0), verifier_ms: AtomicU64::new(0), + repair_ms: AtomicU64::new(0), ssd_enabled: AtomicBool::new(false), ssd_resident_bytes: AtomicU64::new(0), ssd_cache_bytes: AtomicU64::new(0), @@ -472,23 +481,32 @@ impl Metrics { self.source.store(WorkSource::None as u8, Ordering::Relaxed); } + #[allow(clippy::too_many_arguments)] pub(crate) fn speculative_stats( &self, mode: u8, cycles: u64, drafted: u64, accepted: u64, + accepted_depth_total: u64, + accepted_depth_max: u64, verifier_passes: u64, verifier_ms: u64, + repair_ms: u64, ) { self.speculative_mode.store(mode, Ordering::Relaxed); self.speculative_cycles.store(cycles, Ordering::Relaxed); self.drafted_tokens.store(drafted, Ordering::Relaxed); self.accepted_draft_tokens .store(accepted, Ordering::Relaxed); + self.accepted_depth_total + .store(accepted_depth_total, Ordering::Relaxed); + self.accepted_depth_max + .store(accepted_depth_max, Ordering::Relaxed); self.verifier_passes .store(verifier_passes, Ordering::Relaxed); self.verifier_ms.store(verifier_ms, Ordering::Relaxed); + self.repair_ms.store(repair_ms, Ordering::Relaxed); } pub(crate) fn ssd_stats(&self, stats: SsdStats) { @@ -540,7 +558,7 @@ impl Metrics { self.decode_tps.store(0, Ordering::Relaxed); self.prefill_tps.store(0, Ordering::Relaxed); self.prefill_sample.store(0, Ordering::Relaxed); - self.speculative_stats(0, 0, 0, 0, 0, 0); + self.speculative_stats(0, 0, 0, 0, 0, 0, 0, 0, 0); self.ssd_stats(SsdStats::default()); self.model_unloads.fetch_add(1, Ordering::Relaxed); } @@ -695,8 +713,11 @@ impl Metrics { speculative_cycles: self.speculative_cycles.load(Ordering::Relaxed), drafted_tokens: self.drafted_tokens.load(Ordering::Relaxed), accepted_draft_tokens: self.accepted_draft_tokens.load(Ordering::Relaxed), + accepted_depth_total: self.accepted_depth_total.load(Ordering::Relaxed), + accepted_depth_max: self.accepted_depth_max.load(Ordering::Relaxed), verifier_passes: self.verifier_passes.load(Ordering::Relaxed), verifier_ms: self.verifier_ms.load(Ordering::Relaxed), + repair_ms: self.repair_ms.load(Ordering::Relaxed), ssd_enabled: self.ssd_enabled.load(Ordering::Relaxed), ssd_resident_bytes: self.ssd_resident_bytes.load(Ordering::Relaxed), ssd_cache_bytes: self.ssd_cache_bytes.load(Ordering::Relaxed), diff --git a/src/model.rs b/src/model.rs index d766f81..4470bab 100644 --- a/src/model.rs +++ b/src/model.rs @@ -128,8 +128,8 @@ impl ModelChoice { self == Self::Qwen38FlashNext } - pub(crate) fn supports_glm_mtp(self) -> bool { - self.is_glm() + pub(crate) fn supports_integrated_mtp(self) -> bool { + self.is_glm() || self.is_qwen38() } pub(crate) fn main_artifact_size(self) -> u64 { diff --git a/src/model/qwen.rs b/src/model/qwen.rs index 532e5e1..aec32f2 100644 --- a/src/model/qwen.rs +++ b/src/model/qwen.rs @@ -13,7 +13,7 @@ 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_SHA256: &str = "6f1172de47fa30b9602e13fc7ad14e578a813bb5b1fab8f0ef320041aac6c19a"; const MANIFEST_BYTES: &[u8] = include_bytes!("../../assets/models/qwen38-flash-next-bare-speed.json"); diff --git a/src/settings.rs b/src/settings.rs index 4c842fb..0ec6049 100644 --- a/src/settings.rs +++ b/src/settings.rs @@ -45,10 +45,10 @@ pub(crate) struct SpeculativePreferences { impl SpeculativePreferences { pub(crate) fn validate(&self, model: ModelChoice) -> Result<(), String> { if self.glm_mtp_timing && !self.glm_mtp { - return Err("GLM MTP timing requires GLM MTP.".into()); + return Err("MTP timing requires integrated MTP.".into()); } - if !model.supports_glm_mtp() && (self.glm_mtp || self.glm_mtp_timing) { - return Err("GLM MTP is available only for GLM models.".into()); + if !model.supports_integrated_mtp() && (self.glm_mtp || self.glm_mtp_timing) { + return Err("Integrated MTP is unavailable for the selected model.".into()); } if self.keep_vision_loaded && model != ModelChoice::Glm53Flash { return Err("Persistent vision weights are available only for GLM 5.3 Flash.".into()); @@ -83,7 +83,7 @@ impl SpeculativePreferences { } } -#[derive(Clone, Copy, Debug, PartialEq)] +#[derive(Clone, Copy, Debug, Default, PartialEq)] pub(crate) struct EngineSpeculativeSettings { pub(crate) glm_mtp: bool, pub(crate) glm_mtp_timing: bool, diff --git a/tools/qwen38-artifacts.rs b/tools/qwen38-artifacts.rs index 5afca3f..a963483 100644 --- a/tools/qwen38-artifacts.rs +++ b/tools/qwen38-artifacts.rs @@ -657,7 +657,7 @@ fn precision_assignments( reject_unassigned_quantized(&core, &assignments)?; let mtp = tensor_headers(manifest, source, Some("mtp"), complete)?; - add_inferred_sidecar(&mtp, 64, &mut assignments)?; + add_inferred_sidecar(&mtp, &mut assignments)?; let ple = tensor_headers(manifest, source, Some("ple"), complete)?; let ple_group = config .pointer("/mtplx_recipe/ngram/group_size") @@ -711,7 +711,6 @@ fn validate_unquantized(base: &str, tensors: &BTreeMap) -> Resul fn add_inferred_sidecar( tensors: &BTreeMap, - group_size: u64, assignments: &mut BTreeMap, ) -> Result<(), String> { for (name, tensor) in tensors { @@ -732,6 +731,11 @@ fn add_inferred_sidecar( .shape .last() .ok_or_else(|| format!("{base}.scales has no dimensions"))?; + let group_size = if base.contains(".switch_mlp.") { + 32 + } else { + 64 + }; let pack = groups .checked_mul(group_size) .and_then(|columns| columns.checked_div(packed))