use std::collections::VecDeque;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Duration, Instant};
use anyhow::{Context, Result};
use mlx_native::MlxBuffer;
use tokenizers::Tokenizer;
use crate::inference::models::qwen35::kv_cache::{
HybridKvCache, HybridKvCacheSnapshot, HybridKvSlotAnchor,
};
use crate::inference::models::qwen35::model::Qwen35Model;
use crate::inference::spec_decode::cost_controller::SpeculationCostController;
use crate::inference::spec_decode::ngram_proposer::{HistoryLookupConfig, HistoryLookupIndex};
use crate::serve::load_info::{
self, ArchFamily, ChatTemplateSource, LoadInfo, LoadInfoBuilder, MoeShape, TokenizerSource,
};
use crate::core::provenance::{self, Provenance};
use crate::serve::multi_seq_kv::SlotId;
use super::engine::{
effective_repetition_penalty, grammar_runtime_for_request, sample_logits_with_grammar,
supervised_gpu_call, DeepstackData, LoadOptions, SamplingParams, SerialStreamEnd,
SerialStreamResult, SoftTokenData, ToolCallPolicy,
};
use super::engine_supervisor::EngineSupervisor;
const QWEN35_WORKER_TRANSACTION_TIMEOUT: Duration = Duration::from_secs(30);
pub const QWEN35_EOS_STOP_NAMES: &[&str] = &["<|im_end|>", "<|endoftext|>"];
const QWEN35_VISION_MARKERS: &[&str] = &["<|vision_start|>", "<|image_pad|>", "<|vision_end|>"];
fn build_qwen35_serving_tokenizer(gguf: &mlx_native::gguf::GgufFile) -> Result<(Tokenizer, bool)> {
let mut tokenizer =
crate::inference::models::qwen35::tokenizer::build_tokenizer_from_gguf(gguf)
.map_err(|e| anyhow::anyhow!("GGUF-driven tokenizer build failed: {e}"))?;
tokenizer
.with_truncation(None)
.map_err(|e| anyhow::anyhow!("Failed to disable tokenizer truncation: {e}"))?;
let vision_special_tokens_present = QWEN35_VISION_MARKERS
.iter()
.all(|token| tokenizer.token_to_id(token).is_some());
Ok((tokenizer, vision_special_tokens_present))
}
pub struct Qwen35LoadedModel {
pub model: Qwen35Model,
pub tokenizer: Tokenizer,
pub chat_template: String,
pub model_id: String,
pub model_path: PathBuf,
pub eos_token_ids: Vec<u32>,
pub hidden_size: usize,
pub vocab_size: usize,
pub context_length: Option<usize>,
pub quant_type: Option<String>,
pub load_duration: Duration,
pub provenance: Provenance,
pub vision_projector_profile: Option<String>,
pub vision_deepstack_output_count: Option<u32>,
pub vision_special_tokens_present: bool,
pub prompt_cache: HybridPromptCache,
pub lcp_registry: crate::serve::kv_persist::lcp_registry::LcpRegistry<
crate::inference::models::qwen35::kv_cache::HybridKvCacheSnapshot,
>,
pub kv_metrics_sink:
Option<std::sync::Arc<dyn crate::serve::kv_persist::metrics::KvCacheMetricsSink>>,
pub disk_persistor: Option<
std::sync::Arc<
crate::serve::kv_persist::families::qwen35_disk_persistor::Qwen35DiskPersistor,
>,
>,
pub lcp_hydrated_for_cfg: std::collections::HashSet<String>,
pub tq_kv_active: bool,
pub persistent_kv_cache: Option<HybridKvCache>,
pub speculation: super::qwen35_speculation::QwenSpeculationController,
}
impl Qwen35LoadedModel {
pub fn load(opts: &LoadOptions) -> Result<Self> {
let load_start = Instant::now();
let model_path = &opts.model_path;
anyhow::ensure!(
model_path.exists(),
"Model not found: {}",
model_path.display()
);
let gguf = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF open: {e}"))?;
let provenance = provenance::detect(&gguf);
let vision_projector_profile = gguf
.metadata_string("hf2q.vision.projector_profile")
.map(str::to_owned);
let vision_deepstack_output_count = gguf.metadata_u32("hf2q.vision.deepstack_output_count");
let stderr_is_tty = std::io::IsTerminal::is_terminal(&std::io::stderr());
let verbosity = if tracing::enabled!(tracing::Level::INFO) {
1
} else {
0
};
let cfg_preview = Qwen35Model::load_config_only(&gguf)
.context("Qwen35Model::load_config_only (progress sizing)")?;
let mut progress = crate::serve::header::LoadProgress::new(
stderr_is_tty,
verbosity,
cfg_preview.num_hidden_layers as usize,
);
let mut model = Qwen35Model::load_from_gguf(&gguf, &mut progress)
.context("Qwen35Model::load_from_gguf")?;
if let Some(overlay_path) = opts.dwq_overlay_path.as_ref() {
let device = mlx_native::MlxDevice::new()
.map_err(|e| anyhow::anyhow!("Qwen35 DWQ overlay device: {e}"))?;
let stacked = model
.apply_dwq_overlay(&device, overlay_path)
.with_context(|| {
format!("Qwen35 DWQ overlay from {} failed", overlay_path.display())
})?;
tracing::info!(
count = stacked,
path = %overlay_path.display(),
"DWQ overlay applied to Qwen35LoadedModel"
);
}
let mut eos_token_ids: Vec<u32> = Vec::with_capacity(2);
if let Some(id) = gguf.metadata_u32("tokenizer.ggml.eos_token_id") {
eos_token_ids.push(id);
}
if let Some(id) = gguf.metadata_u32("tokenizer.ggml.eot_token_id") {
if !eos_token_ids.contains(&id) {
eos_token_ids.push(id);
}
}
if let Some(arr) = gguf.metadata("tokenizer.ggml.tokens") {
if let mlx_native::gguf::MetadataValue::Array(elems) = arr {
for (i, el) in elems.iter().enumerate() {
if let mlx_native::gguf::MetadataValue::String(s) = el {
if QWEN35_EOS_STOP_NAMES.contains(&s.as_str()) {
let id = i as u32;
if !eos_token_ids.contains(&id) {
eos_token_ids.push(id);
}
}
}
}
}
}
if eos_token_ids.is_empty() {
eos_token_ids.push(151_645);
}
let _eos_token: u32 = eos_token_ids[0];
tracing::info!(
count = eos_token_ids.len(),
ids = ?eos_token_ids,
"Qwen35 EOS token set resolved (iter-267 multi-source)"
);
let (tokenizer, vision_special_tokens_present) = build_qwen35_serving_tokenizer(&gguf)?;
let template_arch = gguf
.metadata_string("general.architecture")
.unwrap_or("qwen35moe");
let chat_template = gguf
.metadata_string("tokenizer.chat_template")
.map(str::to_string)
.unwrap_or_else(|| {
tracing::warn!(
"Qwen35 load: no GGUF `tokenizer.chat_template`; using pinned QWEN3_CHATML fallback"
);
crate::core::chat_templates::QWEN3_CHATML.to_string()
});
crate::core::chat_templates::validate_tool_chat_template(template_arch, &chat_template)
.map_err(|error| anyhow::anyhow!("Qwen3.6 chat template contract: {error}"))?;
let model_id = gguf
.metadata_string("general.name")
.map(|s| s.to_string())
.unwrap_or_else(|| {
model_path
.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| "qwen35-model".to_string())
});
let hidden_size = model.cfg.hidden_size as usize;
let vocab_size = model.cfg.vocab_size as usize;
let context_length = if model.cfg.max_position_embeddings > 0 {
Some(model.cfg.max_position_embeddings as usize)
} else {
None
};
let quant_type = crate::serve::load_info::infer_quant_label(&gguf);
let load_duration = load_start.elapsed();
let loaded = Self {
model,
tokenizer,
chat_template,
model_id,
model_path: model_path.clone(),
eos_token_ids,
hidden_size,
vocab_size,
context_length,
quant_type,
load_duration,
provenance,
vision_projector_profile,
vision_deepstack_output_count,
vision_special_tokens_present,
prompt_cache: HybridPromptCache::new(),
lcp_registry:
crate::serve::kv_persist::lcp_registry::LcpRegistry::with_byte_budget(
qwen35_lcp_registry_byte_budget(),
),
kv_metrics_sink: None,
disk_persistor: opts.kv_persist_dir.as_ref().and_then(|cache_dir| {
let budget_bytes: u64 = match std::env::var("HF2Q_KV_PERSIST_BUDGET_BYTES") {
Ok(raw) => match raw.trim().parse::<u64>() {
Ok(parsed) => parsed,
Err(err) => {
tracing::warn!(
raw = %raw,
error = %err,
"ADR-027 23d-γ: HF2Q_KV_PERSIST_BUDGET_BYTES parse failed; \
defaulting to 0 (unlimited)"
);
0
}
},
Err(_) => 0,
};
match crate::serve::kv_persist::families::qwen35_disk_persistor::Qwen35DiskPersistor::new_with_budget(cache_dir.clone(), budget_bytes) {
Ok(p) => {
tracing::info!(
cache_dir = %cache_dir.display(),
budget_bytes = budget_bytes,
"ADR-027 iter-6b.2 + 23d-γ: Qwen35DiskPersistor constructed; \
cold-process LCP resume enabled"
);
Some(std::sync::Arc::new(p))
}
Err(e) => {
tracing::warn!(
cache_dir = %cache_dir.display(),
error = %e,
"ADR-027 iter-6b.2: Qwen35DiskPersistor construction failed; \
falling back to in-process-only LCP"
);
None
}
}
}),
lcp_hydrated_for_cfg: std::collections::HashSet::new(),
tq_kv_active: crate::serve::api::tq_packed_descriptor::is_tq_active_mode(),
persistent_kv_cache: None,
speculation: super::qwen35_speculation::QwenSpeculationController::from_environment(),
};
if loaded.tq_kv_active && opts.kv_persist_dir.is_some() {
tracing::info!(
"ADR-027 sub-iter 23d-γ: HF2Q_TQ_KV=1 + HF2Q_KV_PERSIST both active — \
persist snapshots round-trip the compact TQ substrate (codec v5); \
fingerprint is capacity-independent and substrate-namespaced so \
cross-mode hydration is a clean miss."
);
}
if loaded.tq_kv_active {
tracing::info!(
"ADR-027 sub-iter 23d-γ: HF2Q_TQ_KV=1; LCP resume restores TQ buffers \
per-slot (restore_partial 23d-γ) — coherent under the production \
TQ-only regime."
);
}
Ok(loaded)
}
pub fn store_lcp_with_disk_writeback(
&mut self,
kv_cache: &crate::inference::models::qwen35::kv_cache::HybridKvCache,
key: crate::serve::kv_persist::lcp_registry::LcpKey,
prompt_tokens: Vec<u32>,
snapshot: crate::inference::models::qwen35::kv_cache::HybridKvCacheSnapshot,
sliding_window: usize,
linear_capacity: usize,
) -> Result<(), crate::serve::kv_persist::lcp_registry::LcpStoreError> {
let snapshot = std::sync::Arc::new(snapshot);
let disk_job = if let Some(persistor) = &self.disk_persistor {
match crate::serve::kv_persist::families::qwen35_hybrid_persistor::cfg_from_cache(
kv_cache,
crate::serve::kv_persist::families::qwen35_hybrid_persistor::FullAttnCodec::F32Dense,
) {
Ok(cfg) => {
let key_hex = lcp_key_to_filename_hex(&key);
let sidecar = crate::serve::kv_persist::families::qwen35_hybrid_persistor::LcpSidecarMetadata {
model_fingerprint: key.model_fingerprint.clone(),
tenant_id: key.tenant_id.clone(),
params_hash: key.params_hash,
prompt_tokens: prompt_tokens.clone(),
sliding_window: sliding_window as u64,
linear_capacity: linear_capacity as u64,
};
Some((std::sync::Arc::clone(persistor), cfg, key_hex, sidecar))
}
Err(e) => {
tracing::warn!(
cache_dir = %persistor.cache_dir().display(),
error = %format!("{e:#}"),
"ADR-027 iter-6b.3: cfg_from_cache failed; disk \
write-through skipped (in-memory store still proceeds)"
);
None
}
}
} else {
None
};
self.lcp_registry.store(
key,
prompt_tokens,
vec![std::sync::Arc::clone(&snapshot)],
sliding_window,
linear_capacity,
)?;
if let Some((persistor, cfg, key_hex, sidecar)) = disk_job {
match persistor.enqueue_write(cfg, key_hex.clone(), snapshot, sidecar) {
Ok(replaced_pending) => {
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
key_hex = %key_hex,
replaced_pending,
pending_writes = persistor.pending_writes(),
"Qwen35 checkpoint queued for async disk persistence"
);
}
Err(e) => {
tracing::warn!(
cache_dir = %persistor.cache_dir().display(),
key_hex = %key_hex,
error = %format!("{e:#}"),
"Qwen35 async checkpoint enqueue failed; in-memory checkpoint remains live"
);
}
}
}
Ok(())
}
pub fn hydrate_lcp_registry_from_disk(
&mut self,
kv_cache: &crate::inference::models::qwen35::kv_cache::HybridKvCache,
device: &mlx_native::MlxDevice,
) {
let persistor = match &self.disk_persistor {
Some(p) => std::sync::Arc::clone(p),
None => return, };
let cfg = match crate::serve::kv_persist::families::qwen35_hybrid_persistor::cfg_from_cache(
kv_cache,
crate::serve::kv_persist::families::qwen35_hybrid_persistor::FullAttnCodec::F32Dense,
) {
Ok(c) => c,
Err(e) => {
tracing::warn!(
error = %format!("{e:#}"),
"ADR-027 iter-6b.3: hydrate_lcp_registry: cfg_from_cache failed; skipping"
);
return;
}
};
let fingerprint_hex =
crate::serve::kv_persist::families::qwen35_disk_persistor::Qwen35DiskPersistor::fingerprint_hex_for(&cfg);
if self.lcp_hydrated_for_cfg.contains(&fingerprint_hex) {
return; }
self.lcp_hydrated_for_cfg.insert(fingerprint_hex.clone());
let triples = match persistor.hydrate_for_cfg(&cfg, device) {
Ok(v) => v,
Err(e) => {
tracing::warn!(
cache_dir = %persistor.cache_dir().display(),
fingerprint = %fingerprint_hex,
error = %format!("{e:#}"),
"ADR-027 iter-6b.3: hydrate_lcp_registry: hydrate_for_cfg failed; \
skipping (registry remains empty for this cfg)"
);
return;
}
};
let total = triples.len();
let mut inserted = 0usize;
for (key_hex, snap, sidecar) in triples {
let key = crate::serve::kv_persist::lcp_registry::LcpKey {
model_fingerprint: sidecar.model_fingerprint.clone(),
tenant_id: sidecar.tenant_id.clone(),
params_hash: sidecar.params_hash,
};
let sliding_window = sidecar.sliding_window as usize;
let linear_capacity = sidecar.linear_capacity as usize;
match self.lcp_registry.store(
key,
sidecar.prompt_tokens.clone(),
vec![std::sync::Arc::new(snap)],
sliding_window,
linear_capacity,
) {
Ok(()) => inserted += 1,
Err(e) => {
tracing::warn!(
key_hex = %key_hex,
error = ?e,
"ADR-027 iter-6b.3: lcp_registry.store rejected hydrated entry; \
skipping (file may exceed byte budget or be empty)"
);
}
}
}
tracing::info!(
cache_dir = %persistor.cache_dir().display(),
fingerprint = %fingerprint_hex,
files_on_disk = total,
entries_inserted = inserted,
registry_len = self.lcp_registry.len(),
"ADR-027 iter-6b.3: hydrate_lcp_registry_from_disk complete"
);
}
pub fn provision_multi_seq_kv_for_slot_aware(&mut self, max_slots: u32) -> anyhow::Result<()> {
use anyhow::Context;
if max_slots == 0 {
anyhow::bail!(
"ADR-040 C2d: provision_multi_seq_kv_for_slot_aware called with \
max_slots == 0; spawn_with_mode invariant is max_slots >= 1 \
(EngineMode::SlotAware variant enforces this at the API \
boundary — caller violated)"
);
}
let device = mlx_native::MlxDevice::new()
.context("ADR-040 C2d: MlxDevice::new for Qwen35 multi-seq KV provisioning")?;
let max_seq_len = self.model.cfg.max_position_embeddings;
let cache = HybridKvCache::new_with_options(
&self.model.cfg,
&device,
max_seq_len,
max_slots,
self.tq_kv_active,
)
.with_context(|| {
format!(
"ADR-040 C2d: HybridKvCache::new_with_options(max_seq_len={}, \
n_seqs={}, tq_kv_active={}) for Qwen35 SlotAware provisioning",
max_seq_len, max_slots, self.tq_kv_active
)
})?;
self.persistent_kv_cache = Some(cache);
Ok(())
}
}
fn lcp_key_to_filename_hex(key: &crate::serve::kv_persist::lcp_registry::LcpKey) -> String {
use sha2::{Digest, Sha256};
let mut h = Sha256::new();
h.update(b"QH35-lcp-key-fname-v1");
h.update(&key.model_fingerprint.0);
h.update(key.tenant_id.as_bytes());
h.update(&key.params_hash.to_le_bytes());
let digest = h.finalize();
hex::encode(&digest[..16])
}
impl LoadInfoBuilder for Qwen35LoadedModel {
fn build_load_info(
&self,
gguf: &mlx_native::gguf::GgufFile,
load_wall_clock: Duration,
kv_cache_budget_bytes: Option<u64>,
kv_spill_active: bool,
) -> LoadInfo {
let cfg = &self.model.cfg;
LoadInfo {
model_id: self.model_id.clone(),
arch_str: load_info::arch_str_from_gguf(gguf),
arch_family: ArchFamily::Qwen35,
model_path: self.model_path.clone(),
on_disk_bytes: load_info::on_disk_bytes(&self.model_path),
backend_chip: mlx_native::MlxDevice::new()
.map(|d| d.name())
.unwrap_or_else(|_| "Apple GPU".to_string()),
backend: "mlx-native",
n_layers: cfg.num_hidden_layers,
hidden_size: self.hidden_size as u32,
vocab_size: self.vocab_size as u32,
n_attention_heads: cfg.num_attention_heads,
n_key_value_heads: cfg.num_key_value_heads,
head_dim: cfg.head_dim,
sliding_window: None,
full_attention_interval: Some(cfg.full_attention_interval),
max_context_length: self.context_length.map(|v| v as u32),
moe: cfg.moe.as_ref().map(|m| MoeShape {
n_experts: m.num_experts,
n_experts_per_tok: m.num_experts_per_tok,
}),
quant_label: self.quant_type.clone(),
quant_bpw: load_info::compute_bpw(gguf),
tokenizer_source: TokenizerSource::GgufEmbedded,
eos_token_ids: self.eos_token_ids.clone(),
bos_token_id: gguf.metadata_u32("tokenizer.ggml.bos_token_id"),
chat_template_source: if gguf.metadata_string("tokenizer.chat_template").is_some() {
ChatTemplateSource::GgufEmbedded
} else {
ChatTemplateSource::HardcodedFallback {
name: "QWEN3_CHATML",
}
},
provenance: self.provenance.clone(),
vision_projector: None,
load_wall_clock,
resident_weight_bytes: None,
kv_cache_budget_bytes,
kv_spill_active,
tq_kv_active: self.tq_kv_active,
kv_bytes_per_token_override: Some(load_info::qwen35_slot_kv_bytes_per_token(
cfg,
self.tq_kv_active,
)),
kv_fixed_bytes_per_slot_override: Some(load_info::qwen35_fixed_kv_bytes_per_slot(cfg)),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct HybridPromptCacheKey {
pub max_tokens: usize,
pub stop_strings: Vec<String>,
pub tool_argument_wire_kinds: Option<Arc<super::registry::ToolArgumentWireKinds>>,
pub vision_fingerprint: Option<[u8; 32]>,
pub thinking_token_budget: Option<usize>,
pub reasoning_end_tokens: Option<Arc<Vec<u32>>>,
pub reasoning_close_tokens: Option<Arc<Vec<u32>>>,
}
impl HybridPromptCacheKey {
pub fn from_params(params: &SamplingParams) -> Self {
Self {
max_tokens: params.max_tokens,
stop_strings: params.stop_strings.clone(),
tool_argument_wire_kinds: params.tool_argument_wire_kinds.clone(),
vision_fingerprint: params.vision_fingerprint,
thinking_token_budget: params.thinking_token_budget,
reasoning_end_tokens: params.reasoning_end_tokens.clone(),
reasoning_close_tokens: params.reasoning_close_tokens.clone(),
}
}
}
#[derive(Debug)]
pub struct HybridPromptCache {
cached_prompt_tokens: Vec<u32>,
snapshot: Option<HybridKvCacheSnapshot>,
first_decoded_token: u32,
gen_params: Option<HybridPromptCacheKey>,
}
impl Default for HybridPromptCache {
fn default() -> Self {
Self::new()
}
}
impl HybridPromptCache {
pub fn new() -> Self {
Self {
cached_prompt_tokens: Vec::new(),
snapshot: None,
first_decoded_token: 0,
gen_params: None,
}
}
pub fn try_match(&self, new_prompt: &[u32], new_params: &SamplingParams) -> Option<usize> {
if !is_greedy_eligible(new_params) {
return None;
}
if self.snapshot.is_none() || self.gen_params.is_none() {
return None;
}
if self.cached_prompt_tokens.is_empty() {
return None;
}
if self.cached_prompt_tokens.as_slice() != new_prompt {
return None;
}
let request_key = HybridPromptCacheKey::from_params(new_params);
if self.gen_params.as_ref() != Some(&request_key) {
return None;
}
Some(new_prompt.len())
}
pub fn snapshot(&self) -> Option<&HybridKvCacheSnapshot> {
self.snapshot.as_ref()
}
pub fn first_decoded_token(&self) -> u32 {
self.first_decoded_token
}
pub fn update(
&mut self,
prompt: Vec<u32>,
snapshot: HybridKvCacheSnapshot,
first_decoded_token: u32,
params: &SamplingParams,
) {
if !is_greedy_eligible(params) {
return;
}
self.cached_prompt_tokens = prompt;
self.snapshot = Some(snapshot);
self.first_decoded_token = first_decoded_token;
self.gen_params = Some(HybridPromptCacheKey::from_params(params));
}
#[allow(dead_code)]
pub fn clear(&mut self) {
self.cached_prompt_tokens.clear();
self.snapshot = None;
self.first_decoded_token = 0;
self.gen_params = None;
}
pub fn has_entry(&self) -> bool {
self.snapshot.is_some()
}
}
fn is_greedy_eligible(params: &SamplingParams) -> bool {
!(params.temperature > 0.0
|| params.top_k > 0
|| params.top_p < 1.0
|| effective_repetition_penalty(params) != 1.0
|| params.seed.is_some()
|| params.logprobs
|| !params.logit_bias.is_empty()
|| params.grammar.is_some())
}
pub(crate) fn is_qwen_server_speculation_exact_eligible(params: &SamplingParams) -> bool {
!(params.temperature > 0.0
|| params.top_k > 0
|| params.top_p < 1.0
|| params.seed.is_some()
|| params.logprobs
|| !params.logit_bias.is_empty())
&& params.stop_strings.is_empty()
&& params.frequency_penalty == 0.0
&& params.presence_penalty == 0.0
&& params.min_p == 0.0
&& params.top_logprobs == 0
&& !params.parallel_tool_calls
&& (params.tool_call_policy == ToolCallPolicy::Auto || params.grammar.is_some())
}
fn is_serial_mtp_exact_eligible(params: &SamplingParams) -> bool {
is_qwen_server_speculation_exact_eligible(params)
&& !params.reasoning_forced_open
&& params.thinking_token_budget.is_none()
&& params.reasoning_end_tokens.is_none()
&& params.reasoning_close_tokens.is_none()
&& params.grammar.is_none()
&& params.tool_call_policy == ToolCallPolicy::Auto
}
static LCP_STORE_NOTIFY: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
fn lcp_notify_once(bit: u64, msg: std::string::String) {
let prev = LCP_STORE_NOTIFY.fetch_or(bit, std::sync::atomic::Ordering::Relaxed);
if prev & bit == 0 {
eprintln!("{msg}");
}
}
fn lcp_store_skip_notify(stride_aligned: bool, lcp_resume_enabled: bool, mid_store_disabled: bool) {
if !stride_aligned {
return;
}
if !lcp_resume_enabled {
lcp_notify_once(
1,
"[hf2q qwen35 lcp store] SKIPPED: HF2Q_KV_LCP_RESUME is off — \
no prefix checkpoints will be written (once-per-process notice)"
.to_string(),
);
} else if mid_store_disabled {
lcp_notify_once(
2,
"[hf2q qwen35 lcp store] SKIPPED: HF2Q_KV_LCP_DISABLE_MID_STORE=1 \
— mid-prefill checkpoints disabled (once-per-process notice)"
.to_string(),
);
}
}
fn lcp_snapshot_error_notify<E: std::fmt::Debug>(phase: &str, chunk_pos: usize, err: &E) {
lcp_notify_once(
4,
format!(
"[hf2q qwen35 lcp store] snapshot FAILED ({phase}, chunk_pos={chunk_pos}): \
{err:?} — checkpoints are NOT being written; further occurrences \
suppressed (once-per-process notice)"
),
);
}
fn lcp_store_error_notify<E: std::fmt::Debug>(phase: &str, chunk_pos: usize, err: &E) {
lcp_notify_once(
8,
format!(
"[hf2q qwen35 lcp store] store FAILED ({phase}, chunk_pos={chunk_pos}): \
{err:?} — check the registry byte budget; further occurrences \
suppressed (once-per-process notice)"
),
);
}
use crate::inference::models::qwen35::io_heads::greedy_argmax_last_token;
use crate::serve::sampler_pure::{self, SamplingParams as SamplerPureParams};
use mlx_native::MlxDevice;
use super::engine::GenerationResult;
use super::registry::{
ModelRegistration, ReasoningSplitter, SplitSlot, ToolCallEvent, ToolCallSplitter,
};
use super::sse::{DeltaKind, GenerationEvent, StreamStats};
fn prefill_positions_for(prompt_len: usize) -> Vec<i32> {
prefill_positions_from(0, prompt_len)
}
fn prefill_positions_from(start: usize, len: usize) -> Vec<i32> {
let mut flat = vec![0i32; 4 * len];
for axis in 0..4 {
for t in 0..len {
flat[axis * len + t] = (start + t) as i32;
}
}
flat
}
fn requested_kv_cache_capacity(
qwen: &Qwen35LoadedModel,
prompt_len: usize,
max_tokens: usize,
) -> usize {
(prompt_len + max_tokens + 64)
.max(128)
.min(qwen.model.cfg.max_position_embeddings as usize)
}
fn serial_kv_cache_capacity(required: usize, current: usize, maximum: usize) -> usize {
let rounded_required = required
.checked_next_power_of_two()
.unwrap_or(maximum)
.min(maximum);
rounded_required
.max(current.saturating_mul(2).min(maximum))
.max(required)
.min(maximum)
}
fn alloc_serial_kv_cache(
qwen: &Qwen35LoadedModel,
device: &MlxDevice,
capacity: usize,
) -> Result<HybridKvCache> {
let mut cache = HybridKvCache::new_with_options(
&qwen.model.cfg,
device,
capacity as u32,
1,
qwen.tq_kv_active,
)
.context("HybridKvCache::new_with_options (SerialFifo)")?;
cache
.ensure_la_capture(
&qwen.model.cfg,
device,
QWEN35_RECOVERY_CAPTURE_MAX_SUFFIX_TOKENS as u32,
)
.context("SerialFifo recovery capture preallocation")?;
cache.clear_la_capture();
Ok(cache)
}
fn alloc_kv_cache_for_request(
qwen: &Qwen35LoadedModel,
device: &MlxDevice,
prompt_len: usize,
max_tokens: usize,
) -> Result<HybridKvCache> {
let max_seq = requested_kv_cache_capacity(qwen, prompt_len, max_tokens);
HybridKvCache::new_with_options(
&qwen.model.cfg,
device,
max_seq as u32,
1,
qwen.tq_kv_active,
)
.context("HybridKvCache::new_with_options")
}
fn take_serial_kv_cache(
qwen: &mut Qwen35LoadedModel,
device: &MlxDevice,
prompt_len: usize,
max_tokens: usize,
) -> Result<(HybridKvCache, bool)> {
let required = requested_kv_cache_capacity(qwen, prompt_len, max_tokens);
let maximum = qwen.model.cfg.max_position_embeddings as usize;
let current = qwen.persistent_kv_cache.take();
if let Some(cache) = current {
if cache.n_seqs == 1 && cache.max_seq_len as usize >= required {
return Ok((cache, true));
}
let capacity = serial_kv_cache_capacity(required, cache.max_seq_len as usize, maximum);
drop(cache);
return alloc_serial_kv_cache(qwen, device, capacity).map(|cache| (cache, false));
}
let capacity = serial_kv_cache_capacity(required, 0, maximum);
alloc_serial_kv_cache(qwen, device, capacity).map(|cache| (cache, false))
}
fn sample_logits_qwen35(logits: &mut [f32], params: &SamplingParams, generated: &[u32]) -> u32 {
let sp = SamplerPureParams {
temperature: params.temperature as f64,
top_p: params.top_p as f64,
top_k: params.top_k,
min_p: 0.0,
repetition_penalty: effective_repetition_penalty(params),
max_tokens: params.max_tokens,
seed: params.seed,
};
sampler_pure::sample_token(logits, &sp, generated)
}
fn sample_logits_qwen35_with_logprob(
logits: &mut [f32],
params: &SamplingParams,
generated: &[u32],
) -> (u32, f32) {
let sp = SamplerPureParams {
temperature: params.temperature as f64,
top_p: params.top_p as f64,
top_k: params.top_k,
min_p: 0.0,
repetition_penalty: effective_repetition_penalty(params),
max_tokens: params.max_tokens,
seed: params.seed,
};
sampler_pure::sample_token_with_logprob(logits, &sp, generated)
}
fn sample_logits_qwen35_constrained(
logits: &mut [f32],
params: &SamplingParams,
generated: &[u32],
runtime: Option<&super::grammar::GrammarRuntime>,
want_logprobs: bool,
) -> Result<(u32, Option<f32>)> {
for (&token, &bias) in ¶ms.logit_bias {
if let Some(logit) = logits.get_mut(token as usize) {
*logit += bias;
}
}
let sampler = SamplerPureParams {
temperature: params.temperature as f64,
top_p: params.top_p as f64,
top_k: params.top_k,
min_p: params.min_p as f64,
repetition_penalty: effective_repetition_penalty(params),
max_tokens: params.max_tokens,
seed: params.seed,
};
sample_logits_with_grammar(
logits,
&sampler,
generated,
runtime,
params.token_bytes.as_deref().map(Vec::as_slice),
want_logprobs,
)
}
fn advance_qwen35_grammar(
runtime: &mut Option<super::grammar::GrammarRuntime>,
params: &SamplingParams,
token: u32,
) {
if let (Some(runtime), Some(token_bytes)) = (runtime.as_mut(), params.token_bytes.as_deref()) {
if let Some(bytes) = token_bytes.get(token as usize) {
if !bytes.is_empty() {
runtime.accept_bytes(bytes);
}
}
}
}
fn qwen35_grammar_terminal_token(
runtime: Option<&super::grammar::GrammarRuntime>,
params: &SamplingParams,
token: u32,
) -> bool {
let Some(runtime) = runtime else {
return false;
};
if runtime.is_dead() {
return true;
}
runtime.is_accepted()
&& params.token_bytes.as_deref().is_some_and(|token_bytes| {
token_bytes
.get(token as usize)
.map(|bytes| bytes.is_empty())
.unwrap_or(true)
})
}
fn build_lcp_key_for_qwen35(
qwen: &Qwen35LoadedModel,
_params: &SamplingParams,
) -> crate::serve::kv_persist::lcp_registry::LcpKey {
use crate::serve::kv_persist::format::compute_model_fingerprint;
let (producer_version, source_sha256) = match &qwen.provenance {
crate::core::provenance::Provenance::Hf2q {
producer_version,
source_sha256,
..
} => (producer_version.as_str(), source_sha256.as_str()),
crate::core::provenance::Provenance::External => ("", ""),
};
let chat_template_hash = match &qwen.provenance {
crate::core::provenance::Provenance::Hf2q { .. } => {
super::kv_spill_descriptor::KvSpillProvenance::hash_chat_template(&qwen.chat_template)
}
crate::core::provenance::Provenance::External => String::new(),
};
let quant = qwen.quant_type.as_deref().unwrap_or("");
let fp = compute_model_fingerprint(
&qwen.model_id,
quant,
producer_version,
source_sha256,
&chat_template_hash,
);
crate::serve::kv_persist::lcp_registry::LcpKey {
model_fingerprint: fp,
tenant_id: String::new(),
params_hash: 0,
}
}
fn build_lcp_key_for_qwen35_chunk(
qwen: &Qwen35LoadedModel,
params: &SamplingParams,
chunk_position: usize,
) -> crate::serve::kv_persist::lcp_registry::LcpKey {
let mut key = build_lcp_key_for_qwen35(qwen, params);
if chunk_position > 0 {
key.tenant_id = format!("qwen35:lcp_chunk:{chunk_position}");
}
key
}
fn lookup_qwen35_resume_checkpoint<T>(
registry: &mut crate::serve::kv_persist::lcp_registry::LcpRegistry<T>,
base_key: &crate::serve::kv_persist::lcp_registry::LcpKey,
new_tokens: &[u32],
stride: usize,
) -> Option<(crate::serve::kv_persist::lcp_registry::LcpPrefix<T>, usize)>
where
T: Send + Sync + 'static + crate::serve::kv_persist::lcp_registry::ByteSized,
{
if new_tokens.is_empty() {
return None;
}
let base_match = registry
.lookup(base_key, new_tokens)
.filter(|prefix| prefix.k == prefix.cached_prompt_len && prefix.k < new_tokens.len());
let base_len = base_match.as_ref().map(|prefix| prefix.k).unwrap_or(0);
if stride > 0 {
let mut chunk_pos = (new_tokens.len() / stride).saturating_mul(stride);
while chunk_pos >= stride && chunk_pos > base_len {
let mut chunk_key = base_key.clone();
chunk_key.tenant_id = format!("qwen35:lcp_chunk:{chunk_pos}");
if let Some(prefix) = registry.lookup(&chunk_key, new_tokens) {
if prefix.k == prefix.cached_prompt_len && prefix.k < new_tokens.len() {
return Some((prefix, chunk_pos));
}
}
if chunk_pos == stride {
break;
}
chunk_pos -= stride;
}
}
base_match.map(|prefix| (prefix, 0))
}
const QWEN35_RECOVERY_TAIL_FALLBACK_TOKENS: usize = 64;
const QWEN35_RECOVERY_CAPTURE_MAX_SUFFIX_TOKENS: usize = 32;
fn qwen35_recovery_tail_tokens(
qwen: &Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
) -> usize {
let unstable_suffix = if params.reasoning_forced_open {
"<think>\n"
} else {
"<think>\n\n</think>\n\n"
};
let encoded = qwen.tokenizer.encode(unstable_suffix, false);
let suffix_tokens = encoded.as_ref().map(|e| e.get_ids()).unwrap_or(&[]);
recovery_tail_for_suffix(
prompt_tokens,
suffix_tokens,
QWEN35_RECOVERY_TAIL_FALLBACK_TOKENS,
)
}
fn recovery_tail_for_suffix(
prompt_tokens: &[u32],
suffix_tokens: &[u32],
fallback: usize,
) -> usize {
if !suffix_tokens.is_empty() && prompt_tokens.ends_with(suffix_tokens) {
suffix_tokens.len()
} else {
fallback
}
}
fn stride_checkpoint_superseded_by_recovery_anchor(
recovery_eligible: bool,
k_end: usize,
stride: usize,
recovery_anchor: usize,
) -> bool {
recovery_eligible && k_end <= recovery_anchor && k_end.saturating_add(stride) > recovery_anchor
}
fn qwen35_recovery_capture_plan(
lcp_resume_start: usize,
recovery_anchor: usize,
prompt_len: usize,
recovery_eligible: bool,
chunked_eligible: bool,
) -> Option<(usize, usize)> {
if !recovery_eligible
|| chunked_eligible
|| recovery_anchor <= lcp_resume_start
|| prompt_len <= recovery_anchor
{
return None;
}
let suffix_len = prompt_len.checked_sub(lcp_resume_start)?;
if suffix_len == 0 || suffix_len > QWEN35_RECOVERY_CAPTURE_MAX_SUFFIX_TOKENS {
return None;
}
let capture_index = recovery_anchor.checked_sub(lcp_resume_start + 1)?;
Some((suffix_len, capture_index))
}
fn store_qwen35_latest_turn_checkpoint(
qwen: &mut Qwen35LoadedModel,
kv_cache: &HybridKvCache,
device: &MlxDevice,
params: &SamplingParams,
prompt_tokens: &[u32],
anchor: usize,
phase: &str,
capture_index: Option<usize>,
) {
let snapshot_started = Instant::now();
let snapshot_result = match capture_index {
Some(index) => kv_cache.snapshot_prefix_from_capture(device, anchor, index),
None => kv_cache.snapshot_prefix(device, anchor),
};
match snapshot_result {
Ok(snapshot) => {
let snapshot_elapsed = snapshot_started.elapsed();
let snapshot_bytes = snapshot.total_bytes();
let base_key = build_lcp_key_for_qwen35(qwen, params);
let linear_capacity = kv_cache
.linear_attn
.first()
.map(|slot| slot.recurrent.byte_len())
.unwrap_or(0);
let disk_enabled = qwen.disk_persistor.is_some();
let store_started = Instant::now();
let store_result = qwen.store_lcp_with_disk_writeback(
kv_cache,
base_key,
prompt_tokens[..anchor].to_vec(),
snapshot,
0,
linear_capacity,
);
let store_elapsed = store_started.elapsed();
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
phase,
anchor_tokens = anchor,
snapshot_bytes,
snapshot_ms = snapshot_elapsed.as_secs_f64() * 1_000.0,
store_ms = store_elapsed.as_secs_f64() * 1_000.0,
disk_enabled,
capture_index,
"Qwen35 latest-turn checkpoint prepared"
);
if let Err(error) = store_result {
lcp_store_error_notify(phase, anchor, &error);
}
}
Err(error) => lcp_snapshot_error_notify(phase, anchor, &error),
}
}
pub fn qwen35_lcp_registry_byte_budget() -> u64 {
crate::serve::kv_persist::lcp_registry::default_lcp_byte_budget()
}
fn qwen35_reported_cached_tokens(
prompt_len: usize,
prompt_cache_hit: bool,
lcp_resume_start: usize,
) -> usize {
if prompt_cache_hit {
prompt_len
} else {
lcp_resume_start.min(prompt_len)
}
}
fn qwen35_hit_stop_string(text: &str, stops: &[String]) -> bool {
if stops.is_empty() {
return false;
}
stops
.iter()
.any(|s| !s.is_empty() && text.ends_with(s.as_str()))
}
fn qwen35_strip_trailing_stop(text: &mut String, stops: &[String]) {
for s in stops {
if !s.is_empty() && text.ends_with(s) {
let new_len = text.len() - s.len();
text.truncate(new_len);
return;
}
}
}
fn generate_qwen35_once_ordinary(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
supervisor: &EngineSupervisor,
) -> Result<GenerationResult> {
let request_start = Instant::now();
let _disk_request_guard = qwen
.disk_persistor
.as_ref()
.map(|persistor| persistor.begin_request());
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen35_once: empty prompt_tokens"
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible(params);
let want_logprobs = params.logprobs;
let mut logprobs_vec: Option<Vec<f32>> = if want_logprobs {
Some(Vec::with_capacity(max_tokens))
} else {
None
};
let mut grammar_runtime = grammar_runtime_for_request(params, registration)?;
let cache_alloc_start = Instant::now();
let device =
MlxDevice::new().map_err(|e| anyhow::anyhow!("MlxDevice::new (qwen35 generate): {e}"))?;
let (mut kv_cache, cache_reused) = take_serial_kv_cache(qwen, &device, prompt_len, max_tokens)?;
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "unary",
prompt_tokens = prompt_len,
max_tokens,
cache_capacity_tokens = kv_cache.max_seq_len,
cache_reused,
tq_kv = qwen.tq_kv_active,
elapsed_ms = cache_alloc_start.elapsed().as_secs_f64() * 1000.0,
"Qwen35 request cache ready"
);
qwen.hydrate_lcp_registry_from_disk(&kv_cache, &device);
let prompt_cache_hit = qwen.prompt_cache.try_match(prompt_tokens, params).is_some();
let mut lcp_resume_start: usize = 0;
if !prompt_cache_hit {
let stride_for_observe = crate::debug::INVESTIGATION_ENV.kv_lcp_deltanet_checkpoint_stride;
let base_key_for_observe = build_lcp_key_for_qwen35(qwen, params);
let detected = lookup_qwen35_resume_checkpoint(
&mut qwen.lcp_registry,
&base_key_for_observe,
prompt_tokens,
stride_for_observe,
)
.map(|(prefix, _chunk_pos)| prefix.k);
if let Some(sink) = qwen.kv_metrics_sink.as_ref() {
sink.record_lcp_probe(detected);
}
let _ = detected;
let lcp_resume_enabled = crate::serve::api::engine::effective_kv_lcp_resume(
crate::debug::INVESTIGATION_ENV.kv_lcp_resume,
true,
);
if lcp_resume_enabled {
let stride = crate::debug::INVESTIGATION_ENV.kv_lcp_deltanet_checkpoint_stride;
let base_key = build_lcp_key_for_qwen35(qwen, params);
eprintln!(
"[hf2q qwen35 lcp probe] enabled, registry_len={}, prompt_len={}, \
stride={}, scanning latest-turn + stride checkpoints",
qwen.lcp_registry.len(),
prompt_tokens.len(),
stride,
);
if let Some((prefix, chunk_pos)) = lookup_qwen35_resume_checkpoint(
&mut qwen.lcp_registry,
&base_key,
prompt_tokens,
stride,
) {
let snapshot: &HybridKvCacheSnapshot = &prefix.dense_kvs[0];
let restore_start = Instant::now();
kv_cache
.restore_partial(snapshot, prefix.k)
.context("qwen35 lcp_registry restore_partial")?;
let restore_ms = restore_start.elapsed().as_micros() as f64 / 1000.0;
lcp_resume_start = prefix.k;
let checkpoint = if chunk_pos == 0 {
"LATEST-TURN"
} else {
"STRIDE-ALIGNED"
};
eprintln!(
"[hf2q qwen35 lcp resume] {checkpoint} HIT — restoring at \
k={} (cached_prompt_len={}, chunk_pos={}, restore_ms={:.3})",
prefix.k, prefix.cached_prompt_len, chunk_pos, restore_ms
);
} else {
eprintln!(
"[hf2q qwen35 lcp probe] no compatible checkpoint \
(registry_len={})",
qwen.lcp_registry.len()
);
}
}
}
if cache_reused && !prompt_cache_hit && lcp_resume_start == 0 {
kv_cache.reset();
}
let prefill_start = Instant::now();
let mut next_token: u32;
if prompt_cache_hit {
let snap = qwen
.prompt_cache
.snapshot()
.expect("try_match returned Some implies snapshot Some");
kv_cache
.restore_partial(snap, prompt_len)
.context("prompt_cache restore_partial")?;
next_token = qwen.prompt_cache.first_decoded_token();
tracing::debug!(
"qwen35 prompt_cache: HIT — {} tokens; prefill skipped",
prompt_len
);
} else {
let stride = crate::debug::INVESTIGATION_ENV.kv_lcp_deltanet_checkpoint_stride;
let lcp_resume_enabled = crate::debug::INVESTIGATION_ENV.kv_lcp_resume;
let chunked_eligible = stride > 0
&& prompt_len > stride
&& (lcp_resume_start == 0 || lcp_resume_start % stride == 0)
&& crate::debug::INVESTIGATION_ENV.kv_lcp_chunked_prefill;
let recovery_tail_tokens = qwen35_recovery_tail_tokens(qwen, prompt_tokens, params);
let recovery_anchor = prompt_len.saturating_sub(recovery_tail_tokens);
let recovery_eligible = lcp_resume_enabled
&& recovery_anchor > lcp_resume_start
&& recovery_anchor >= 16;
let recovery_capture_plan = qwen35_recovery_capture_plan(
lcp_resume_start,
recovery_anchor,
prompt_len,
recovery_eligible,
chunked_eligible,
);
let prefill_logits = if let Some((suffix_len, capture_index)) = recovery_capture_plan {
let capture_alloc_start = Instant::now();
kv_cache
.ensure_la_capture(&qwen.model.cfg, &device, suffix_len as u32)
.context("Qwen35 latest-turn recovery capture allocation")?;
let capture_alloc_ms = capture_alloc_start.elapsed().as_secs_f64() * 1000.0;
let position_build_start = Instant::now();
let suffix_tokens = &prompt_tokens[lcp_resume_start..];
let mut suffix_positions = vec![0i32; 4 * suffix_len];
for axis in 0..4 {
for token in 0..suffix_len {
suffix_positions[axis * suffix_len + token] = (lcp_resume_start + token) as i32;
}
}
let position_build_ms = position_build_start.elapsed().as_secs_f64() * 1000.0;
let forward_start = Instant::now();
let logits = supervised_gpu_call(supervisor, "qwen35_serial_prefill", || {
qwen.model
.forward_gpu_last_logits(
suffix_tokens,
&suffix_positions,
&mut kv_cache,
SlotId(0),
)
.context("Qwen35 captured latest-turn suffix prefill")
})?;
let forward_ms = forward_start.elapsed().as_secs_f64() * 1000.0;
let checkpoint_start = Instant::now();
store_qwen35_latest_turn_checkpoint(
qwen,
&kv_cache,
&device,
params,
prompt_tokens,
recovery_anchor,
"captured latest-turn recovery-anchor",
Some(capture_index),
);
let checkpoint_ms = checkpoint_start.elapsed().as_secs_f64() * 1000.0;
let capture_clear_start = Instant::now();
kv_cache.clear_la_capture();
let capture_clear_ms = capture_clear_start.elapsed().as_secs_f64() * 1000.0;
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "unary",
suffix_tokens = suffix_len,
capture_index,
capture_alloc_ms,
position_build_ms,
forward_ms,
checkpoint_ms,
capture_clear_ms,
"Qwen35 captured recovery prefill phase timing"
);
eprintln!(
"[hf2q qwen35 lcp store] captured latest-turn recovery anchor={} \
suffix_tokens={} capture_index={}",
recovery_anchor, suffix_len, capture_index
);
logits
} else if recovery_eligible && !chunked_eligible {
let prefix_tokens = &prompt_tokens[lcp_resume_start..recovery_anchor];
let prefix_len = prefix_tokens.len();
let mut prefix_positions = vec![0i32; 4 * prefix_len];
for axis in 0..4 {
for token in 0..prefix_len {
prefix_positions[axis * prefix_len + token] = (lcp_resume_start + token) as i32;
}
}
supervised_gpu_call(supervisor, "qwen35_serial_prefill", || {
qwen.model
.forward_gpu_last_logits(
prefix_tokens,
&prefix_positions,
&mut kv_cache,
SlotId(0),
)
.context("Qwen35 latest-turn recovery-anchor prefix prefill")
})?;
store_qwen35_latest_turn_checkpoint(
qwen,
&kv_cache,
&device,
params,
prompt_tokens,
recovery_anchor,
"latest-turn recovery-anchor",
None,
);
let tail_tokens = &prompt_tokens[recovery_anchor..];
let tail_len = tail_tokens.len();
let mut tail_positions = vec![0i32; 4 * tail_len];
for axis in 0..4 {
for token in 0..tail_len {
tail_positions[axis * tail_len + token] = (recovery_anchor + token) as i32;
}
}
eprintln!(
"[hf2q qwen35 lcp store] latest-turn recovery anchor={} tail_tokens={}",
recovery_anchor, tail_len
);
supervised_gpu_call(supervisor, "qwen35_serial_prefill", || {
qwen.model
.forward_gpu_last_logits(tail_tokens, &tail_positions, &mut kv_cache, SlotId(0))
.context("Qwen35 latest-turn recovery-anchor tail prefill")
})?
} else if chunked_eligible {
anyhow::ensure!(
lcp_resume_start % stride == 0,
"qwen35 chunked prefill: lcp_resume_start ({}) must be \
stride-aligned ({}) — snapshots are stored only at \
stride boundaries, so this should be guaranteed by the \
probe site",
lcp_resume_start,
stride
);
let first_chunk_idx = lcp_resume_start / stride;
let chunked_prefill_end = if recovery_eligible {
recovery_anchor
} else {
prompt_len
};
let n_chunks = (chunked_prefill_end + stride - 1) / stride;
eprintln!(
"[hf2q qwen35 chunked prefill] {} chunks (stride={}, \
prompt_len={}, prefill_end={}, first_chunk_idx={})",
n_chunks, stride, prompt_len, chunked_prefill_end, first_chunk_idx
);
let mut last_logits: Option<Vec<f32>> = None;
for chunk_idx in first_chunk_idx..n_chunks {
let k_start = chunk_idx * stride;
let k_end = ((chunk_idx + 1) * stride).min(chunked_prefill_end);
let chunk_seq_len = k_end - k_start;
let chunk_tokens = &prompt_tokens[k_start..k_end];
let mut chunk_positions = vec![0i32; 4 * chunk_seq_len];
for axis in 0..4 {
for t in 0..chunk_seq_len {
chunk_positions[axis * chunk_seq_len + t] = (k_start + t) as i32;
}
}
let logits = supervised_gpu_call(
supervisor,
"qwen35_serial_prefill_chunk",
|| {
qwen.model
.forward_gpu_last_logits(
chunk_tokens,
&chunk_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!(
"qwen35 chunked prefill: chunk {}/{} (k_start={}, k_end={}, seq_len={})",
chunk_idx + 1, n_chunks, k_start, k_end, chunk_seq_len
)
})
},
)?;
if chunk_idx == n_chunks - 1 {
last_logits = Some(logits.clone());
}
let stride_aligned = k_end % stride == 0;
let superseded_by_recovery_anchor = stride_checkpoint_superseded_by_recovery_anchor(
recovery_eligible,
k_end,
stride,
recovery_anchor,
);
let mid_store_disabled =
std::env::var("HF2Q_KV_LCP_DISABLE_MID_STORE").as_deref() == Ok("1");
lcp_store_skip_notify(
stride_aligned && !superseded_by_recovery_anchor,
lcp_resume_enabled,
mid_store_disabled,
);
if lcp_resume_enabled
&& stride_aligned
&& !superseded_by_recovery_anchor
&& !mid_store_disabled
{
match kv_cache.snapshot_prefix(&device, k_end) {
Ok(snap) => {
let chunk_key = build_lcp_key_for_qwen35_chunk(qwen, params, k_end);
let linear_capacity = kv_cache
.linear_attn
.first()
.map(|s| s.recurrent.byte_len())
.unwrap_or(0);
if let Err(e) = qwen.store_lcp_with_disk_writeback(
&kv_cache,
chunk_key,
prompt_tokens[..k_end].to_vec(),
snap,
0,
linear_capacity,
) {
lcp_store_error_notify("mid-prefill", k_end, &e);
} else {
eprintln!(
"[hf2q qwen35 lcp store] mid-prefill snapshot \
at chunk_pos={k_end} (registry_len_after={})",
qwen.lcp_registry.len()
);
}
}
Err(e) => {
lcp_snapshot_error_notify("mid-prefill", k_end, &e);
}
}
}
}
if recovery_eligible {
store_qwen35_latest_turn_checkpoint(
qwen,
&kv_cache,
&device,
params,
prompt_tokens,
recovery_anchor,
"chunked latest-turn recovery-anchor",
None,
);
let tail_tokens = &prompt_tokens[recovery_anchor..];
let tail_len = tail_tokens.len();
let mut tail_positions = vec![0i32; 4 * tail_len];
for axis in 0..4 {
for token in 0..tail_len {
tail_positions[axis * tail_len + token] = (recovery_anchor + token) as i32;
}
}
eprintln!(
"[hf2q qwen35 lcp store] latest-turn recovery anchor={} tail_tokens={}",
recovery_anchor, tail_len
);
last_logits = Some(supervised_gpu_call(
supervisor,
"qwen35_serial_prefill",
|| {
qwen.model
.forward_gpu_last_logits(
tail_tokens,
&tail_positions,
&mut kv_cache,
SlotId(0),
)
.context("Qwen35 chunked recovery-anchor tail prefill")
},
)?);
}
last_logits.expect("at least one chunk in chunked prefill")
} else if lcp_resume_start > 0 {
let suffix_tokens = &prompt_tokens[lcp_resume_start..];
let suffix_len = suffix_tokens.len();
let mut suffix_positions = vec![0i32; 4 * suffix_len];
for axis in 0..4 {
for t in 0..suffix_len {
suffix_positions[axis * suffix_len + t] = (lcp_resume_start + t) as i32;
}
}
eprintln!(
"[hf2q qwen35 lcp resume] suffix prefill {} tokens \
(lcp_resume_start={}, prompt_len={})",
suffix_len, lcp_resume_start, prompt_len
);
supervised_gpu_call(supervisor, "qwen35_serial_prefill", || {
qwen.model
.forward_gpu_last_logits(
suffix_tokens,
&suffix_positions,
&mut kv_cache,
SlotId(0),
)
.context("Qwen35Model::forward_gpu_last_logits (LCP resume suffix)")
})?
} else {
let positions = prefill_positions_for(prompt_len);
supervised_gpu_call(supervisor, "qwen35_serial_prefill", || {
qwen.model
.forward_gpu_last_logits(prompt_tokens, &positions, &mut kv_cache, SlotId(0))
.context("Qwen35Model::forward_gpu_last_logits (prefill)")
})?
};
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
if is_greedy {
next_token = greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32);
} else {
let mut logits = prefill_logits.clone();
let (token, logprob) = sample_logits_qwen35_constrained(
&mut logits,
params,
&[],
grammar_runtime.as_ref(),
want_logprobs,
)?;
next_token = token;
if let (Some(values), Some(logprob)) = (logprobs_vec.as_mut(), logprob) {
values.push(logprob);
}
}
advance_qwen35_grammar(&mut grammar_runtime, params, next_token);
if is_greedy {
let prompt_snapshot_start = Instant::now();
match kv_cache.snapshot_prefix(&device, prompt_len) {
Ok(snap) => {
let snapshot_bytes = snap.total_bytes();
qwen.prompt_cache
.update(prompt_tokens.to_vec(), snap, next_token, params);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
phase = "full-prompt replay",
snapshot_bytes,
elapsed_ms = prompt_snapshot_start.elapsed().as_secs_f64() * 1000.0,
"Qwen35 prompt-cache checkpoint prepared"
);
}
Err(e) => {
eprintln!("[hf2q qwen35 lcp store] prompt_cache snapshot failed: {e}");
}
}
}
}
let prefill_duration = prefill_start.elapsed();
let reported_cached_tokens =
qwen35_reported_cached_tokens(prompt_len, prompt_cache_hit, lcp_resume_start);
let prefill_work_tokens = prompt_len.saturating_sub(reported_cached_tokens);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "unary",
prompt_tokens = prompt_len,
cached_tokens = reported_cached_tokens,
work_tokens = prefill_work_tokens,
elapsed_ms = prefill_duration.as_secs_f64() * 1000.0,
tokens_per_second = if prefill_duration.is_zero() {
0.0
} else {
prefill_work_tokens as f64 / prefill_duration.as_secs_f64()
},
"Qwen35 prefill complete"
);
let decode_start = Instant::now();
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let first_fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut decoded_text = first_fragment.clone();
let mut finish_reason: &'static str = "length";
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
} else {
for step in 1..max_tokens {
let pos = (prompt_len + step - 1) as i32;
if pos as u32 >= kv_cache.max_seq_len {
tracing::warn!(
pos,
max_seq = kv_cache.max_seq_len,
"qwen35 decode: hit kv-cache bound; stopping with finish=length",
);
break;
}
let decode_positions = vec![pos; 4];
next_token = if is_greedy {
supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_greedy(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| format!("forward_gpu_greedy decode step {step}"))
})?
} else {
let logits_full = supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_last_logits(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| format!("forward_gpu_last_logits decode step {step}"))
})?;
let mut logits = logits_full;
let (token, logprob) = sample_logits_qwen35_constrained(
&mut logits,
params,
&generated_tokens,
grammar_runtime.as_ref(),
want_logprobs,
)?;
if let (Some(values), Some(logprob)) = (logprobs_vec.as_mut(), logprob) {
values.push(logprob);
}
token
};
advance_qwen35_grammar(&mut grammar_runtime, params, next_token);
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
if qwen35_grammar_terminal_token(grammar_runtime.as_ref(), params, next_token) {
finish_reason = "stop";
break;
}
generated_tokens.push(next_token);
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
decoded_text.push_str(&fragment);
if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
break;
}
}
}
let decode_duration = decode_start.elapsed();
qwen.persistent_kv_cache = Some(kv_cache);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "unary",
generated_tokens = generated_tokens.len(),
elapsed_ms = decode_duration.as_secs_f64() * 1000.0,
tokens_per_second = if decode_duration.is_zero() {
0.0
} else {
generated_tokens.len() as f64 / decode_duration.as_secs_f64()
},
"Qwen35 decode complete"
);
let (content, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (decoded_text, None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &tok in &generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "unary",
prompt_tokens = prompt_len,
cached_tokens = reported_cached_tokens,
completion_tokens = generated_tokens.len(),
total_ms = request_start.elapsed().as_secs_f64() * 1000.0,
"Qwen35 request complete"
);
Ok(GenerationResult {
text: content,
reasoning_text,
prompt_tokens: prompt_len,
completion_tokens: generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: reported_cached_tokens,
logprobs: logprobs_vec,
})
}
pub(super) fn generate_qwen35_once(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
supervisor: &EngineSupervisor,
) -> Result<GenerationResult> {
use super::qwen35_speculation::{self, QwenSpeculationDecision};
let mut decision = qwen.speculation.decide(
prompt_tokens,
is_serial_mtp_exact_eligible(params),
qwen.model.mtp.is_some(),
qwen.prompt_cache.try_match(prompt_tokens, params).is_some(),
);
if decision == QwenSpeculationDecision::Eligible {
match generate_qwen35_once_mtp(qwen, prompt_tokens, params, registration) {
Ok((result, stats)) => {
qwen35_speculation::record_outcome(
stats.proposed,
stats.accepted,
stats.rejected,
stats.target_forwards,
result.cached_tokens,
);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::speculation",
drafted_tokens = stats.proposed,
accepted_tokens = stats.accepted,
rejected_tokens = stats.rejected,
target_forwards = stats.target_forwards,
cached_tokens = result.cached_tokens,
"Qwen native MTP transaction complete"
);
return Ok(result);
}
Err(error) => {
tracing::warn!(
error = %error,
"Qwen native MTP unavailable; falling back to ordinary decode"
);
decision = QwenSpeculationDecision::RuntimeUnavailable;
}
}
}
qwen35_speculation::record_fallback(decision);
let result =
generate_qwen35_once_ordinary(qwen, prompt_tokens, params, registration, supervisor)?;
qwen35_speculation::record_outcome(0, 0, 0, 0, result.cached_tokens);
tracing::debug!(
target: "hf2q::serve::api::engine_qwen35::speculation",
?decision,
cached_tokens = result.cached_tokens,
"Qwen ordinary decode selected"
);
Ok(result)
}
fn generate_qwen35_once_mtp(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
) -> Result<(
GenerationResult,
crate::inference::models::qwen35::spec_decode::SpecDecodeStats,
)> {
use crate::inference::models::qwen35::spec_decode::SpecDecode;
anyhow::ensure!(!prompt_tokens.is_empty(), "Qwen MTP: empty prompt");
anyhow::ensure!(
is_serial_mtp_exact_eligible(params),
"Qwen MTP: request has unsupported sampling, grammar, tool, or thinking semantics"
);
anyhow::ensure!(
qwen.model.mtp.is_some(),
"Qwen MTP: model has no MTP weights"
);
let max_tokens = params.max_tokens.max(1);
let result = SpecDecode::run_with_eos_set(
&qwen.model,
prompt_tokens,
max_tokens,
qwen.eos_token_ids.clone(),
qwen.model.cfg.max_position_embeddings,
)?;
let mut decoded_text = qwen
.tokenizer
.decode(&result.tokens, false)
.unwrap_or_default();
let finish_reason = if result.tokens.len() < max_tokens {
"stop"
} else {
"length"
};
let (content, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (std::mem::take(&mut decoded_text), None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut splitter =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &token in &result.tokens {
let fragment = qwen.tokenizer.decode(&[token], false).unwrap_or_default();
if let Some(splitter) = splitter.as_mut() {
let _ = splitter.feed(&fragment);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
Ok((
GenerationResult {
text: content,
reasoning_text,
prompt_tokens: prompt_tokens.len(),
completion_tokens: result.tokens.len(),
reasoning_tokens: (reasoning_token_count > 0).then_some(reasoning_token_count),
finish_reason,
prefill_duration: result.stats.prefill_elapsed,
decode_duration: result.stats.decode_elapsed,
cached_tokens: 0,
logprobs: None,
},
result.stats,
))
}
pub fn generate_qwen35_once_slot_aware(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
) -> Result<GenerationResult> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen35_once_slot_aware: empty prompt_tokens"
);
anyhow::ensure!(
slot_id.0 < kv_cache.n_seqs,
"generate_qwen35_once_slot_aware: SlotOutOfRange slot={} max_slots={} \
(ADR-040 iter-C2d-cont-kernel iter-1)",
slot_id.0,
kv_cache.n_seqs,
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let need_seq = prompt_len + max_tokens + 64;
if need_seq > kv_cache.max_seq_len as usize {
return Err(anyhow::anyhow!(
"generate_qwen35_once_slot_aware: per-request need_seq={} exceeds \
persistent cache max_seq_len={} (slot={} prompt_len={} max_tokens={}). \
ADR-040 iter-C2d-cont-kernel iter-1 sizes the persistent cache to \
cfg.max_position_embeddings; reduce max_tokens or use a shorter prompt.",
need_seq,
kv_cache.max_seq_len,
slot_id.0,
prompt_len,
max_tokens
));
}
let is_greedy = is_greedy_eligible(params);
let want_logprobs = params.logprobs;
let mut logprobs_vec: Option<Vec<f32>> = if want_logprobs {
Some(Vec::with_capacity(max_tokens))
} else {
None
};
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-1: reset_for_slot at entry")?;
let prompt_cache_hit = qwen.prompt_cache.try_match(prompt_tokens, params).is_some();
let prefill_start = Instant::now();
let next_token: u32;
if prompt_cache_hit {
let snap = qwen
.prompt_cache
.snapshot()
.expect("try_match returned Some implies snapshot Some");
kv_cache
.restore_partial(snap, prompt_len)
.context("ADR-040 iter-C2d-cont-kernel iter-1: prompt_cache restore_partial")?;
next_token = qwen.prompt_cache.first_decoded_token();
tracing::debug!(
"qwen35 slot-aware prompt_cache: HIT slot={} prompt_len={} prefill skipped",
slot_id.0,
prompt_len
);
} else {
let positions = prefill_positions_for(prompt_len);
let prefill_logits = qwen
.model
.forward_gpu_last_logits(prompt_tokens, &positions, kv_cache, slot_id)
.context("Qwen35Model::forward_gpu_last_logits (slot-aware prefill)")?;
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 slot-aware prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
if is_greedy && !want_logprobs {
next_token = greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32);
} else {
let mut logits = prefill_logits;
if let Some(ref mut lps) = logprobs_vec {
let (tok, lp) = sample_logits_qwen35_with_logprob(&mut logits, params, &[]);
lps.push(lp);
next_token = tok;
} else {
next_token = sample_logits_qwen35(&mut logits, params, &[]);
}
}
}
let prefill_duration = prefill_start.elapsed();
let device = MlxDevice::new()
.map_err(|e| anyhow::anyhow!("MlxDevice::new (qwen35 slot-aware decode): {e}"))?;
let _ = &device; let decode_start = Instant::now();
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let mut decoded_text = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let stops = ¶ms.stop_strings;
let mut finish_reason: &'static str = "length";
if qwen.eos_token_ids.contains(&next_token) {
generated_tokens.pop();
decoded_text.clear();
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, stops) {
qwen35_strip_trailing_stop(&mut decoded_text, stops);
finish_reason = "stop";
}
let mut step = 1usize;
while step < max_tokens && finish_reason == "length" {
let pos = prompt_len + step - 1;
let pos_i32 = pos as i32;
let positions: Vec<i32> = vec![pos_i32; 4];
let last_input = &generated_tokens[generated_tokens.len() - 1..];
let tok = if is_greedy && !want_logprobs {
qwen.model
.forward_gpu_greedy(last_input, &positions, kv_cache, slot_id)
.with_context(|| {
format!(
"Qwen35Model::forward_gpu_greedy (slot-aware decode step {step}; \
ADR-040 §6.1.50 iter-G)"
)
})?
} else {
let logits = qwen
.model
.forward_gpu_last_logits(last_input, &positions, kv_cache, slot_id)
.with_context(|| {
format!("Qwen35Model::forward_gpu_last_logits (slot-aware decode step {step})")
})?;
anyhow::ensure!(
logits.len() == qwen.vocab_size,
"qwen35 slot-aware decode logits len {} != vocab_size {}",
logits.len(),
qwen.vocab_size
);
let mut logits = logits;
if let Some(ref mut lps) = logprobs_vec {
let (tok, lp) =
sample_logits_qwen35_with_logprob(&mut logits, params, &generated_tokens);
lps.push(lp);
tok
} else {
sample_logits_qwen35(&mut logits, params, &generated_tokens)
}
};
if qwen.eos_token_ids.contains(&tok) {
finish_reason = "stop";
break;
}
generated_tokens.push(tok);
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
decoded_text.push_str(&frag);
if qwen35_hit_stop_string(&decoded_text, stops) {
qwen35_strip_trailing_stop(&mut decoded_text, stops);
finish_reason = "stop";
break;
}
step += 1;
}
let decode_duration = decode_start.elapsed();
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-1: reset_for_slot at exit")?;
let (content_text, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (decoded_text, None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &tok in &generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
Ok(GenerationResult {
text: content_text,
reasoning_text,
prompt_tokens: prompt_len,
completion_tokens: generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: if prompt_cache_hit { prompt_len } else { 0 },
logprobs: logprobs_vec,
})
}
pub(crate) struct Qwen35TickOutcome {
pub fragment: String,
pub is_reasoning: bool,
pub finished: bool,
}
pub(crate) enum Qwen35PrefillAdvance {
Pending {
state: Qwen35PrefillState,
advanced_tokens: usize,
checkpoint: Option<Qwen35StablePromptCheckpoint>,
},
Ready {
state: Qwen35DecodeState,
prefill_logits: Vec<f32>,
advanced_tokens: usize,
checkpoint: Option<Qwen35StablePromptCheckpoint>,
},
}
pub(crate) struct Qwen35StablePromptCheckpoint {
pub prompt_tokens: Vec<u32>,
pub kv: HybridKvSlotAnchor,
pub prefill_logits: Vec<f32>,
pub vision_fingerprint: Option<[u8; 32]>,
pub spec: Option<Qwen35SpecPrefixBoundary>,
}
#[derive(Clone)]
pub(crate) struct Qwen35SpecPrefixBoundary {
pub token_count: usize,
pub pending_target_hidden: MlxBuffer,
}
#[derive(Debug)]
pub(crate) struct Qwen35VisionPrefillData {
soft_tokens: Vec<SoftTokenData>,
deepstack: Option<DeepstackData>,
positions_flat: Option<Vec<i32>>,
}
impl Qwen35VisionPrefillData {
pub(crate) fn new(
soft_tokens: Vec<SoftTokenData>,
deepstack: Option<DeepstackData>,
positions_flat: Option<Vec<i32>>,
) -> Self {
Self {
soft_tokens,
deepstack,
positions_flat,
}
}
pub(crate) fn text_anchor_reuse_limit(&self, prompt_len: usize) -> Option<usize> {
let soft_ranges: Vec<_> = self
.soft_tokens
.iter()
.map(|soft| soft.range.clone())
.collect();
super::engine::qwen35_text_anchor_reuse_limit(
prompt_len,
&soft_ranges,
self.deepstack
.as_ref()
.map(|data| data.image_token_positions.as_slice()),
self.positions_flat.as_deref(),
)
}
pub(crate) fn validate(&self, prompt_len: usize, hidden_size: usize) -> Result<()> {
anyhow::ensure!(
!self.soft_tokens.is_empty()
|| self.deepstack.is_some()
|| self.positions_flat.is_some(),
"Qwen35VisionPrefillData must carry at least one multimodal extension"
);
if let Some(positions) = self.positions_flat.as_ref() {
anyhow::ensure!(
positions.len() == 4 * prompt_len,
"Qwen35 vision positions len {} != 4 * prompt_len {}",
positions.len(),
4 * prompt_len
);
}
let row_bytes = hidden_size
.checked_mul(std::mem::size_of::<f32>())
.context("Qwen35 vision hidden row byte overflow")?;
let mut previous_end = 0usize;
for (index, soft) in self.soft_tokens.iter().enumerate() {
anyhow::ensure!(
soft.range.start < soft.range.end && soft.range.end <= prompt_len,
"Qwen35 vision soft token [{index}] range {:?} outside prompt_len={prompt_len}",
soft.range
);
anyhow::ensure!(
index == 0 || soft.range.start >= previous_end,
"Qwen35 vision soft token ranges overlap or are unsorted at index {index}"
);
let needed = soft
.range
.len()
.checked_mul(row_bytes)
.context("Qwen35 vision soft-token byte-size overflow")?;
anyhow::ensure!(
soft.embeddings.byte_len() >= needed,
"Qwen35 vision soft token [{index}] has {} bytes, needs at least {needed}",
soft.embeddings.byte_len()
);
previous_end = soft.range.end;
}
if let Some(deepstack) = self.deepstack.as_ref() {
anyhow::ensure!(
deepstack
.image_token_positions
.windows(2)
.all(|pair| pair[0] < pair[1]),
"Qwen35 vision deepstack positions must be strictly increasing"
);
anyhow::ensure!(
deepstack
.image_token_positions
.last()
.is_none_or(|position| (*position as usize) < prompt_len),
"Qwen35 vision deepstack position outside prompt_len={prompt_len}"
);
let needed = deepstack
.image_token_positions
.len()
.checked_mul(row_bytes)
.context("Qwen35 vision deepstack byte-size overflow")?;
for (index, chunk) in deepstack.chunks.iter().enumerate() {
anyhow::ensure!(
chunk.byte_len() >= needed,
"Qwen35 vision deepstack chunk [{index}] has {} bytes, needs at least {needed}",
chunk.byte_len()
);
}
}
Ok(())
}
fn chunk(&self, start: usize, end: usize, hidden_size: usize) -> Result<Qwen35VisionChunk> {
anyhow::ensure!(start < end, "Qwen35 vision chunk must be non-empty");
let row_bytes = hidden_size
.checked_mul(std::mem::size_of::<f32>())
.context("Qwen35 vision chunk row byte overflow")?;
let mut soft_tokens = Vec::new();
for soft in &self.soft_tokens {
let intersection_start = soft.range.start.max(start);
let intersection_end = soft.range.end.min(end);
if intersection_start >= intersection_end {
continue;
}
let source_row = intersection_start - soft.range.start;
let rows = intersection_end - intersection_start;
let byte_offset = source_row
.checked_mul(row_bytes)
.context("Qwen35 vision soft-token view offset overflow")?;
let elements = rows
.checked_mul(hidden_size)
.context("Qwen35 vision soft-token view length overflow")?;
soft_tokens.push((
(intersection_start - start)..(intersection_end - start),
soft.embeddings.slice_view(byte_offset as u64, elements),
));
}
let deepstack = self.deepstack.as_ref().and_then(|deepstack| {
let selected: Vec<(usize, u32)> = deepstack
.image_token_positions
.iter()
.copied()
.enumerate()
.filter(|(_, position)| {
let position = *position as usize;
position >= start && position < end
})
.collect();
let (first, _) = selected.first().copied()?;
let rows = selected.len();
debug_assert!(selected
.iter()
.enumerate()
.all(|(offset, (index, _))| *index == first + offset));
let byte_offset = first.checked_mul(row_bytes)?;
let elements = rows.checked_mul(hidden_size)?;
let positions = selected
.into_iter()
.map(|(_, position)| position - start as u32)
.collect();
let chunks = deepstack
.chunks
.iter()
.map(|chunk| chunk.slice_view(byte_offset as u64, elements))
.collect();
Some((positions, chunks))
});
let positions_flat = self.positions_flat.as_ref().map(|positions| {
let full_len = positions.len() / 4;
let mut chunk_positions = Vec::with_capacity(4 * (end - start));
for axis in 0..4 {
chunk_positions
.extend_from_slice(&positions[axis * full_len + start..axis * full_len + end]);
}
chunk_positions
});
Ok(Qwen35VisionChunk {
soft_tokens,
deepstack,
positions_flat,
})
}
fn decode_position_base(&self, prompt_len: usize) -> usize {
self.positions_flat
.as_ref()
.map(|positions| {
positions[..prompt_len]
.iter()
.copied()
.max()
.unwrap_or(0)
.saturating_add(1)
.max(0) as usize
})
.unwrap_or(prompt_len)
}
}
struct Qwen35VisionChunk {
soft_tokens: Vec<(std::ops::Range<usize>, mlx_native::MlxBuffer)>,
deepstack: Option<(Vec<u32>, Vec<mlx_native::MlxBuffer>)>,
positions_flat: Option<Vec<i32>>,
}
pub(crate) struct Qwen35PrefillState {
slot_id: SlotId,
prompt_tokens: Vec<u32>,
params: SamplingParams,
cached_tokens: usize,
next_token_index: usize,
cached_prefill_logits: Option<Vec<f32>>,
stable_prompt_prefix_tokens: Option<usize>,
vision: Option<Qwen35VisionPrefillData>,
mtp_pending_hidden: Option<MlxBuffer>,
speculation_unavailable: bool,
prefill_started: Instant,
}
fn qwen35_next_prefill_end(
cursor: usize,
prompt_len: usize,
max_chunk_tokens: usize,
stable_prompt_prefix_tokens: Option<usize>,
) -> usize {
let mut end = cursor.saturating_add(max_chunk_tokens).min(prompt_len);
if let Some(boundary) = stable_prompt_prefix_tokens {
if cursor < boundary && boundary < end {
end = boundary;
}
}
end
}
impl Qwen35PrefillState {
pub(crate) fn prompt_tokens(&self) -> &[u32] {
&self.prompt_tokens
}
pub(crate) fn vision_fingerprint(&self) -> Option<[u8; 32]> {
self.params.vision_fingerprint
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn begin(
prompt_tokens: Vec<u32>,
params: SamplingParams,
registration: Option<&ModelRegistration>,
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
cached_tokens: usize,
cached_prefill_logits: Option<Vec<f32>>,
cached_spec: Option<Qwen35SpecPrefixBoundary>,
vision: Option<Qwen35VisionPrefillData>,
hidden_size: usize,
) -> Result<Self> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"Qwen35PrefillState::begin: empty prompt_tokens"
);
anyhow::ensure!(
slot_id.0 < kv_cache.n_seqs,
"Qwen35PrefillState::begin: SlotOutOfRange slot={} max_slots={}",
slot_id.0,
kv_cache.n_seqs,
);
let prompt_len = prompt_tokens.len();
anyhow::ensure!(
cached_tokens <= prompt_len,
"Qwen35PrefillState::begin: cached_tokens={} exceeds prompt_len={}",
cached_tokens,
prompt_len,
);
anyhow::ensure!(
(cached_tokens == prompt_len) == cached_prefill_logits.is_some(),
"Qwen35PrefillState::begin: a full-prompt cache hit requires prompt-boundary logits, and partial/cold prefill must not supply them"
);
if let Some(spec) = cached_spec.as_ref() {
anyhow::ensure!(
spec.token_count == cached_tokens,
"Qwen cached speculative boundary token_count={} != cached_tokens={cached_tokens}",
spec.token_count,
);
anyhow::ensure!(
vision.is_none(),
"Qwen speculative prefix reuse is text-only"
);
kv_cache
.validate_speculative_cursors_for_slot(slot_id, cached_tokens)
.context("Qwen cached speculative boundary cursor equality")?;
}
let max_tokens = params.max_tokens.max(1);
let need_seq = prompt_len
.checked_add(max_tokens)
.and_then(|tokens| tokens.checked_add(64))
.context("Qwen35PrefillState::begin: prompt + completion capacity overflow")?;
anyhow::ensure!(
need_seq <= kv_cache.max_seq_len as usize,
"Qwen35PrefillState::begin: per-request need_seq={} exceeds persistent cache max_seq_len={} (slot={} prompt_len={} max_tokens={})",
need_seq,
kv_cache.max_seq_len,
slot_id.0,
prompt_len,
max_tokens,
);
let _ = grammar_runtime_for_request(¶ms, registration)?;
if let Some(vision) = vision.as_ref() {
vision.validate(prompt_len, hidden_size)?;
anyhow::ensure!(
params.vision_fingerprint.is_some(),
"Qwen35 SlotAware multimodal prefill requires an exact vision fingerprint"
);
}
if cached_tokens == 0 {
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 full-context slots: Qwen35 cold reset_for_slot at entry")?;
}
kv_cache
.validate_sequence_len_for_slot(slot_id, cached_tokens)
.context("Qwen35PrefillState::begin: validate cache/ledger boundary")?;
let stable_prompt_prefix_tokens = params
.stable_prompt_prefix_tokens
.filter(|&boundary| boundary > cached_tokens && boundary < prompt_len);
Ok(Self {
slot_id,
prompt_tokens,
params,
cached_tokens,
next_token_index: cached_tokens,
cached_prefill_logits,
stable_prompt_prefix_tokens,
vision,
mtp_pending_hidden: cached_spec.map(|spec| spec.pending_target_hidden),
speculation_unavailable: false,
prefill_started: Instant::now(),
})
}
pub(super) fn advance(
mut self,
qwen: &mut Qwen35LoadedModel,
registration: Option<&ModelRegistration>,
kv_cache: &mut HybridKvCache,
max_chunk_tokens: usize,
supervisor: &EngineSupervisor,
) -> Result<Qwen35PrefillAdvance> {
let (prefill_logits, mtp_hidden, advanced_tokens, checkpoint) = if let Some(logits) =
self.cached_prefill_logits.take()
{
anyhow::ensure!(
logits.len() == qwen.vocab_size,
"qwen35 cached prompt-boundary logits len {} != vocab_size {}",
logits.len(),
qwen.vocab_size
);
(logits, None, 0, None)
} else {
anyhow::ensure!(
max_chunk_tokens > 0,
"Qwen35PrefillState::advance requires a non-zero chunk"
);
kv_cache
.validate_sequence_len_for_slot(self.slot_id, self.next_token_index)
.context("validate Qwen35 slot cursors before bounded prefill")?;
let end = qwen35_next_prefill_end(
self.next_token_index,
self.prompt_tokens.len(),
max_chunk_tokens,
self.stable_prompt_prefix_tokens,
);
anyhow::ensure!(
end > self.next_token_index,
"Qwen35 bounded prefill has no suffix without cached logits"
);
let chunk_start = self.next_token_index;
let chunk = &self.prompt_tokens[self.next_token_index..end];
let vision_chunk = self
.vision
.as_ref()
.map(|vision| vision.chunk(self.next_token_index, end, qwen.hidden_size))
.transpose()?;
let positions = vision_chunk
.as_ref()
.and_then(|vision| vision.positions_flat.clone())
.unwrap_or_else(|| prefill_positions_from(self.next_token_index, chunk.len()));
let chunk_started = Instant::now();
let lease =
supervisor.arm("Qwen35 bounded prefill", QWEN35_WORKER_TRANSACTION_TIMEOUT)?;
let mtp_prefill = (self.cached_tokens == 0 || self.mtp_pending_hidden.is_some())
&& !self.speculation_unavailable
&& self.vision.is_none()
&& qwen.speculation.policy()
== super::qwen35_speculation::QwenSpeculationPolicy::Auto
&& is_qwen_server_speculation_exact_eligible(&self.params)
&& qwen.model.mtp.is_some()
&& kv_cache.mtp_slot.is_some();
let (forward, mtp_hidden) = if let Some(vision) = vision_chunk.as_ref() {
let soft_tokens: Vec<_> = vision
.soft_tokens
.iter()
.map(
|(range, embeddings)| crate::serve::forward_prefill::SoftTokenInjection {
range: range.clone(),
embeddings,
},
)
.collect();
let deepstack = vision.deepstack.as_ref().map(|(positions, chunks)| {
crate::serve::forward_prefill::DeepstackInjection {
image_token_positions: positions.clone(),
chunks: chunks.iter().collect(),
}
});
(
qwen.model
.forward_gpu_last_logits_with_soft_tokens_and_deepstack(
chunk,
&positions,
&soft_tokens,
deepstack.as_ref(),
kv_cache,
self.slot_id,
),
None,
)
} else if mtp_prefill {
match qwen.model.forward_gpu_last_logits_with_hidden(
chunk,
&positions,
kv_cache,
self.slot_id,
) {
Err(error) => (
Err(error.context("Qwen SlotAware MTP target prefill")),
None,
),
Ok((logits, target_nextn)) => {
let catchup = (|| -> Result<MlxBuffer> {
let shared_embed_rows = qwen.model.embed_tokens_gpu(chunk)?;
let mtp =
qwen.model.mtp.as_ref().context(
"Qwen SlotAware MTP prompt catch-up weights missing",
)?;
qwen.model.with_gpu_cache_mut(|device, registry| {
mtp.process_target_batch(
chunk,
self.mtp_pending_hidden.as_ref(),
&target_nextn,
&shared_embed_rows,
kv_cache,
self.slot_id,
&positions,
device,
registry,
&qwen.model.cfg,
)
})?;
kv_cache
.validate_speculative_cursors_for_slot(self.slot_id, end)
.context("Qwen SlotAware prompt target/MTP cursor equality")?;
let hidden =
crate::inference::models::qwen35::spec_decode::last_hidden_row(
&target_nextn,
qwen.model.cfg.hidden_size,
)?;
anyhow::ensure!(
hidden.element_count() == qwen.model.cfg.hidden_size as usize,
"Qwen SlotAware MTP prefill hidden must be one row"
);
Ok(hidden)
})();
match catchup {
Ok(hidden) => {
super::qwen35_speculation::record_outcome(0, 0, 0, 1, 0);
self.mtp_pending_hidden = Some(hidden);
(Ok(logits), None)
}
Err(error) => {
tracing::warn!(
slot = self.slot_id.0,
error = %error,
"Qwen SlotAware MTP prompt catch-up unavailable; replaying bounded ordinary prefill"
);
super::qwen35_speculation::record_fallback(
super::qwen35_speculation::QwenSpeculationDecision::RuntimeUnavailable,
);
rollback_slot_mtp_transaction(
kv_cache,
self.slot_id,
chunk_start as u32,
chunk_start as u32,
)
.context("Qwen SlotAware MTP prompt-catch-up bounded rollback")?;
self.mtp_pending_hidden = None;
self.speculation_unavailable = true;
(
qwen.model.forward_gpu_last_logits(
chunk,
&positions,
kv_cache,
self.slot_id,
),
None,
)
}
}
}
}
} else {
(
qwen.model
.forward_gpu_last_logits(chunk, &positions, kv_cache, self.slot_id),
None,
)
};
lease.finish()?;
let logits = forward
.context("Qwen35Model::forward_gpu_last_logits (slot-aware bounded prefill)")?;
kv_cache
.validate_sequence_len_for_slot(self.slot_id, end)
.context("validate Qwen35 slot cursors after bounded prefill")?;
tracing::info!(
slot = self.slot_id.0,
chunk_start,
chunk_end = end,
chunk_tokens = chunk.len(),
prompt_tokens = self.prompt_tokens.len(),
elapsed_seconds = chunk_started.elapsed().as_secs_f64(),
"Qwen35 bounded prefill chunk complete"
);
let advanced = end - self.next_token_index;
self.next_token_index = end;
let checkpoint = if self.stable_prompt_prefix_tokens == Some(end) {
let spec =
self.mtp_pending_hidden
.as_ref()
.map(|hidden| Qwen35SpecPrefixBoundary {
token_count: end,
pending_target_hidden: hidden.clone(),
});
Some(Qwen35StablePromptCheckpoint {
prompt_tokens: self.prompt_tokens[..end].to_vec(),
kv: kv_cache
.snapshot_slot_anchor(self.slot_id, end)
.context("capture Qwen35 stable prompt boundary")?,
prefill_logits: logits.clone(),
vision_fingerprint: self.params.vision_fingerprint,
spec,
})
} else {
None
};
(logits, mtp_hidden, advanced, checkpoint)
};
if self.next_token_index < self.prompt_tokens.len() {
return Ok(Qwen35PrefillAdvance::Pending {
state: self,
advanced_tokens,
checkpoint,
});
}
let prefill_duration = self.prefill_started.elapsed();
let decode_position_base = self
.vision
.as_ref()
.map(|vision| vision.decode_position_base(self.prompt_tokens.len()))
.unwrap_or(self.prompt_tokens.len());
let mtp_hidden = mtp_hidden.or_else(|| self.mtp_pending_hidden.take());
let state = Qwen35DecodeState::from_prefill_logits(
qwen,
self.prompt_tokens,
self.params,
registration,
self.slot_id,
self.cached_tokens,
&prefill_logits,
prefill_duration,
decode_position_base,
mtp_hidden,
)?;
Ok(Qwen35PrefillAdvance::Ready {
state,
prefill_logits,
advanced_tokens,
checkpoint,
})
}
pub(crate) fn operator_progress(&self) -> (usize, usize, f64) {
let completed = self.next_token_index.saturating_sub(self.cached_tokens);
let work = self.prompt_tokens.len().saturating_sub(self.cached_tokens);
let rate = completed as f64
/ self
.prefill_started
.elapsed()
.as_secs_f64()
.max(f64::EPSILON);
(completed, work, rate)
}
}
pub(crate) struct Qwen35DecodeState {
slot_id: SlotId,
prompt_tokens: Vec<u32>,
prompt_len: usize,
decode_position_base: usize,
max_tokens: usize,
is_greedy: bool,
want_logprobs: bool,
logprobs_vec: Option<Vec<f32>>,
cached_tokens: usize,
params: SamplingParams,
grammar_runtime: Option<super::grammar::GrammarRuntime>,
tool_splitter: Option<ToolCallSplitter>,
next_token: u32,
generated_tokens: Vec<u32>,
sampling_history: Vec<u32>,
decoded_text: String,
stop_strings: Vec<String>,
finish_reason: &'static str,
step: usize,
answer_event_reported: bool,
thinking_budget: Option<Qwen35ThinkingBudgetState>,
prefill_duration: Duration,
decode_start: Instant,
mtp: Option<Qwen35SlotMtpState>,
history_lookup: Option<HistoryLookupIndex>,
pending_speculation_output: VecDeque<u32>,
terminal_after_pending: bool,
mtp_cost: SpeculationCostController,
history_cost: SpeculationCostController,
}
struct Qwen35SlotMtpState {
verifier_hidden: MlxBuffer,
}
const QWEN35_REPETITION_WINDOW: usize = 64;
fn qwen35_prompt_sampling_history(prompt_tokens: &[u32]) -> Vec<u32> {
let start = prompt_tokens.len().saturating_sub(QWEN35_REPETITION_WINDOW);
prompt_tokens[start..].to_vec()
}
fn qwen35_observe_sampling_history(history: &mut Vec<u32>, token: u32) {
if history.len() == QWEN35_REPETITION_WINDOW {
history.remove(0);
}
history.push(token);
}
fn take_pending_speculation_output(queue: &mut VecDeque<u32>) -> Option<u32> {
queue.pop_front()
}
fn equivalent_target_decisions(queue: &VecDeque<u32>, terminal_after_pending: bool) -> usize {
queue.len() + usize::from(terminal_after_pending)
}
fn may_route_history_miss_to_mtp(mtp_available: bool, cost: &SpeculationCostController) -> bool {
mtp_available && cost.may_speculate()
}
fn mtp_cursor_for_slot(kv_cache: &HybridKvCache, slot_id: SlotId) -> Result<u32> {
kv_cache
.mtp_slot
.as_ref()
.and_then(|slot| slot.current_len.get(slot_id.0 as usize).copied())
.context("Qwen SlotAware MTP cursor missing")
}
fn rollback_slot_mtp_transaction(
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
target_cursor: u32,
mtp_cursor: u32,
) -> Result<()> {
let mut failures = Vec::new();
if let Err(error) = kv_cache.truncate_full_attn_to_for_slot(slot_id, target_cursor) {
failures.push(format!("full attention: {error:#}"));
}
if let Err(error) = kv_cache.rewind_la_ping_pong_for_slot(slot_id) {
failures.push(format!("linear attention: {error:#}"));
}
if let Err(error) = kv_cache.truncate_mtp_to_for_slot(slot_id, mtp_cursor) {
failures.push(format!("MTP cursor: {error:#}"));
}
if failures.is_empty() {
return Ok(());
}
if let Err(error) = kv_cache.reset_for_slot(slot_id) {
failures.push(format!("fail-closed slot reset: {error:#}"));
}
anyhow::bail!(
"Qwen SlotAware MTP rollback failed: {}",
failures.join("; ")
)
}
fn rollback_slot_target_transaction(
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
target_cursor: u32,
) -> Result<()> {
let mut failures = Vec::new();
if let Err(error) = kv_cache.truncate_full_attn_to_for_slot(slot_id, target_cursor) {
failures.push(format!("full attention: {error:#}"));
}
if let Err(error) = kv_cache.rewind_la_ping_pong_for_slot(slot_id) {
failures.push(format!("linear attention: {error:#}"));
}
if failures.is_empty() {
return Ok(());
}
if let Err(error) = kv_cache.reset_for_slot(slot_id) {
failures.push(format!("fail-closed slot reset: {error:#}"));
}
anyhow::bail!(
"Qwen SlotAware target rollback failed: {}",
failures.join("; ")
)
}
fn rollback_slot_mtp_draft_error<T>(
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
mtp_cursor: u32,
error: anyhow::Error,
context: &'static str,
) -> Result<T> {
let original = error.context(context);
match kv_cache.truncate_mtp_to_for_slot(slot_id, mtp_cursor) {
Ok(()) => Err(original),
Err(rollback) => {
let reset = kv_cache.reset_for_slot(slot_id);
Err(anyhow::anyhow!(
"{original:#}; MTP draft rollback failure: {rollback:#}; fail-closed reset: {reset:?}"
))
}
}
}
fn rollback_slot_target_error<T>(
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
target_cursor: u32,
error: anyhow::Error,
context: &'static str,
) -> Result<T> {
let original = error.context(context);
match rollback_slot_target_transaction(kv_cache, slot_id, target_cursor) {
Ok(()) => Err(original),
Err(rollback) => Err(anyhow::anyhow!(
"{original:#}; rollback failure: {rollback:#}"
)),
}
}
fn rollback_slot_mtp_error<T>(
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
target_cursor: u32,
mtp_cursor: u32,
error: anyhow::Error,
context: &'static str,
) -> Result<T> {
let original = error.context(context);
match rollback_slot_mtp_transaction(kv_cache, slot_id, target_cursor, mtp_cursor) {
Ok(()) => Err(original),
Err(rollback) => Err(anyhow::anyhow!(
"{original:#}; rollback failure: {rollback:#}"
)),
}
}
#[derive(Debug, Clone)]
struct Qwen35ThinkingBudgetState {
limit: usize,
reasoning_tokens: usize,
forced_tokens: Arc<Vec<u32>>,
close_tokens: Arc<Vec<u32>>,
forced_cursor: Option<usize>,
closed: bool,
}
impl Qwen35ThinkingBudgetState {
fn from_params(params: &SamplingParams) -> Option<Self> {
let limit = params.thinking_token_budget?;
let forced_tokens = params.reasoning_end_tokens.clone()?;
let close_tokens = params.reasoning_close_tokens.clone()?;
(params.reasoning_forced_open
&& limit > 0
&& !forced_tokens.is_empty()
&& !close_tokens.is_empty())
.then_some(Self {
limit,
reasoning_tokens: 0,
forced_tokens,
close_tokens,
forced_cursor: None,
closed: false,
})
}
fn next_forced_token(&mut self) -> Option<(u32, bool)> {
if self.closed {
return None;
}
let started = self.forced_cursor.is_none() && self.reasoning_tokens >= self.limit;
if started {
self.forced_cursor = Some(0);
}
let cursor = self.forced_cursor?;
let token = self.forced_tokens.get(cursor).copied()?;
self.forced_cursor = Some(cursor + 1);
Some((token, started))
}
fn observe_generated(&mut self, generated_tokens: &[u32], tool_opened: bool) {
if self.closed {
return;
}
if tool_opened || generated_tokens.ends_with(self.close_tokens.as_slice()) {
self.closed = true;
return;
}
self.reasoning_tokens = self.reasoning_tokens.saturating_add(1);
}
fn was_forced_closed(&self) -> bool {
self.forced_cursor.is_some() && self.closed
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct Qwen35CanonicalDecision {
token: u32,
terminal: bool,
}
#[derive(Debug, PartialEq, Eq)]
struct Qwen35VerifiedBlockPlan {
output: VecDeque<u32>,
terminal_after_pending: bool,
matched_drafts: usize,
rejected_drafts: usize,
valid_input_tokens: usize,
carry_hidden_row: usize,
}
fn plan_qwen35_verified_block(
drafts: &[u32],
mut canonical_at: impl FnMut(usize) -> Result<Qwen35CanonicalDecision>,
) -> Result<Qwen35VerifiedBlockPlan> {
anyhow::ensure!(
!drafts.is_empty(),
"verified block requires at least one draft"
);
let mut output = VecDeque::with_capacity(drafts.len() + 1);
for (index, &draft) in drafts.iter().enumerate() {
let decision = canonical_at(index)?;
if decision.token != draft {
if !decision.terminal {
output.push_back(decision.token);
}
return Ok(Qwen35VerifiedBlockPlan {
output,
terminal_after_pending: decision.terminal,
matched_drafts: index,
rejected_drafts: 1,
valid_input_tokens: index + 1,
carry_hidden_row: index,
});
}
if decision.terminal {
return Ok(Qwen35VerifiedBlockPlan {
output,
terminal_after_pending: true,
matched_drafts: index + 1,
rejected_drafts: 0,
valid_input_tokens: index + 1,
carry_hidden_row: index,
});
}
output.push_back(draft);
}
let bonus = canonical_at(drafts.len())?;
if !bonus.terminal {
output.push_back(bonus.token);
}
Ok(Qwen35VerifiedBlockPlan {
output,
terminal_after_pending: bonus.terminal,
matched_drafts: drafts.len(),
rejected_drafts: 0,
valid_input_tokens: drafts.len() + 1,
carry_hidden_row: drafts.len(),
})
}
#[derive(Clone)]
struct Qwen35SpecSemanticState {
generated_tokens: Vec<u32>,
sampling_history: Vec<u32>,
grammar_runtime: Option<super::grammar::GrammarRuntime>,
tool_splitter: Option<ToolCallSplitter>,
thinking_budget: Option<Qwen35ThinkingBudgetState>,
}
impl Qwen35SpecSemanticState {
fn from_decode(state: &Qwen35DecodeState) -> Self {
Self {
generated_tokens: state.generated_tokens.clone(),
sampling_history: state.sampling_history.clone(),
grammar_runtime: state.grammar_runtime.clone(),
tool_splitter: state.tool_splitter.clone(),
thinking_budget: state.thinking_budget.clone(),
}
}
fn select_and_observe(
&mut self,
qwen: &Qwen35LoadedModel,
params: &SamplingParams,
logits: &mut [f32],
) -> Result<Qwen35CanonicalDecision> {
let forced = self
.thinking_budget
.as_mut()
.and_then(Qwen35ThinkingBudgetState::next_forced_token)
.map(|(token, _)| token);
let token = if let Some(forced) = forced {
forced
} else {
sample_logits_qwen35_constrained(
logits,
params,
&self.sampling_history,
self.grammar_runtime.as_ref(),
false,
)?
.0
};
advance_qwen35_grammar(&mut self.grammar_runtime, params, token);
let terminal = qwen.eos_token_ids.contains(&token)
|| qwen35_grammar_terminal_token(self.grammar_runtime.as_ref(), params, token);
if terminal {
return Ok(Qwen35CanonicalDecision { token, terminal });
}
self.generated_tokens.push(token);
qwen35_observe_sampling_history(&mut self.sampling_history, token);
let fragment = qwen.tokenizer.decode(&[token], false).unwrap_or_default();
let tool_opened = self.tool_splitter.as_mut().is_some_and(|splitter| {
splitter
.feed(&fragment)
.iter()
.any(|event| matches!(event, ToolCallEvent::ToolCallOpen))
});
if tool_opened {
if let Some(runtime) = self.grammar_runtime.as_mut() {
runtime.trigger();
}
}
if let Some(budget) = self.thinking_budget.as_mut() {
budget.observe_generated(&self.generated_tokens, tool_opened);
}
Ok(Qwen35CanonicalDecision { token, terminal })
}
}
impl Qwen35DecodeState {
pub(crate) fn prefill_seed(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
cached_tokens: usize,
cached_prefill_logits: Option<&[f32]>,
) -> Result<(Self, Vec<f32>)> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"Qwen35DecodeState::prefill_seed: empty prompt_tokens"
);
anyhow::ensure!(
slot_id.0 < kv_cache.n_seqs,
"Qwen35DecodeState::prefill_seed: SlotOutOfRange slot={} max_slots={} \
(ADR-040 Phase F M1)",
slot_id.0,
kv_cache.n_seqs,
);
let prompt_len = prompt_tokens.len();
anyhow::ensure!(
cached_tokens <= prompt_len,
"Qwen35DecodeState::prefill_seed: cached_tokens={} exceeds prompt_len={}",
cached_tokens,
prompt_len,
);
anyhow::ensure!(
(cached_tokens == prompt_len) == cached_prefill_logits.is_some(),
"Qwen35DecodeState::prefill_seed: a full-prompt cache hit requires prompt-boundary logits, and partial/cold prefill must not supply them"
);
let max_tokens = params.max_tokens.max(1);
let need_seq = prompt_len + max_tokens + 64;
if need_seq > kv_cache.max_seq_len as usize {
return Err(anyhow::anyhow!(
"Qwen35DecodeState::prefill_seed: per-request need_seq={} exceeds \
persistent cache max_seq_len={} (slot={} prompt_len={} max_tokens={}). \
ADR-040 Phase F M1; reduce max_tokens or use a shorter prompt.",
need_seq,
kv_cache.max_seq_len,
slot_id.0,
prompt_len,
max_tokens
));
}
let _ = grammar_runtime_for_request(params, registration)?;
if cached_tokens == 0 {
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 full-context slots: Qwen35 cold reset_for_slot at entry")?;
} else {
let cursor = kv_cache
.sequence_len_for_slot(slot_id)
.context("ADR-040 full-context slots: read Qwen35 retained cursor")?
as usize;
anyhow::ensure!(
cursor == cached_tokens,
"Qwen35DecodeState::prefill_seed: retained token ledger/cache cursor mismatch for slot {} (ledger={}, cursor={}); refusing unsafe resume",
slot_id.0,
cached_tokens,
cursor,
);
}
let prefill_start = Instant::now();
let prefill_logits = if let Some(logits) = cached_prefill_logits {
anyhow::ensure!(
logits.len() == qwen.vocab_size,
"qwen35 cached prompt-boundary logits len {} != vocab_size {}",
logits.len(),
qwen.vocab_size
);
logits.to_vec()
} else {
let suffix = &prompt_tokens[cached_tokens..];
anyhow::ensure!(
!suffix.is_empty(),
"Qwen35DecodeState::prefill_seed: empty suffix without cached prompt-boundary logits"
);
let positions = prefill_positions_from(cached_tokens, suffix.len());
qwen.model
.forward_gpu_last_logits(suffix, &positions, kv_cache, slot_id)
.context("Qwen35Model::forward_gpu_last_logits (slot-aware prefill)")?
};
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 slot-aware prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
let prefill_duration = prefill_start.elapsed();
let state = Self::from_prefill_logits(
qwen,
prompt_tokens.to_vec(),
params.clone(),
registration,
slot_id,
cached_tokens,
&prefill_logits,
prefill_duration,
prompt_len,
None,
)?;
Ok((state, prefill_logits))
}
#[allow(clippy::too_many_arguments)]
fn from_prefill_logits(
qwen: &Qwen35LoadedModel,
prompt_tokens: Vec<u32>,
params: SamplingParams,
registration: Option<&ModelRegistration>,
slot_id: SlotId,
cached_tokens: usize,
prefill_logits: &[f32],
prefill_duration: Duration,
decode_position_base: usize,
mtp_hidden: Option<MlxBuffer>,
) -> Result<Self> {
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 slot-aware prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible(¶ms);
let want_logprobs = params.logprobs;
let mut logprobs_vec = want_logprobs.then(|| Vec::with_capacity(max_tokens));
let mut grammar_runtime = grammar_runtime_for_request(¶ms, registration)?;
let mut tool_splitter = registration.and_then(ToolCallSplitter::from_registration);
let mut sampling_history = qwen35_prompt_sampling_history(&prompt_tokens);
let next_token = if is_greedy && !want_logprobs {
greedy_argmax_last_token(prefill_logits, qwen.vocab_size as u32)
} else {
let mut logits = prefill_logits.to_vec();
let (token, logprob) = sample_logits_qwen35_constrained(
&mut logits,
¶ms,
&sampling_history,
grammar_runtime.as_ref(),
want_logprobs,
)?;
if let (Some(values), Some(logprob)) = (logprobs_vec.as_mut(), logprob) {
values.push(logprob);
}
token
};
advance_qwen35_grammar(&mut grammar_runtime, ¶ms, next_token);
let mut generated_tokens = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
qwen35_observe_sampling_history(&mut sampling_history, next_token);
let mut decoded_text = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut tool_opened = false;
if let Some(splitter) = tool_splitter.as_mut() {
let marker_events = splitter.feed(&decoded_text);
tool_opened = marker_events
.iter()
.any(|event| matches!(event, ToolCallEvent::ToolCallOpen));
if tool_opened {
if let Some(runtime) = grammar_runtime.as_mut() {
runtime.trigger();
}
}
}
let stop_strings = params.stop_strings.clone();
let mut finish_reason = "length";
if qwen.eos_token_ids.contains(&next_token)
|| qwen35_grammar_terminal_token(grammar_runtime.as_ref(), ¶ms, next_token)
{
generated_tokens.pop();
decoded_text.clear();
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, &stop_strings) {
qwen35_strip_trailing_stop(&mut decoded_text, &stop_strings);
finish_reason = "stop";
}
let mut thinking_budget = Qwen35ThinkingBudgetState::from_params(¶ms);
if finish_reason == "length" {
if let Some(budget) = thinking_budget.as_mut() {
budget.observe_generated(&generated_tokens, tool_opened);
}
}
let history_lookup = if finish_reason == "length"
&& qwen.speculation.policy() == super::qwen35_speculation::QwenSpeculationPolicy::Auto
&& is_qwen_server_speculation_exact_eligible(¶ms)
&& cached_tokens != prompt_len
{
let mut lookup = HistoryLookupIndex::new(HistoryLookupConfig {
min_match: 6,
max_match: 12,
max_draft_tokens: 3,
max_model_len: qwen.model.cfg.max_position_embeddings as usize,
});
lookup.reset(&prompt_tokens);
lookup.extend_verified(&generated_tokens);
Some(lookup)
} else {
None
};
Ok(Self {
slot_id,
prompt_tokens,
prompt_len,
decode_position_base,
max_tokens,
is_greedy,
want_logprobs,
logprobs_vec,
cached_tokens,
params,
grammar_runtime,
tool_splitter,
next_token,
generated_tokens,
sampling_history,
decoded_text,
stop_strings,
finish_reason,
step: 1,
answer_event_reported: false,
thinking_budget,
prefill_duration,
decode_start: Instant::now(),
mtp: mtp_hidden.map(|verifier_hidden| Qwen35SlotMtpState { verifier_hidden }),
history_lookup,
pending_speculation_output: VecDeque::new(),
terminal_after_pending: false,
mtp_cost: SpeculationCostController::new(),
history_cost: SpeculationCostController::new(),
})
}
pub(crate) fn finished_at_seed(&self) -> bool {
self.finish_reason != "length"
}
pub(crate) fn seed_fragment(&self) -> String {
self.decoded_text.clone()
}
pub(crate) fn retained_prefix(&self, valid_tokens: usize) -> Vec<u32> {
self.prompt_tokens
.iter()
.chain(self.generated_tokens.iter())
.copied()
.take(valid_tokens)
.collect()
}
pub(crate) fn spec_prefix_candidate(
&self,
valid_tokens: usize,
) -> Option<Qwen35SpecPrefixBoundary> {
if valid_tokens == 0
|| valid_tokens > self.prompt_tokens.len() + self.generated_tokens.len()
|| !self.pending_speculation_output.is_empty()
{
return None;
}
self.mtp.as_ref().map(|mtp| Qwen35SpecPrefixBoundary {
token_count: valid_tokens,
pending_target_hidden: mtp.verifier_hidden.clone(),
})
}
pub(crate) fn prompt_cache_identity(&self) -> (&[u32], &SamplingParams) {
(&self.prompt_tokens, &self.params)
}
pub(crate) fn mark_first_answer_event(&mut self) -> bool {
if self.answer_event_reported {
false
} else {
self.answer_event_reported = true;
true
}
}
pub(crate) fn operator_progress(&self) -> (usize, usize, f64) {
let generated = self.generated_tokens.len();
let rate = generated as f64 / self.decode_start.elapsed().as_secs_f64().max(f64::EPSILON);
(generated, self.max_tokens, rate)
}
pub(crate) fn operator_thinking_progress(&self) -> (Option<usize>, Option<usize>, bool, bool) {
self.thinking_budget.as_ref().map_or(
(None, None, false, self.answer_event_reported),
|budget| {
(
Some(budget.reasoning_tokens.min(budget.limit)),
Some(budget.limit),
budget.was_forced_closed(),
self.answer_event_reported,
)
},
)
}
pub(crate) fn operator_prefill_progress(&self) -> (usize, usize, f64) {
let work = self.prompt_len.saturating_sub(self.cached_tokens);
let rate = work as f64 / self.prefill_duration.as_secs_f64().max(f64::EPSILON);
(work, work, rate)
}
pub(super) fn decode_tick(
&mut self,
qwen: &mut Qwen35LoadedModel,
kv_cache: &mut HybridKvCache,
supervisor: &EngineSupervisor,
) -> Result<Qwen35TickOutcome> {
if self.step >= self.max_tokens {
return Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
});
}
if let Some(token) = take_pending_speculation_output(&mut self.pending_speculation_output) {
return self.commit_speculation_output(qwen, token);
}
if self.history_lookup.is_some() {
return self.decode_tick_history_lookup(qwen, kv_cache, supervisor);
}
if self.mtp.is_some() {
return self.decode_tick_mtp_k3(qwen, kv_cache, supervisor);
}
let ordinary_target_started = Instant::now();
let pos = self.decode_position_base + self.step - 1;
let pos_i32 = pos as i32;
let positions: Vec<i32> = vec![pos_i32; 4];
let last_input = &self.generated_tokens[self.generated_tokens.len() - 1..];
let forced_token = self
.thinking_budget
.as_mut()
.and_then(Qwen35ThinkingBudgetState::next_forced_token);
if forced_token.is_some_and(|(_, started)| started) {
tracing::warn!(
slot = self.slot_id.0,
budget = self.params.thinking_token_budget,
generated_tokens = self.generated_tokens.len(),
"Qwen35 thinking token budget reached; forcing reasoning close and continuing answer"
);
}
let tok = if self.is_greedy && !self.want_logprobs {
let lease =
supervisor.arm("Qwen35 decode greedy", QWEN35_WORKER_TRANSACTION_TIMEOUT)?;
let forward =
qwen.model
.forward_gpu_greedy(last_input, &positions, kv_cache, self.slot_id);
lease.finish()?;
let predicted = forward.with_context(|| {
format!(
"Qwen35Model::forward_gpu_greedy (slot-aware decode step {}; \
ADR-040 Phase F M1)",
self.step
)
})?;
forced_token.map_or(predicted, |(token, _)| token)
} else {
let lease =
supervisor.arm("Qwen35 decode logits", QWEN35_WORKER_TRANSACTION_TIMEOUT)?;
let forward =
qwen.model
.forward_gpu_last_logits(last_input, &positions, kv_cache, self.slot_id);
lease.finish()?;
let logits = forward.with_context(|| {
format!(
"Qwen35Model::forward_gpu_last_logits (slot-aware decode step {})",
self.step
)
})?;
anyhow::ensure!(
logits.len() == qwen.vocab_size,
"qwen35 slot-aware decode logits len {} != vocab_size {}",
logits.len(),
qwen.vocab_size
);
let (token, logprob) = if let Some((token, _)) = forced_token {
(token, self.want_logprobs.then_some(0.0))
} else {
let mut logits = logits;
sample_logits_qwen35_constrained(
&mut logits,
&self.params,
&self.sampling_history,
self.grammar_runtime.as_ref(),
self.want_logprobs,
)?
};
if let (Some(values), Some(logprob)) = (self.logprobs_vec.as_mut(), logprob) {
values.push(logprob);
}
token
};
let ordinary_target_elapsed = ordinary_target_started.elapsed();
self.mtp_cost
.observe_ordinary_target(ordinary_target_elapsed);
self.history_cost
.observe_ordinary_target(ordinary_target_elapsed);
advance_qwen35_grammar(&mut self.grammar_runtime, &self.params, tok);
if qwen.eos_token_ids.contains(&tok) {
self.finish_reason = "stop";
return Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
});
}
if qwen35_grammar_terminal_token(self.grammar_runtime.as_ref(), &self.params, tok) {
self.finish_reason = "stop";
return Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
});
}
self.generated_tokens.push(tok);
qwen35_observe_sampling_history(&mut self.sampling_history, tok);
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
self.decoded_text.push_str(&frag);
let mut tool_opened = false;
if let Some(splitter) = self.tool_splitter.as_mut() {
let marker_events = splitter.feed(&frag);
tool_opened = marker_events
.iter()
.any(|event| matches!(event, ToolCallEvent::ToolCallOpen));
if tool_opened {
if let Some(runtime) = self.grammar_runtime.as_mut() {
runtime.trigger();
}
}
}
if let Some(budget) = self.thinking_budget.as_mut() {
budget.observe_generated(&self.generated_tokens, tool_opened);
}
if qwen35_hit_stop_string(&self.decoded_text, &self.stop_strings) {
qwen35_strip_trailing_stop(&mut self.decoded_text, &self.stop_strings);
self.finish_reason = "stop";
return Ok(Qwen35TickOutcome {
fragment: frag,
is_reasoning: false,
finished: true,
});
}
self.step += 1;
let finished = self.step >= self.max_tokens;
Ok(Qwen35TickOutcome {
fragment: frag,
is_reasoning: false,
finished,
})
}
fn decode_tick_history_lookup(
&mut self,
qwen: &mut Qwen35LoadedModel,
kv_cache: &mut HybridKvCache,
supervisor: &EngineSupervisor,
) -> Result<Qwen35TickOutcome> {
if let Some(token) = take_pending_speculation_output(&mut self.pending_speculation_output) {
return self.commit_speculation_output(qwen, token);
}
if !self.history_cost.may_speculate() {
let lookup = self
.history_lookup
.take()
.expect("history lookup checked above");
let outcome = self.decode_tick(qwen, kv_cache, supervisor);
self.history_lookup = Some(lookup);
return outcome;
}
let expected_len = self.prompt_tokens.len() + self.generated_tokens.len();
let lookup = self
.history_lookup
.as_mut()
.expect("history lookup checked above");
if lookup.verified_len() < expected_len {
let generated_cursor = lookup
.verified_len()
.checked_sub(self.prompt_tokens.len())
.context("Qwen history lookup precedes prompt boundary")?;
lookup.extend_verified(&self.generated_tokens[generated_cursor..]);
}
anyhow::ensure!(
lookup.verified_len() == expected_len,
"Qwen history lookup ledger cursor mismatch"
);
let remaining = self.max_tokens.saturating_sub(self.generated_tokens.len());
let mut drafts = lookup.propose();
drafts.truncate(remaining.saturating_sub(1).min(3));
if drafts.is_empty() {
super::qwen35_speculation::record_history_lookup_no_match();
let mut lookup = self
.history_lookup
.take()
.expect("history lookup checked above");
if self.mtp.is_some() && !may_route_history_miss_to_mtp(true, &self.mtp_cost) {
self.mtp = None;
}
let committed_before = self.generated_tokens.len();
let outcome = self.decode_tick(qwen, kv_cache, supervisor);
if self.generated_tokens.len() > committed_before {
lookup.extend_verified(&self.generated_tokens[committed_before..]);
}
self.history_lookup = Some(lookup);
return outcome;
}
let round_started = Instant::now();
let next_token = *self
.generated_tokens
.last()
.context("Qwen history lookup needs a seed")?;
let next_pos = (self.decode_position_base + self.generated_tokens.len() - 1) as i32;
let prior_target_len = kv_cache.sequence_len_for_slot(self.slot_id)?;
let prior_mtp_len = self
.mtp
.as_ref()
.map(|_| mtp_cursor_for_slot(kv_cache, self.slot_id))
.transpose()?;
if let Some(mtp_len) = prior_mtp_len {
anyhow::ensure!(
mtp_len == prior_target_len,
"Qwen history verifier requires equal target/MTP cursors (target={prior_target_len}, mtp={mtp_len})"
);
}
let verify_rows = drafts.len() + 1;
if let Err(error) = qwen.model.with_gpu_cache_mut(|device, _registry| {
kv_cache.ensure_la_capture(&qwen.model.cfg, device, verify_rows as u32)
}) {
kv_cache.clear_la_capture();
return Err(error.context("Qwen history recurrent capture allocation"));
}
let mut verify_input = Vec::with_capacity(verify_rows);
verify_input.push(next_token);
verify_input.extend_from_slice(&drafts);
let verify_positions = crate::inference::models::qwen35::spec_decode::positions_for_range(
next_pos,
verify_rows,
);
let lease = match supervisor.arm(
"Qwen35 SlotAware history block verify",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
) {
Ok(lease) => lease,
Err(error) => {
kv_cache.clear_la_capture();
return Err(error.context("Qwen history verifier admission"));
}
};
let verified = qwen.model.forward_gpu_with_nextn_hidden_buffer(
&verify_input,
&verify_positions,
kv_cache,
self.slot_id,
);
let supervision = lease.finish();
let (mut verify_logits, verify_hidden) = match (supervision, verified) {
(Ok(()), Ok(value)) => value,
(supervision, forward) => {
let error = supervision
.err()
.or_else(|| forward.err())
.expect("failed history verification has an error");
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history verifier fail-closed reset")?;
return Err(error.context("Qwen history block verify"));
}
};
let vocab = qwen.vocab_size;
if verify_logits.element_count() != verify_rows * vocab
|| verify_hidden.element_count() != verify_rows * qwen.model.cfg.hidden_size as usize
{
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history verify shape reset")?;
anyhow::bail!("Qwen history verify output shape mismatch");
}
if let (Some(mtp_state), Some(prior_mtp_len)) = (self.mtp.as_ref(), prior_mtp_len) {
let shared_embed_rows = match qwen.model.embed_tokens_gpu(&verify_input) {
Ok(rows) => rows,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history MTP embedding reset")?;
return Err(error.context("Qwen history MTP embeddings"));
}
};
let mtp = match qwen.model.mtp.as_ref() {
Some(mtp) => mtp,
None => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history missing-MTP reset")?;
anyhow::bail!("Qwen history MTP state exists without weights");
}
};
let lease = match supervisor.arm(
"Qwen35 SlotAware history MTP catch-up",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
) {
Ok(lease) => lease,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history MTP catch-up admission reset")?;
return Err(error.context("Qwen history MTP catch-up admission"));
}
};
let caught_up = qwen.model.with_gpu_cache_mut(|device, registry| {
mtp.process_target_batch(
&verify_input,
Some(&mtp_state.verifier_hidden),
&verify_hidden,
&shared_embed_rows,
kv_cache,
self.slot_id,
&verify_positions,
device,
registry,
&qwen.model.cfg,
)
});
if let Err(error) = lease.finish().and(caught_up) {
kv_cache.clear_la_capture();
kv_cache.reset_for_slot(self.slot_id).with_context(|| {
format!("Qwen history MTP catch-up failed ({error:#}); fail-closed reset")
})?;
return Err(error.context("Qwen history MTP catch-up"));
}
debug_assert_eq!(prior_mtp_len, prior_target_len);
}
let verify_logits = match verify_logits.as_mut_slice::<f32>() {
Ok(logits) => logits,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history logits view failure reset")?;
return Err(anyhow::anyhow!("{error}").context("Qwen history logits view"));
}
};
let mut semantic = Qwen35SpecSemanticState::from_decode(self);
let plan = match plan_qwen35_verified_block(&drafts, |row| {
let start = row * vocab;
semantic.select_and_observe(
qwen,
&self.params,
&mut verify_logits[start..start + vocab],
)
}) {
Ok(plan) => plan,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history semantic failure reset")?;
return Err(error.context("Qwen history canonical accept walk"));
}
};
let committed_cursor = prior_target_len + plan.valid_input_tokens as u32;
let state_result = (|| -> Result<()> {
kv_cache.truncate_full_attn_to_for_slot(self.slot_id, committed_cursor)?;
if prior_mtp_len.is_some() {
kv_cache.truncate_mtp_to_for_slot(self.slot_id, committed_cursor)?;
}
if plan.valid_input_tokens < verify_rows {
kv_cache.rollback_la_to(self.slot_id, plan.carry_hidden_row as u32)?;
}
kv_cache.clear_la_capture();
if prior_mtp_len.is_some() {
kv_cache.validate_speculative_cursors_for_slot(
self.slot_id,
committed_cursor as usize,
)?;
} else {
kv_cache.validate_sequence_len_for_slot(self.slot_id, committed_cursor as usize)?;
}
Ok(())
})();
if let Err(error) = state_result {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history commit failure reset")?;
return Err(error.context("Qwen history commit verified prefix"));
}
if let Some(mtp_state) = self.mtp.as_mut() {
mtp_state.verifier_hidden =
match crate::inference::models::qwen35::spec_decode::nth_hidden_row(
&verify_hidden,
qwen.model.cfg.hidden_size,
plan.carry_hidden_row as u64,
) {
Ok(hidden) => hidden,
Err(error) => {
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen history carry-hidden failure reset")?;
return Err(error.context("Qwen history carry hidden"));
}
};
}
self.pending_speculation_output = plan.output;
self.terminal_after_pending = plan.terminal_after_pending;
super::qwen35_speculation::record_proposer_outcome(
super::qwen35_speculation::QwenSpeculationProposer::HistoryLookup,
drafts.len(),
plan.matched_drafts,
plan.rejected_drafts,
1,
0,
);
let equivalent_decisions = equivalent_target_decisions(
&self.pending_speculation_output,
self.terminal_after_pending,
);
let round_elapsed = round_started.elapsed();
if let Some(equivalent_ordinary) = self
.history_cost
.equivalent_ordinary_elapsed(equivalent_decisions)
{
super::qwen35_speculation::record_proposer_timing(
super::qwen35_speculation::QwenSpeculationProposer::HistoryLookup,
round_elapsed,
equivalent_ordinary,
);
}
let remains_profitable = self
.history_cost
.observe_speculative_round(equivalent_decisions, round_elapsed);
if !remains_profitable {
super::qwen35_speculation::record_cost_disabled(
super::qwen35_speculation::QwenSpeculationProposer::HistoryLookup,
);
self.history_lookup = None;
}
if let Some(token) = take_pending_speculation_output(&mut self.pending_speculation_output) {
self.commit_speculation_output(qwen, token)
} else {
self.finish_reason = "stop";
Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
})
}
}
fn decode_tick_mtp_k3(
&mut self,
qwen: &mut Qwen35LoadedModel,
kv_cache: &mut HybridKvCache,
supervisor: &EngineSupervisor,
) -> Result<Qwen35TickOutcome> {
const DRAFT_DEPTH: usize = 3;
const VERIFY_ROWS: usize = DRAFT_DEPTH + 1;
if let Some(token) = take_pending_speculation_output(&mut self.pending_speculation_output) {
return self.commit_speculation_output(qwen, token);
}
let remaining = self.max_tokens.saturating_sub(self.generated_tokens.len());
if remaining < VERIFY_ROWS || !self.mtp_cost.may_speculate() {
return self.decode_tick_mtp_warmup(qwen, kv_cache, supervisor);
}
let round_started = Instant::now();
let phase_profile = std::env::var("HF2Q_MTP_PHASE_PROFILE").as_deref() == Ok("1");
let next_token = *self
.generated_tokens
.last()
.context("Qwen SlotAware MTP K3 needs a seeded output token")?;
let next_pos = (self.decode_position_base + self.generated_tokens.len() - 1) as i32;
let prior_target_len = kv_cache
.sequence_len_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 read target cursor")?;
let prior_mtp_len = mtp_cursor_for_slot(kv_cache, self.slot_id)?;
anyhow::ensure!(
prior_target_len == prior_mtp_len,
"Qwen SlotAware MTP K3 requires equal entry cursors (target={prior_target_len}, mtp={prior_mtp_len})"
);
let mtp = qwen
.model
.mtp
.as_ref()
.context("Qwen SlotAware MTP K3 weights missing")?;
let initial_hidden = &self
.mtp
.as_ref()
.context("Qwen SlotAware MTP K3 state missing")?
.verifier_hidden;
if let Err(error) = qwen.model.with_gpu_cache_mut(|device, _registry| {
kv_cache
.ensure_la_capture(&qwen.model.cfg, device, VERIFY_ROWS as u32)
.context("Qwen SlotAware MTP K3 allocate recurrent capture")
}) {
kv_cache.clear_la_capture();
return Err(error);
}
let draft_lease = match supervisor.arm(
"Qwen35 SlotAware MTP K3 draft",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
) {
Ok(lease) => lease,
Err(error) => {
kv_cache.clear_la_capture();
return Err(error.context("Qwen SlotAware MTP K3 draft admission"));
}
};
let drafted = qwen.model.with_gpu_cache_mut(|device, registry| {
let mut drafts = Vec::with_capacity(DRAFT_DEPTH);
let mut chain_hidden: Option<MlxBuffer> = None;
let mut chain_token = next_token;
for depth in 0..DRAFT_DEPTH {
let shared_embed =
qwen.model
.embed_tokens_gpu_in_context(&[chain_token], device, registry)?;
let previous = chain_hidden.as_ref().unwrap_or(initial_hidden);
let (token, next_hidden) = mtp.forward_draft_greedy_for_token(
previous,
chain_token,
&shared_embed,
kv_cache,
self.slot_id,
&[next_pos + depth as i32; 4],
device,
registry,
&qwen.model.cfg,
)?;
drafts.push(token);
chain_token = token;
chain_hidden = Some(next_hidden);
}
Ok::<_, anyhow::Error>(drafts)
});
let draft_supervision = draft_lease.finish();
if let Err(error) = draft_supervision.and(
drafted
.as_ref()
.map(|_| ())
.map_err(|e| anyhow::anyhow!("{e:#}")),
) {
kv_cache.clear_la_capture();
return rollback_slot_mtp_draft_error(
kv_cache,
self.slot_id,
prior_mtp_len,
error,
"Qwen SlotAware MTP K3 draft",
);
}
let drafts = drafted.expect("draft result checked above");
if let Err(error) = kv_cache.truncate_mtp_to_for_slot(self.slot_id, prior_mtp_len) {
kv_cache.clear_la_capture();
return rollback_slot_mtp_draft_error(
kv_cache,
self.slot_id,
prior_mtp_len,
error,
"Qwen SlotAware MTP K3 discard speculative draft cache",
);
}
let draft_elapsed = round_started.elapsed();
let mut verify_input = Vec::with_capacity(VERIFY_ROWS);
verify_input.push(next_token);
verify_input.extend_from_slice(&drafts);
let verify_positions = crate::inference::models::qwen35::spec_decode::positions_for_range(
next_pos,
VERIFY_ROWS,
);
let verify_lease = match supervisor.arm(
"Qwen35 SlotAware MTP K3 verify",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
) {
Ok(lease) => lease,
Err(error) => {
kv_cache.clear_la_capture();
return Err(error.context("Qwen SlotAware MTP K3 verify admission"));
}
};
let verify_started = Instant::now();
let verified = qwen.model.forward_gpu_with_nextn_hidden_buffer(
&verify_input,
&verify_positions,
kv_cache,
self.slot_id,
);
let verify_supervision = verify_lease.finish();
let (mut verify_logits, verify_hidden) = match (verify_supervision, verified) {
(Ok(()), Ok(value)) => value,
(supervision, forward) => {
let error = supervision
.err()
.or_else(|| forward.err())
.expect("failed verification has an error");
kv_cache.clear_la_capture();
kv_cache.reset_for_slot(self.slot_id).with_context(|| {
format!("Qwen SlotAware MTP K3 verify failed ({error:#}); fail-closed reset")
})?;
return Err(error.context("Qwen SlotAware MTP K3 verify"));
}
};
let verify_elapsed = verify_started.elapsed();
let vocab = qwen.vocab_size;
if verify_logits.element_count() != VERIFY_ROWS * vocab
|| verify_hidden.element_count() != VERIFY_ROWS * qwen.model.cfg.hidden_size as usize
{
let error = anyhow::anyhow!(
"Qwen SlotAware MTP K3 verify shape mismatch: logits={} expected={}, hidden={} expected={}",
verify_logits.element_count(),
VERIFY_ROWS * vocab,
verify_hidden.element_count(),
VERIFY_ROWS * qwen.model.cfg.hidden_size as usize,
);
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 shape failure reset")?;
return Err(error);
}
let embed_started = Instant::now();
let shared_embed_rows = match qwen.model.embed_tokens_gpu(&verify_input) {
Ok(rows) => rows,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 embedding failure reset")?;
return Err(error.context("Qwen SlotAware MTP K3 verifier embeddings"));
}
};
let embed_elapsed = embed_started.elapsed();
let catchup_lease = match supervisor.arm(
"Qwen35 SlotAware MTP K3 target catch-up",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
) {
Ok(lease) => lease,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 catch-up admission reset")?;
return Err(error.context("Qwen SlotAware MTP K3 catch-up admission"));
}
};
let catchup_started = Instant::now();
let caught_up = qwen.model.with_gpu_cache_mut(|device, registry| {
mtp.process_target_batch(
&verify_input,
Some(initial_hidden),
&verify_hidden,
&shared_embed_rows,
kv_cache,
self.slot_id,
&verify_positions,
device,
registry,
&qwen.model.cfg,
)
});
let catchup_supervision = catchup_lease.finish();
if let Err(error) = catchup_supervision.and(caught_up) {
kv_cache.clear_la_capture();
kv_cache.reset_for_slot(self.slot_id).with_context(|| {
format!("Qwen SlotAware MTP K3 catch-up failed ({error:#}); fail-closed reset")
})?;
return Err(error.context("Qwen SlotAware MTP K3 target catch-up"));
}
let catchup_elapsed = catchup_started.elapsed();
let verify_logits = match verify_logits.as_mut_slice::<f32>() {
Ok(logits) => logits,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 logits view failure reset")?;
return Err(anyhow::anyhow!("{error}").context("Qwen SlotAware MTP K3 logits view"));
}
};
let semantic_started = Instant::now();
let mut semantic = Qwen35SpecSemanticState::from_decode(self);
let plan = match plan_qwen35_verified_block(&drafts, |row| {
let start = row * vocab;
semantic.select_and_observe(
qwen,
&self.params,
&mut verify_logits[start..start + vocab],
)
}) {
Ok(plan) => plan,
Err(error) => {
kv_cache.clear_la_capture();
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 semantic failure reset")?;
return Err(error.context("Qwen SlotAware MTP K3 canonical accept walk"));
}
};
let semantic_elapsed = semantic_started.elapsed();
let committed_cursor = prior_target_len + plan.valid_input_tokens as u32;
let state_started = Instant::now();
let mut rollback_elapsed = Duration::ZERO;
let state_result = (|| -> Result<()> {
kv_cache
.truncate_full_attn_to_for_slot(self.slot_id, committed_cursor)
.context("Qwen SlotAware MTP K3 target truncate")?;
kv_cache
.truncate_mtp_to_for_slot(self.slot_id, committed_cursor)
.context("Qwen SlotAware MTP K3 MTP truncate")?;
if plan.valid_input_tokens < VERIFY_ROWS {
let rollback_started = Instant::now();
kv_cache
.rollback_la_to(self.slot_id, plan.carry_hidden_row as u32)
.context("Qwen SlotAware MTP K3 recurrent rollback")?;
rollback_elapsed = rollback_started.elapsed();
}
kv_cache.clear_la_capture();
kv_cache
.validate_speculative_cursors_for_slot(self.slot_id, committed_cursor as usize)
.context("Qwen SlotAware MTP K3 committed cursor equality")?;
Ok(())
})();
let state_elapsed = state_started.elapsed();
if let Err(error) = state_result {
kv_cache.clear_la_capture();
kv_cache.reset_for_slot(self.slot_id).with_context(|| {
format!("Qwen SlotAware MTP K3 commit failed ({error:#}); fail-closed reset")
})?;
return Err(error);
}
let carry_hidden = match crate::inference::models::qwen35::spec_decode::nth_hidden_row(
&verify_hidden,
qwen.model.cfg.hidden_size,
plan.carry_hidden_row as u64,
) {
Ok(hidden) => hidden,
Err(error) => {
kv_cache
.reset_for_slot(self.slot_id)
.context("Qwen SlotAware MTP K3 carry-hidden failure reset")?;
return Err(error.context("Qwen SlotAware MTP K3 carry hidden"));
}
};
self.mtp
.as_mut()
.expect("MTP state checked above")
.verifier_hidden = carry_hidden;
self.pending_speculation_output = plan.output;
self.terminal_after_pending = plan.terminal_after_pending;
super::qwen35_speculation::record_proposer_outcome(
super::qwen35_speculation::QwenSpeculationProposer::Mtp,
DRAFT_DEPTH,
plan.matched_drafts,
plan.rejected_drafts,
1,
0,
);
let equivalent_decisions = equivalent_target_decisions(
&self.pending_speculation_output,
self.terminal_after_pending,
);
let round_elapsed = round_started.elapsed();
if phase_profile {
eprintln!(
"[MTP_PHASE] draft={:.2}ms verify={:.2}ms embed={:.2}ms catchup={:.2}ms semantic={:.2}ms state={:.2}ms rollback={:.2}ms partial={} remainder={:.2}ms total={:.2}ms",
draft_elapsed.as_secs_f64() * 1000.0,
verify_elapsed.as_secs_f64() * 1000.0,
embed_elapsed.as_secs_f64() * 1000.0,
catchup_elapsed.as_secs_f64() * 1000.0,
semantic_elapsed.as_secs_f64() * 1000.0,
state_elapsed.as_secs_f64() * 1000.0,
rollback_elapsed.as_secs_f64() * 1000.0,
plan.valid_input_tokens < VERIFY_ROWS,
round_elapsed
.saturating_sub(draft_elapsed)
.saturating_sub(verify_elapsed)
.saturating_sub(embed_elapsed)
.saturating_sub(catchup_elapsed)
.saturating_sub(semantic_elapsed)
.saturating_sub(state_elapsed)
.as_secs_f64()
* 1000.0,
round_elapsed.as_secs_f64() * 1000.0,
);
}
if let Some(equivalent_ordinary) = self
.mtp_cost
.equivalent_ordinary_elapsed(equivalent_decisions)
{
super::qwen35_speculation::record_proposer_timing(
super::qwen35_speculation::QwenSpeculationProposer::Mtp,
round_elapsed,
equivalent_ordinary,
);
}
let remains_profitable = self
.mtp_cost
.observe_speculative_round(equivalent_decisions, round_elapsed);
if !remains_profitable {
super::qwen35_speculation::record_cost_disabled(
super::qwen35_speculation::QwenSpeculationProposer::Mtp,
);
self.mtp = None;
}
if let Some(token) = take_pending_speculation_output(&mut self.pending_speculation_output) {
self.commit_speculation_output(qwen, token)
} else {
self.finish_reason = "stop";
Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
})
}
}
fn decode_tick_mtp_warmup(
&mut self,
qwen: &mut Qwen35LoadedModel,
kv_cache: &mut HybridKvCache,
supervisor: &EngineSupervisor,
) -> Result<Qwen35TickOutcome> {
let next_token = *self
.generated_tokens
.last()
.context("Qwen SlotAware coherent ordinary decode needs a seed")?;
let next_pos = (self.decode_position_base + self.generated_tokens.len() - 1) as i32;
let prior_target = kv_cache.sequence_len_for_slot(self.slot_id)?;
let prior_mtp = mtp_cursor_for_slot(kv_cache, self.slot_id)?;
anyhow::ensure!(
prior_target == prior_mtp,
"Qwen coherent ordinary decode cursor mismatch (target={prior_target}, mtp={prior_mtp})"
);
let pending_hidden = self
.mtp
.as_ref()
.context("Qwen coherent ordinary decode missing MTP state")?
.verifier_hidden
.clone();
let ordinary_decision_started = Instant::now();
let lease = supervisor.arm(
"Qwen35 coherent ordinary target",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
)?;
let forward = qwen.model.forward_gpu_with_nextn_hidden(
&[next_token],
&[next_pos; 4],
kv_cache,
self.slot_id,
);
if let Err(error) = lease.finish() {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary target supervision",
);
}
let (mut logits, nextn_hidden) = match forward {
Ok(value) => value,
Err(error) => {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary target",
);
}
};
let shared_embed = match qwen.model.embed_tokens_gpu(&[next_token]) {
Ok(embed) => embed,
Err(error) => {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary shared embedding",
);
}
};
let mtp = match qwen.model.mtp.as_ref() {
Some(mtp) => mtp,
None => {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
anyhow::anyhow!("Qwen coherent ordinary decode MTP weights missing"),
"Qwen coherent ordinary MTP lookup",
);
}
};
let lease = match supervisor.arm(
"Qwen35 coherent ordinary MTP catch-up",
QWEN35_WORKER_TRANSACTION_TIMEOUT,
) {
Ok(lease) => lease,
Err(error) => {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary MTP catch-up admission",
);
}
};
let caught_up = qwen.model.with_gpu_cache_mut(|device, registry| {
mtp.process_target_batch(
&[next_token],
Some(&pending_hidden),
&nextn_hidden,
&shared_embed,
kv_cache,
self.slot_id,
&[next_pos; 4],
device,
registry,
&qwen.model.cfg,
)
});
if let Err(error) = lease.finish().and(caught_up) {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary MTP catch-up",
);
}
if let Err(error) = kv_cache
.validate_speculative_cursors_for_slot(self.slot_id, (prior_target + 1) as usize)
{
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary committed cursor equality",
);
}
let carry_hidden = match crate::inference::models::qwen35::spec_decode::nth_hidden_row(
&nextn_hidden,
qwen.model.cfg.hidden_size,
0,
) {
Ok(hidden) => hidden,
Err(error) => {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary carry hidden",
);
}
};
self.mtp
.as_mut()
.expect("MTP state checked above")
.verifier_hidden = carry_hidden;
let mut semantic = Qwen35SpecSemanticState::from_decode(self);
let decision = match semantic.select_and_observe(qwen, &self.params, &mut logits) {
Ok(decision) => decision,
Err(error) => {
return rollback_slot_mtp_error(
kv_cache,
self.slot_id,
prior_target,
prior_mtp,
error,
"Qwen coherent ordinary canonical decision",
);
}
};
let ordinary_decision_elapsed = ordinary_decision_started.elapsed();
self.mtp_cost
.observe_ordinary_target(ordinary_decision_elapsed);
self.history_cost
.observe_ordinary_target(ordinary_decision_elapsed);
super::qwen35_speculation::record_outcome(0, 0, 0, 1, 0);
if decision.terminal {
self.finish_reason = "stop";
return Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
});
}
self.pending_speculation_output.push_back(decision.token);
let token = self
.pending_speculation_output
.pop_front()
.expect("coherent ordinary decision queued above");
self.commit_speculation_output(qwen, token)
}
fn commit_speculation_output(
&mut self,
qwen: &Qwen35LoadedModel,
token: u32,
) -> Result<Qwen35TickOutcome> {
if qwen.eos_token_ids.contains(&token) {
self.finish_reason = "stop";
return Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
});
}
let forced_token = self
.thinking_budget
.as_mut()
.and_then(Qwen35ThinkingBudgetState::next_forced_token);
if let Some((forced, started)) = forced_token {
debug_assert_eq!(forced, token, "verified speculative force token drift");
if started {
tracing::warn!(
slot = self.slot_id.0,
budget = self.params.thinking_token_budget,
generated_tokens = self.generated_tokens.len(),
"Qwen35 thinking token budget reached; forcing reasoning close and continuing answer"
);
}
}
advance_qwen35_grammar(&mut self.grammar_runtime, &self.params, token);
if qwen35_grammar_terminal_token(self.grammar_runtime.as_ref(), &self.params, token) {
self.finish_reason = "stop";
return Ok(Qwen35TickOutcome {
fragment: String::new(),
is_reasoning: false,
finished: true,
});
}
self.generated_tokens.push(token);
qwen35_observe_sampling_history(&mut self.sampling_history, token);
self.next_token = token;
let fragment = qwen.tokenizer.decode(&[token], false).unwrap_or_default();
self.decoded_text.push_str(&fragment);
let mut tool_opened = false;
if let Some(splitter) = self.tool_splitter.as_mut() {
let marker_events = splitter.feed(&fragment);
tool_opened = marker_events
.iter()
.any(|event| matches!(event, ToolCallEvent::ToolCallOpen));
if tool_opened {
if let Some(runtime) = self.grammar_runtime.as_mut() {
runtime.trigger();
}
}
}
if let Some(budget) = self.thinking_budget.as_mut() {
budget.observe_generated(&self.generated_tokens, tool_opened);
}
self.step = self.step.saturating_add(1);
let terminal = self.terminal_after_pending && self.pending_speculation_output.is_empty();
if terminal {
self.finish_reason = "stop";
}
Ok(Qwen35TickOutcome {
fragment,
is_reasoning: false,
finished: terminal || self.generated_tokens.len() >= self.max_tokens,
})
}
pub(crate) fn reset_at_exit(&self, kv_cache: &mut HybridKvCache) -> Result<()> {
kv_cache
.reset_for_slot(self.slot_id)
.context("ADR-040 Phase F M1: Qwen35 reset_for_slot at exit")
}
pub(crate) fn finish(
self,
qwen: &Qwen35LoadedModel,
registration: Option<&ModelRegistration>,
) -> GenerationResult {
let (content_text, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&self.decoded_text,
self.params.reasoning_forced_open,
),
_ => (self.decoded_text.clone(), None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp = super::registry::make_reasoning_splitter(
reg,
self.params.reasoning_forced_open,
);
let mut count = 0usize;
for &tok in &self.generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
let decode_duration = self.decode_start.elapsed();
GenerationResult {
text: content_text,
reasoning_text,
prompt_tokens: self.prompt_len,
completion_tokens: self.generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason: self.finish_reason,
prefill_duration: self.prefill_duration,
decode_duration,
cached_tokens: self.cached_tokens,
logprobs: self.logprobs_vec,
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn generate_stream_qwen35_once_extended_slot_aware(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[crate::serve::forward_prefill::SoftTokenInjection<'_>],
deepstack: Option<&crate::serve::forward_prefill::DeepstackInjection<'_>>,
positions_flat: Option<&[i32]>,
params: &SamplingParams,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
registration: Option<&ModelRegistration>,
cancellation_counter: Option<&std::sync::atomic::AtomicU64>,
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
) {
macro_rules! send {
($ev:expr) => {
if events.blocking_send($ev).is_err() {
tracing::info!("SSE stream dropped by client; aborting qwen35 slot-aware decode");
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return;
}
};
}
if prompt_tokens.is_empty() {
send!(GenerationEvent::Error(
"generate_stream_qwen35_once_extended_slot_aware: empty prompt_tokens".into()
));
return;
}
if slot_id.0 >= kv_cache.n_seqs {
send!(GenerationEvent::Error(format!(
"capability_unsupported: ADR-040 iter-C2d-cont-kernel iter-2 — \
SlotOutOfRange slot={} max_slots={} (generate_stream_qwen35_\
once_extended_slot_aware)",
slot_id.0, kv_cache.n_seqs,
)));
return;
}
let has_extension = !soft_tokens.is_empty() || deepstack.is_some() || positions_flat.is_some();
if let Some(p) = positions_flat {
if p.len() != 4 * prompt_tokens.len() {
send!(GenerationEvent::Error(format!(
"qwen35 stream slot-aware (iter-4): positions_flat.len() = {} \
!= 4 * prompt_len = {}",
p.len(),
4 * prompt_tokens.len()
)));
return;
}
}
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let need_seq = prompt_len + max_tokens + 64;
if need_seq > kv_cache.max_seq_len as usize {
send!(GenerationEvent::Error(format!(
"capability_unsupported: ADR-040 iter-C2d-cont-kernel iter-2 — \
per-request need_seq={} exceeds persistent cache \
max_seq_len={} (slot={} prompt_len={} max_tokens={}). \
Persistent cache is sized to cfg.max_position_embeddings; \
reduce max_tokens or use a shorter prompt.",
need_seq, kv_cache.max_seq_len, slot_id.0, prompt_len, max_tokens
)));
return;
}
let is_greedy = is_greedy_eligible(params);
let device = match MlxDevice::new() {
Ok(d) => d,
Err(e) => {
send!(GenerationEvent::Error(format!(
"qwen35 stream slot-aware: MlxDevice::new failed: {e}"
)));
return;
}
};
let _ = &device;
if let Err(e) = kv_cache.reset_for_slot(slot_id) {
send!(GenerationEvent::Error(format!(
"ADR-040 iter-C2d-cont-kernel iter-2: reset_for_slot at entry \
failed: {e:#}"
)));
return;
}
let pre_dispatches = mlx_native::dispatch_count();
let pre_syncs = mlx_native::sync_count();
let prompt_cache_hit =
!has_extension && qwen.prompt_cache.try_match(prompt_tokens, params).is_some();
let prefill_start = Instant::now();
let mut next_token: u32;
if prompt_cache_hit {
let snap = qwen
.prompt_cache
.snapshot()
.expect("try_match Some implies snapshot Some");
if let Err(e) = kv_cache.restore_partial(snap, prompt_len) {
send!(GenerationEvent::Error(format!(
"ADR-040 iter-C2d-cont-kernel iter-2: prompt_cache \
restore_partial failed: {e:#}"
)));
return;
}
next_token = qwen.prompt_cache.first_decoded_token();
tracing::debug!(
"qwen35 stream slot-aware prompt_cache: HIT slot={} prompt_len={} \
prefill skipped",
slot_id.0,
prompt_len
);
} else {
let positions_owned: Vec<i32>;
let positions_slice: &[i32] = match positions_flat {
Some(p) => p,
None => {
positions_owned = prefill_positions_for(prompt_len);
&positions_owned
}
};
let prefill_logits_res: Result<Vec<f32>> = if has_extension {
qwen.model
.forward_gpu_last_logits_with_soft_tokens_and_deepstack(
prompt_tokens,
positions_slice,
soft_tokens,
deepstack,
kv_cache,
slot_id,
)
} else {
qwen.model
.forward_gpu_last_logits(prompt_tokens, positions_slice, kv_cache, slot_id)
};
let prefill_logits = match prefill_logits_res {
Ok(l) => l,
Err(e) => {
send!(GenerationEvent::Error(format!(
"qwen35 stream slot-aware prefill failed: {e:#}"
)));
return;
}
};
if prefill_logits.len() != qwen.vocab_size {
send!(GenerationEvent::Error(format!(
"qwen35 stream slot-aware prefill logits len {} != \
vocab_size {}",
prefill_logits.len(),
qwen.vocab_size,
)));
return;
}
if is_greedy {
next_token = greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32);
} else {
let mut logits = prefill_logits;
next_token = sample_logits_qwen35(&mut logits, params, &[]);
}
}
let prefill_duration = prefill_start.elapsed();
let t_post: i32 = match positions_flat {
Some(p) => {
let mut max_t = 0i32;
for i in 0..prompt_len {
let v = p[i]; if v > max_t {
max_t = v;
}
}
max_t.saturating_add(1)
}
None => prompt_len as i32,
};
let mut reasoning_splitter = registration
.and_then(|r| super::registry::make_reasoning_splitter(r, params.reasoning_forced_open));
let mut tool_splitter = registration.and_then(ToolCallSplitter::from_registration);
let mut tool_call_body: String = String::new();
let mut tool_call_index: usize = 0;
let mut saw_tool_call: bool = false;
fn route_content_qwen35_slot_aware(
tool_splitter: &mut Option<ToolCallSplitter>,
body: &mut String,
tc_index: &mut usize,
saw_tc: &mut bool,
registration: Option<&ModelRegistration>,
wire_kinds: Option<&super::registry::ToolArgumentWireKinds>,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
text: &str,
) -> bool {
if text.is_empty() {
return true;
}
let Some(tcs) = tool_splitter.as_mut() else {
return events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: text.to_string(),
})
.is_ok();
};
for ev in tcs.feed(text) {
match ev {
ToolCallEvent::Content(t) => {
if !t.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: t,
})
.is_err()
{
return false;
}
}
ToolCallEvent::ToolCallOpen => {
body.clear();
}
ToolCallEvent::ToolCallText(t) => {
body.push_str(&t);
}
ToolCallEvent::ToolCallClose => {
let parsed = registration.and_then(|r| {
super::registry::parse_tool_call_body_with_wire_kinds(r, body, wire_kinds)
});
let body_dump = std::mem::take(body);
let sink = super::engine::EventSink::new(events);
if super::engine::emit_streaming_tool_call_close(
parsed,
body_dump,
params_tool_call_policy_for_qwen35_stream(),
tc_index,
saw_tc,
&sink,
)
.is_err()
{
return false;
}
}
}
}
true
}
fn emit_fragment_qwen35_slot_aware(
reasoning_splitter: &mut Option<ReasoningSplitter>,
tool_splitter: &mut Option<ToolCallSplitter>,
body: &mut String,
tc_index: &mut usize,
saw_tc: &mut bool,
registration: Option<&ModelRegistration>,
wire_kinds: Option<&super::registry::ToolArgumentWireKinds>,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
fragment: &str,
) -> bool {
if fragment.is_empty() {
return true;
}
if let Some(rs) = reasoning_splitter.as_mut() {
for (slot, text) in rs.feed(fragment) {
match slot {
SplitSlot::Reasoning => {
if !text.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Reasoning,
text,
})
.is_err()
{
return false;
}
}
SplitSlot::Content => {
if !route_content_qwen35_slot_aware(
tool_splitter,
body,
tc_index,
saw_tc,
registration,
wire_kinds,
events,
&text,
) {
return false;
}
}
}
}
true
} else {
route_content_qwen35_slot_aware(
tool_splitter,
body,
tc_index,
saw_tc,
registration,
wire_kinds,
events,
fragment,
)
}
}
let decode_start = Instant::now();
let mut completion_tokens = 0usize;
let mut generated_tokens = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let mut accumulated_text = String::new();
let mut reasoning_token_count = 0usize;
let mut finish_reason: &'static str = "length";
let first_text = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut is_eos_first = qwen.eos_token_ids.contains(&next_token);
if !is_eos_first && !first_text.is_empty() {
accumulated_text.push_str(&first_text);
if !emit_fragment_qwen35_slot_aware(
&mut reasoning_splitter,
&mut tool_splitter,
&mut tool_call_body,
&mut tool_call_index,
&mut saw_tool_call,
registration,
params.tool_argument_wire_kinds.as_deref(),
events,
&first_text,
) {
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
}
completion_tokens += 1;
if reasoning_splitter
.as_ref()
.map(|s| s.in_reasoning())
.unwrap_or(false)
{
reasoning_token_count += 1;
}
if is_eos_first {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&accumulated_text, ¶ms.stop_strings) {
finish_reason = "stop";
is_eos_first = true;
}
if !is_eos_first {
for step in 1..max_tokens {
let pos = t_post + (step as i32 - 1);
if pos as u32 >= kv_cache.max_seq_len {
break;
}
let decode_positions = vec![pos; 4];
let dec_result: Result<u32, anyhow::Error> = if is_greedy {
qwen.model
.forward_gpu_greedy(&[next_token], &decode_positions, kv_cache, slot_id)
.map_err(|e| {
anyhow::anyhow!(
"qwen35 stream slot-aware forward_gpu_greedy (ADR-040 \
§6.1.50 iter-G) step {step}: {e}"
)
})
} else {
match qwen.model.forward_gpu_last_logits(
&[next_token],
&decode_positions,
kv_cache,
slot_id,
) {
Ok(logits) => {
if logits.len() != qwen.vocab_size {
Err(anyhow::anyhow!(
"qwen35 stream slot-aware decode logits len {} \
!= vocab_size {}",
logits.len(),
qwen.vocab_size,
))
} else {
let mut tmp = logits;
Ok(sample_logits_qwen35(&mut tmp, params, &generated_tokens))
}
}
Err(e) => Err(e),
}
};
next_token = match dec_result {
Ok(t) => t,
Err(e) => {
send!(GenerationEvent::Error(format!(
"qwen35 stream slot-aware decode step {step} failed: {e:#}"
)));
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
generated_tokens.push(next_token);
completion_tokens += 1;
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
accumulated_text.push_str(&fragment);
if !emit_fragment_qwen35_slot_aware(
&mut reasoning_splitter,
&mut tool_splitter,
&mut tool_call_body,
&mut tool_call_index,
&mut saw_tool_call,
registration,
params.tool_argument_wire_kinds.as_deref(),
events,
&fragment,
) {
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
if reasoning_splitter
.as_ref()
.map(|s| s.in_reasoning())
.unwrap_or(false)
{
reasoning_token_count += 1;
}
if qwen35_hit_stop_string(&accumulated_text, ¶ms.stop_strings) {
finish_reason = "stop";
break;
}
}
}
if let Some(rs) = reasoning_splitter.as_mut() {
if let Some((slot, tail)) = rs.finish() {
match slot {
SplitSlot::Reasoning => {
if !tail.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Reasoning,
text: tail,
})
.is_err()
{
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
}
SplitSlot::Content => {
if !route_content_qwen35_slot_aware(
&mut tool_splitter,
&mut tool_call_body,
&mut tool_call_index,
&mut saw_tool_call,
registration,
params.tool_argument_wire_kinds.as_deref(),
events,
&tail,
) {
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
}
}
}
}
if let Some(tcs) = tool_splitter.as_mut() {
if let Some(ev) = tcs.finish() {
match ev {
ToolCallEvent::Content(t) => {
if !t.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: t,
})
.is_err()
{
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
}
ToolCallEvent::ToolCallText(t) => {
let prefix = registration.and_then(|r| r.tool_open).unwrap_or("");
let fallback = format!("{prefix}{t}");
if !fallback.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: fallback,
})
.is_err()
{
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let _ = kv_cache.reset_for_slot(slot_id);
return;
}
}
ToolCallEvent::ToolCallOpen | ToolCallEvent::ToolCallClose => {
}
}
}
}
if saw_tool_call {
finish_reason = "tool_calls";
}
let decode_duration = decode_start.elapsed();
if let Err(e) = kv_cache.reset_for_slot(slot_id) {
send!(GenerationEvent::Error(format!(
"ADR-040 iter-C2d-cont-kernel iter-2: reset_for_slot at exit \
failed: {e:#}"
)));
return;
}
let stats = StreamStats {
prefill_time_secs: Some(prefill_duration.as_secs_f64()),
decode_time_secs: Some(decode_duration.as_secs_f64()),
total_time_secs: Some((prefill_duration + decode_duration).as_secs_f64()),
time_to_first_token_ms: Some(prefill_duration.as_secs_f64() * 1000.0),
prefill_tokens_per_sec: Some(if prefill_duration.as_secs_f64() > 0.0 {
prompt_len as f64 / prefill_duration.as_secs_f64()
} else {
0.0
}),
decode_tokens_per_sec: Some(if decode_duration.as_secs_f64() > 0.0 {
completion_tokens as f64 / decode_duration.as_secs_f64()
} else {
0.0
}),
gpu_sync_count: Some(mlx_native::sync_count().saturating_sub(pre_syncs)),
gpu_dispatch_count: Some(mlx_native::dispatch_count().saturating_sub(pre_dispatches)),
cached_prompt_tokens: Some(if prompt_cache_hit { prompt_len } else { 0 }),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
};
send!(GenerationEvent::Done {
finish_reason,
prompt_tokens: prompt_len,
completion_tokens,
stats,
});
}
pub(super) fn generate_qwen35_once_with_soft_tokens(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[crate::serve::forward_prefill::SoftTokenInjection<'_>],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
supervisor: &EngineSupervisor,
) -> Result<GenerationResult> {
if soft_tokens.is_empty() {
return generate_qwen35_once(qwen, prompt_tokens, params, registration, supervisor);
}
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen35_once_with_soft_tokens: empty prompt_tokens"
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible(params);
let want_logprobs = params.logprobs;
let mut logprobs_vec: Option<Vec<f32>> = if want_logprobs {
Some(Vec::with_capacity(max_tokens))
} else {
None
};
let device = MlxDevice::new()
.map_err(|e| anyhow::anyhow!("MlxDevice::new (qwen35 generate w/ soft tokens): {e}"))?;
let mut kv_cache = alloc_kv_cache_for_request(qwen, &device, prompt_len, max_tokens)?;
qwen.hydrate_lcp_registry_from_disk(&kv_cache, &device);
let prefill_start = Instant::now();
let positions = prefill_positions_for(prompt_len);
let prefill_logits = supervised_gpu_call(supervisor, "qwen35_serial_soft_prefill", || {
qwen.model
.forward_gpu_last_logits_with_soft_tokens(
prompt_tokens,
&positions,
soft_tokens,
&mut kv_cache,
SlotId(0),
)
.context("Qwen35Model::forward_gpu_last_logits_with_soft_tokens (prefill)")
})?;
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 prefill (soft tokens) logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
let mut next_token: u32 = if want_logprobs {
let mut logits = prefill_logits.clone();
let (tok, lp) = sample_logits_qwen35_with_logprob(&mut logits, params, &[]);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32)
} else {
let mut logits = prefill_logits.clone();
sample_logits_qwen35(&mut logits, params, &[])
};
let prefill_duration = prefill_start.elapsed();
let decode_start = Instant::now();
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let first_fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut decoded_text = first_fragment.clone();
let mut finish_reason: &'static str = "length";
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
} else {
for step in 1..max_tokens {
let pos = (prompt_len + step - 1) as i32;
if pos as u32 >= kv_cache.max_seq_len {
tracing::warn!(
pos,
max_seq = kv_cache.max_seq_len,
"qwen35 decode (soft tokens): hit kv-cache bound; stopping with finish=length",
);
break;
}
let decode_positions = vec![pos; 4];
next_token = if want_logprobs {
let logits_full = supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_last_logits(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!(
"forward_gpu_last_logits decode step {step} (soft tokens, logprobs)"
)
})
})?;
let mut logits = logits_full;
let (tok, lp) =
sample_logits_qwen35_with_logprob(&mut logits, params, &generated_tokens);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_greedy(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_greedy decode step {step} (soft tokens)")
})
})?
} else {
let logits_full = supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_last_logits(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_last_logits decode step {step} (soft tokens)")
})
})?;
let mut logits = logits_full;
sample_logits_qwen35(&mut logits, params, &generated_tokens)
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
generated_tokens.push(next_token);
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
decoded_text.push_str(&fragment);
if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
break;
}
}
}
let decode_duration = decode_start.elapsed();
let (content, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (decoded_text, None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &tok in &generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
Ok(GenerationResult {
text: content,
reasoning_text,
prompt_tokens: prompt_len,
completion_tokens: generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: 0,
logprobs: logprobs_vec,
})
}
pub(super) fn generate_qwen35_once_with_soft_tokens_and_deepstack(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[crate::serve::forward_prefill::SoftTokenInjection<'_>],
deepstack: Option<&crate::serve::forward_prefill::DeepstackInjection<'_>>,
positions_flat: Option<&[i32]>,
params: &SamplingParams,
registration: Option<&ModelRegistration>,
supervisor: &EngineSupervisor,
) -> Result<GenerationResult> {
if soft_tokens.is_empty() && deepstack.is_none() && positions_flat.is_none() {
return generate_qwen35_once(qwen, prompt_tokens, params, registration, supervisor);
}
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen35_once_with_soft_tokens_and_deepstack: empty prompt_tokens"
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible(params);
let want_logprobs = params.logprobs;
let mut logprobs_vec: Option<Vec<f32>> = if want_logprobs {
Some(Vec::with_capacity(max_tokens))
} else {
None
};
let device = MlxDevice::new()
.map_err(|e| anyhow::anyhow!("MlxDevice::new (qwen35 vision generate): {e}"))?;
let (mut kv_cache, cache_reused) = take_serial_kv_cache(qwen, &device, prompt_len, max_tokens)?;
qwen.hydrate_lcp_registry_from_disk(&kv_cache, &device);
let prompt_cache_hit = params.vision_fingerprint.is_some()
&& qwen.prompt_cache.try_match(prompt_tokens, params).is_some();
if cache_reused && !prompt_cache_hit {
kv_cache.reset();
}
let prefill_start = Instant::now();
let mut next_token: u32;
if prompt_cache_hit {
let snap = qwen
.prompt_cache
.snapshot()
.expect("vision prompt-cache hit requires snapshot");
kv_cache
.restore_partial(snap, prompt_len)
.context("vision prompt-cache restore_partial")?;
next_token = qwen.prompt_cache.first_decoded_token();
} else {
let positions_owned: Vec<i32>;
let positions: &[i32] = match positions_flat {
Some(p) => {
anyhow::ensure!(
p.len() == 4 * prompt_len,
"generate_qwen35_once_with_soft_tokens_and_deepstack: \
positions_flat.len() = {} != 4 * prompt_len = {}",
p.len(),
4 * prompt_len
);
p
}
None => {
positions_owned = prefill_positions_for(prompt_len);
&positions_owned
}
};
let prefill_logits =
supervised_gpu_call(supervisor, "qwen35_serial_deepstack_prefill", || {
qwen.model
.forward_gpu_last_logits_with_soft_tokens_and_deepstack(
prompt_tokens,
positions,
soft_tokens,
deepstack,
&mut kv_cache,
SlotId(0),
)
.context(
"Qwen35Model::forward_gpu_last_logits_with_soft_tokens_and_deepstack \
(prefill)",
)
})?;
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 vision prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
next_token = if want_logprobs {
let mut logits = prefill_logits.clone();
let (tok, lp) = sample_logits_qwen35_with_logprob(&mut logits, params, &[]);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32)
} else {
let mut logits = prefill_logits.clone();
sample_logits_qwen35(&mut logits, params, &[])
};
if is_greedy && params.vision_fingerprint.is_some() {
match kv_cache.snapshot_prefix(&device, prompt_len) {
Ok(snapshot) => {
qwen.prompt_cache
.update(prompt_tokens.to_vec(), snapshot, next_token, params)
}
Err(error) => tracing::warn!(%error, "Qwen vision prompt-cache snapshot failed"),
}
}
}
let prefill_duration = prefill_start.elapsed();
let t_post: i32 = match positions_flat {
Some(p) => {
let mut max_t = 0i32;
for i in 0..prompt_len {
let v = p[i]; if v > max_t {
max_t = v;
}
}
max_t.saturating_add(1)
}
None => prompt_len as i32,
};
let decode_start = Instant::now();
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let first_fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut decoded_text = first_fragment.clone();
let mut finish_reason: &'static str = "length";
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
} else {
for step in 1..max_tokens {
let pos = t_post + (step as i32 - 1);
if pos as u32 >= kv_cache.max_seq_len {
tracing::warn!(
pos,
max_seq = kv_cache.max_seq_len,
"qwen35 decode (wedge-4d): hit kv-cache bound; stopping with finish=length",
);
break;
}
let decode_positions = vec![pos; 4];
next_token = if want_logprobs {
let logits_full =
supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model.forward_gpu_last_logits(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_last_logits decode step {step} (wedge-4d, logprobs)")
})
})?;
let mut logits = logits_full;
let (tok, lp) =
sample_logits_qwen35_with_logprob(&mut logits, params, &generated_tokens);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_greedy(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_greedy decode step {step} (wedge-4d)")
})
})?
} else {
let logits_full = supervised_gpu_call(supervisor, "qwen35_serial_decode", || {
qwen.model
.forward_gpu_last_logits(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_last_logits decode step {step} (wedge-4d)")
})
})?;
let mut logits = logits_full;
sample_logits_qwen35(&mut logits, params, &generated_tokens)
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
generated_tokens.push(next_token);
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
decoded_text.push_str(&fragment);
if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
break;
}
}
}
let decode_duration = decode_start.elapsed();
qwen.persistent_kv_cache = Some(kv_cache);
let (content, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (decoded_text, None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &tok in &generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
Ok(GenerationResult {
text: content,
reasoning_text,
prompt_tokens: prompt_len,
completion_tokens: generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: if prompt_cache_hit { prompt_len } else { 0 },
logprobs: logprobs_vec,
})
}
pub(super) fn generate_stream_qwen35_once(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
registration: Option<&ModelRegistration>,
cancellation_counter: Option<&std::sync::atomic::AtomicU64>,
supervisor: &EngineSupervisor,
) -> SerialStreamResult {
generate_stream_qwen35_once_extended(
qwen,
prompt_tokens,
&[],
None,
None,
params,
events,
registration,
cancellation_counter,
supervisor,
)
}
#[allow(clippy::too_many_arguments)]
pub(super) fn generate_stream_qwen35_once_extended(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[crate::serve::forward_prefill::SoftTokenInjection<'_>],
deepstack: Option<&crate::serve::forward_prefill::DeepstackInjection<'_>>,
positions_flat: Option<&[i32]>,
params: &SamplingParams,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
registration: Option<&ModelRegistration>,
cancellation_counter: Option<&std::sync::atomic::AtomicU64>,
supervisor: &EngineSupervisor,
) -> SerialStreamResult {
let request_start = Instant::now();
let _disk_request_guard = qwen
.disk_persistor
.as_ref()
.map(|persistor| persistor.begin_request());
macro_rules! send {
($ev:expr) => {
if events.blocking_send($ev).is_err() {
tracing::info!("SSE stream dropped by client; aborting qwen35 decode");
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
};
}
if prompt_tokens.is_empty() {
send!(GenerationEvent::Error(
"generate_stream_qwen35_once: empty prompt_tokens".into()
));
return Ok(SerialStreamEnd::TerminalSent);
}
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible(params);
let mut grammar_runtime = match grammar_runtime_for_request(params, registration) {
Ok(runtime) => runtime,
Err(error) => {
send!(GenerationEvent::Error(format!(
"qwen35 stream grammar initialization failed: {error:#}"
)));
return Ok(SerialStreamEnd::TerminalSent);
}
};
let has_extension = !soft_tokens.is_empty() || deepstack.is_some() || positions_flat.is_some();
if let Some(p) = positions_flat {
if p.len() != 4 * prompt_len {
send!(GenerationEvent::Error(format!(
"qwen35 stream (wedge-4e): positions_flat.len() = {} != 4 * prompt_len = {}",
p.len(),
4 * prompt_len
)));
return Ok(SerialStreamEnd::TerminalSent);
}
}
let cache_alloc_start = Instant::now();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(e) => {
return Err(e).context("Qwen35 SerialFifo streaming device initialization");
}
};
let (mut kv_cache, cache_reused) =
match take_serial_kv_cache(qwen, &device, prompt_len, max_tokens) {
Ok(k) => k,
Err(e) => {
return Err(e).context("Qwen35 SerialFifo streaming KV cache allocation");
}
};
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "stream",
prompt_tokens = prompt_len,
max_tokens,
cache_capacity_tokens = kv_cache.max_seq_len,
cache_reused,
tq_kv = qwen.tq_kv_active,
elapsed_ms = cache_alloc_start.elapsed().as_secs_f64() * 1000.0,
"Qwen35 request cache ready"
);
qwen.hydrate_lcp_registry_from_disk(&kv_cache, &device);
let pre_dispatches = mlx_native::dispatch_count();
let pre_syncs = mlx_native::sync_count();
let prompt_cache_eligible = !has_extension || params.vision_fingerprint.is_some();
let prompt_cache_hit =
prompt_cache_eligible && qwen.prompt_cache.try_match(prompt_tokens, params).is_some();
let mut lcp_resume_start: usize = 0;
if !prompt_cache_hit && !has_extension {
let stride_for_observe = crate::debug::INVESTIGATION_ENV.kv_lcp_deltanet_checkpoint_stride;
let base_key_for_observe = build_lcp_key_for_qwen35(qwen, params);
let detected = lookup_qwen35_resume_checkpoint(
&mut qwen.lcp_registry,
&base_key_for_observe,
prompt_tokens,
stride_for_observe,
)
.map(|(prefix, _chunk_pos)| prefix.k);
if let Some(sink) = qwen.kv_metrics_sink.as_ref() {
sink.record_lcp_probe(detected);
}
let _ = detected;
let lcp_resume_enabled = crate::serve::api::engine::effective_kv_lcp_resume(
crate::debug::INVESTIGATION_ENV.kv_lcp_resume,
true,
);
if lcp_resume_enabled {
let stride = crate::debug::INVESTIGATION_ENV.kv_lcp_deltanet_checkpoint_stride;
let base_key = build_lcp_key_for_qwen35(qwen, params);
eprintln!(
"[hf2q qwen35 stream lcp probe] enabled, registry_len={}, \
prompt_len={}, stride={}, scanning latest-turn + stride checkpoints",
qwen.lcp_registry.len(),
prompt_tokens.len(),
stride,
);
if let Some((prefix, chunk_pos)) = lookup_qwen35_resume_checkpoint(
&mut qwen.lcp_registry,
&base_key,
prompt_tokens,
stride,
) {
let snapshot: &HybridKvCacheSnapshot = &prefix.dense_kvs[0];
let restore_start = Instant::now();
if let Err(e) = kv_cache.restore_partial(snapshot, prefix.k) {
return Err(e).context("Qwen35 SerialFifo streaming LCP checkpoint restore");
}
let restore_ms = restore_start.elapsed().as_micros() as f64 / 1000.0;
lcp_resume_start = prefix.k;
let checkpoint = if chunk_pos == 0 {
"LATEST-TURN"
} else {
"STRIDE-ALIGNED"
};
eprintln!(
"[hf2q qwen35 stream lcp resume] {checkpoint} HIT — restoring \
at k={} (cached_prompt_len={}, chunk_pos={}, restore_ms={:.3})",
prefix.k, prefix.cached_prompt_len, chunk_pos, restore_ms
);
} else {
eprintln!(
"[hf2q qwen35 stream lcp probe] no compatible checkpoint \
(registry_len={})",
qwen.lcp_registry.len()
);
}
}
}
if cache_reused && !prompt_cache_hit && lcp_resume_start == 0 {
kv_cache.reset();
}
let prefill_start = Instant::now();
let mut next_token: u32;
if prompt_cache_hit {
let snap = qwen
.prompt_cache
.snapshot()
.expect("try_match Some implies snapshot Some");
if let Err(e) = kv_cache.restore_partial(snap, prompt_len) {
return Err(e).context("Qwen35 SerialFifo streaming prompt-cache restore");
}
next_token = qwen.prompt_cache.first_decoded_token();
tracing::debug!(
"qwen35 prompt_cache: STREAMING HIT — {} tokens; prefill skipped",
prompt_len
);
} else {
let positions_owned: Vec<i32>;
let positions_slice: &[i32] = match positions_flat {
Some(p) => p,
None => {
positions_owned = prefill_positions_for(prompt_len);
&positions_owned
}
};
let stride = crate::debug::INVESTIGATION_ENV.kv_lcp_deltanet_checkpoint_stride;
let lcp_resume_enabled = crate::debug::INVESTIGATION_ENV.kv_lcp_resume;
let chunked_eligible = !has_extension
&& stride > 0
&& prompt_len > stride
&& (lcp_resume_start == 0 || lcp_resume_start % stride == 0)
&& crate::debug::INVESTIGATION_ENV.kv_lcp_chunked_prefill;
let recovery_tail_tokens = qwen35_recovery_tail_tokens(qwen, prompt_tokens, params);
let recovery_anchor = prompt_len.saturating_sub(recovery_tail_tokens);
let recovery_eligible = !has_extension
&& lcp_resume_enabled
&& recovery_anchor > lcp_resume_start
&& recovery_anchor >= 16;
let recovery_capture_plan = qwen35_recovery_capture_plan(
lcp_resume_start,
recovery_anchor,
prompt_len,
recovery_eligible,
chunked_eligible,
);
let prefill_logits_res = if has_extension {
supervised_gpu_call(supervisor, "qwen35_serial_stream_prefill", || {
qwen.model
.forward_gpu_last_logits_with_soft_tokens_and_deepstack(
prompt_tokens,
positions_slice,
soft_tokens,
deepstack,
&mut kv_cache,
SlotId(0),
)
})
} else if let Some((suffix_len, capture_index)) = recovery_capture_plan {
if let Err(error) =
kv_cache.ensure_la_capture(&qwen.model.cfg, &device, suffix_len as u32)
{
return Err(error)
.context("Qwen35 SerialFifo streaming recovery-capture allocation");
}
let suffix_tokens = &prompt_tokens[lcp_resume_start..];
let mut suffix_positions = vec![0i32; 4 * suffix_len];
for axis in 0..4 {
for token in 0..suffix_len {
suffix_positions[axis * suffix_len + token] = (lcp_resume_start + token) as i32;
}
}
match supervised_gpu_call(supervisor, "qwen35_serial_stream_prefill", || {
qwen.model.forward_gpu_last_logits(
suffix_tokens,
&suffix_positions,
&mut kv_cache,
SlotId(0),
)
}) {
Ok(logits) => {
store_qwen35_latest_turn_checkpoint(
qwen,
&kv_cache,
&device,
params,
prompt_tokens,
recovery_anchor,
"stream captured latest-turn recovery-anchor",
Some(capture_index),
);
kv_cache.clear_la_capture();
eprintln!(
"[hf2q qwen35 stream lcp store] captured latest-turn recovery \
anchor={} suffix_tokens={} capture_index={}",
recovery_anchor, suffix_len, capture_index
);
Ok(logits)
}
Err(error) => Err(error).context("Qwen35 stream captured latest-turn suffix"),
}
} else if recovery_eligible && !chunked_eligible {
let prefix_tokens = &prompt_tokens[lcp_resume_start..recovery_anchor];
let prefix_len = prefix_tokens.len();
let mut prefix_positions = vec![0i32; 4 * prefix_len];
for axis in 0..4 {
for token in 0..prefix_len {
prefix_positions[axis * prefix_len + token] = (lcp_resume_start + token) as i32;
}
}
if let Err(error) =
supervised_gpu_call(supervisor, "qwen35_serial_stream_recovery_prefix", || {
qwen.model.forward_gpu_last_logits(
prefix_tokens,
&prefix_positions,
&mut kv_cache,
SlotId(0),
)
})
{
return Err(error).context("Qwen35 SerialFifo streaming recovery-anchor prefix");
}
store_qwen35_latest_turn_checkpoint(
qwen,
&kv_cache,
&device,
params,
prompt_tokens,
recovery_anchor,
"stream latest-turn recovery-anchor",
None,
);
let tail_tokens = &prompt_tokens[recovery_anchor..];
let tail_len = tail_tokens.len();
let mut tail_positions = vec![0i32; 4 * tail_len];
for axis in 0..4 {
for token in 0..tail_len {
tail_positions[axis * tail_len + token] = (recovery_anchor + token) as i32;
}
}
eprintln!(
"[hf2q qwen35 stream lcp store] latest-turn recovery anchor={} \
tail_tokens={}",
recovery_anchor, tail_len
);
supervised_gpu_call(supervisor, "qwen35_serial_stream_recovery_tail", || {
qwen.model.forward_gpu_last_logits(
tail_tokens,
&tail_positions,
&mut kv_cache,
SlotId(0),
)
})
} else if chunked_eligible {
let first_chunk_idx = lcp_resume_start / stride;
let chunked_prefill_end = if recovery_eligible {
recovery_anchor
} else {
prompt_len
};
let n_chunks = (chunked_prefill_end + stride - 1) / stride;
eprintln!(
"[hf2q qwen35 stream chunked prefill] {} chunks (stride={}, \
prompt_len={}, prefill_end={}, first_chunk_idx={})",
n_chunks, stride, prompt_len, chunked_prefill_end, first_chunk_idx
);
let mut last_logits_res: Result<Vec<f32>> = Err(anyhow::anyhow!("no chunks executed"));
for chunk_idx in first_chunk_idx..n_chunks {
let k_start = chunk_idx * stride;
let k_end = ((chunk_idx + 1) * stride).min(chunked_prefill_end);
let chunk_seq_len = k_end - k_start;
let chunk_tokens = &prompt_tokens[k_start..k_end];
let mut chunk_positions = vec![0i32; 4 * chunk_seq_len];
for axis in 0..4 {
for t in 0..chunk_seq_len {
chunk_positions[axis * chunk_seq_len + t] = (k_start + t) as i32;
}
}
let res =
supervised_gpu_call(supervisor, "qwen35_serial_stream_prefill_chunk", || {
qwen.model.forward_gpu_last_logits(
chunk_tokens,
&chunk_positions,
&mut kv_cache,
SlotId(0),
)
});
let logits = match res {
Ok(l) => l,
Err(e) => {
return Err(e).with_context(|| {
format!(
"Qwen35 SerialFifo streaming prefill chunk {}/{}",
chunk_idx + 1,
n_chunks
)
});
}
};
if chunk_idx == n_chunks - 1 {
last_logits_res = Ok(logits.clone());
}
let stride_aligned = k_end % stride == 0;
let superseded_by_recovery_anchor = stride_checkpoint_superseded_by_recovery_anchor(
recovery_eligible,
k_end,
stride,
recovery_anchor,
);
let mid_store_disabled =
std::env::var("HF2Q_KV_LCP_DISABLE_MID_STORE").as_deref() == Ok("1");
lcp_store_skip_notify(
stride_aligned && !superseded_by_recovery_anchor,
lcp_resume_enabled,
mid_store_disabled,
);
if lcp_resume_enabled
&& stride_aligned
&& !superseded_by_recovery_anchor
&& !mid_store_disabled
{
match kv_cache.snapshot_prefix(&device, k_end) {
Ok(snap) => {
let chunk_key = build_lcp_key_for_qwen35_chunk(qwen, params, k_end);
let linear_capacity = kv_cache
.linear_attn
.first()
.map(|s| s.recurrent.byte_len())
.unwrap_or(0);
if let Err(e) = qwen.store_lcp_with_disk_writeback(
&kv_cache,
chunk_key,
prompt_tokens[..k_end].to_vec(),
snap,
0,
linear_capacity,
) {
lcp_store_error_notify("stream mid-prefill", k_end, &e);
} else {
eprintln!(
"[hf2q qwen35 stream lcp store] mid-prefill \
snapshot at chunk_pos={k_end} \
(registry_len_after={})",
qwen.lcp_registry.len()
);
}
}
Err(e) => {
lcp_snapshot_error_notify("stream mid-prefill", k_end, &e);
}
}
}
}
if recovery_eligible {
store_qwen35_latest_turn_checkpoint(
qwen,
&kv_cache,
&device,
params,
prompt_tokens,
recovery_anchor,
"stream chunked latest-turn recovery-anchor",
None,
);
let tail_tokens = &prompt_tokens[recovery_anchor..];
let tail_len = tail_tokens.len();
let mut tail_positions = vec![0i32; 4 * tail_len];
for axis in 0..4 {
for token in 0..tail_len {
tail_positions[axis * tail_len + token] = (recovery_anchor + token) as i32;
}
}
eprintln!(
"[hf2q qwen35 stream lcp store] latest-turn recovery anchor={} \
tail_tokens={}",
recovery_anchor, tail_len
);
last_logits_res =
supervised_gpu_call(supervisor, "qwen35_serial_stream_recovery_tail", || {
qwen.model.forward_gpu_last_logits(
tail_tokens,
&tail_positions,
&mut kv_cache,
SlotId(0),
)
});
}
last_logits_res
} else if lcp_resume_start > 0 {
let suffix_tokens = &prompt_tokens[lcp_resume_start..];
let suffix_len = suffix_tokens.len();
let mut suffix_positions = vec![0i32; 4 * suffix_len];
for axis in 0..4 {
for t in 0..suffix_len {
suffix_positions[axis * suffix_len + t] = (lcp_resume_start + t) as i32;
}
}
eprintln!(
"[hf2q qwen35 stream lcp resume] suffix prefill {} tokens \
(lcp_resume_start={}, prompt_len={})",
suffix_len, lcp_resume_start, prompt_len
);
supervised_gpu_call(supervisor, "qwen35_serial_stream_prefill", || {
qwen.model.forward_gpu_last_logits(
suffix_tokens,
&suffix_positions,
&mut kv_cache,
SlotId(0),
)
})
} else {
supervised_gpu_call(supervisor, "qwen35_serial_stream_prefill", || {
qwen.model.forward_gpu_last_logits(
prompt_tokens,
positions_slice,
&mut kv_cache,
SlotId(0),
)
})
};
let prefill_logits = match prefill_logits_res {
Ok(l) => l,
Err(e) => {
return Err(e).context("Qwen35 SerialFifo streaming prefill");
}
};
if is_greedy {
next_token = greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32);
} else {
let mut logits = prefill_logits.clone();
next_token = sample_logits_qwen35_constrained(
&mut logits,
params,
&[],
grammar_runtime.as_ref(),
false,
)?
.0;
}
advance_qwen35_grammar(&mut grammar_runtime, params, next_token);
if is_greedy && prompt_cache_eligible {
match kv_cache.snapshot_prefix(&device, prompt_len) {
Ok(snap) => {
qwen.prompt_cache
.update(prompt_tokens.to_vec(), snap, next_token, params)
}
Err(e) => {
eprintln!("[hf2q qwen35 stream lcp store] prompt_cache snapshot failed: {e}")
}
}
}
}
let prefill_duration = prefill_start.elapsed();
let reported_cached_tokens =
qwen35_reported_cached_tokens(prompt_len, prompt_cache_hit, lcp_resume_start);
let prefill_work_tokens = prompt_len.saturating_sub(reported_cached_tokens);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "stream",
prompt_tokens = prompt_len,
cached_tokens = reported_cached_tokens,
work_tokens = prefill_work_tokens,
elapsed_ms = prefill_duration.as_secs_f64() * 1000.0,
tokens_per_second = if prefill_duration.is_zero() {
0.0
} else {
prefill_work_tokens as f64 / prefill_duration.as_secs_f64()
},
"Qwen35 prefill complete"
);
let t_post: i32 = match positions_flat {
Some(p) => {
let mut max_t = 0i32;
for i in 0..prompt_len {
let v = p[i]; if v > max_t {
max_t = v;
}
}
max_t.saturating_add(1)
}
None => prompt_len as i32,
};
let mut reasoning_splitter = registration
.and_then(|r| super::registry::make_reasoning_splitter(r, params.reasoning_forced_open));
let mut tool_splitter = registration.and_then(ToolCallSplitter::from_registration);
let mut tool_call_body: String = String::new();
let mut tool_call_index: usize = 0;
let mut saw_tool_call: bool = false;
fn route_content_qwen35(
tool_splitter: &mut Option<ToolCallSplitter>,
body: &mut String,
tc_index: &mut usize,
saw_tc: &mut bool,
tool_call_policy: super::engine::ToolCallPolicy,
registration: Option<&ModelRegistration>,
wire_kinds: Option<&super::registry::ToolArgumentWireKinds>,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
text: &str,
) -> bool {
if text.is_empty() {
return true;
}
let Some(tcs) = tool_splitter.as_mut() else {
return events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: text.to_string(),
})
.is_ok();
};
for ev in tcs.feed(text) {
match ev {
ToolCallEvent::Content(t) => {
if !t.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: t,
})
.is_err()
{
return false;
}
}
ToolCallEvent::ToolCallOpen => {
body.clear();
}
ToolCallEvent::ToolCallText(t) => {
body.push_str(&t);
}
ToolCallEvent::ToolCallClose => {
let parsed = registration.and_then(|r| {
super::registry::parse_tool_call_body_with_wire_kinds(r, body, wire_kinds)
});
let body_dump = std::mem::take(body);
let sink = super::engine::EventSink::new(events);
if super::engine::emit_streaming_tool_call_close(
parsed,
body_dump,
tool_call_policy,
tc_index,
saw_tc,
&sink,
)
.is_err()
{
return false;
}
}
}
}
true
}
fn emit_fragment_qwen35(
reasoning_splitter: &mut Option<ReasoningSplitter>,
tool_splitter: &mut Option<ToolCallSplitter>,
body: &mut String,
tc_index: &mut usize,
saw_tc: &mut bool,
tool_call_policy: super::engine::ToolCallPolicy,
registration: Option<&ModelRegistration>,
wire_kinds: Option<&super::registry::ToolArgumentWireKinds>,
events: &tokio::sync::mpsc::Sender<GenerationEvent>,
fragment: &str,
) -> bool {
if fragment.is_empty() {
return true;
}
if let Some(rs) = reasoning_splitter.as_mut() {
for (slot, text) in rs.feed(fragment) {
match slot {
SplitSlot::Reasoning => {
if !text.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Reasoning,
text,
})
.is_err()
{
return false;
}
}
SplitSlot::Content => {
if !route_content_qwen35(
tool_splitter,
body,
tc_index,
saw_tc,
tool_call_policy,
registration,
wire_kinds,
events,
&text,
) {
return false;
}
}
}
}
true
} else {
route_content_qwen35(
tool_splitter,
body,
tc_index,
saw_tc,
tool_call_policy,
registration,
wire_kinds,
events,
fragment,
)
}
}
let decode_start = Instant::now();
let mut completion_tokens = 0usize;
let mut accumulated_text = String::new();
let mut reasoning_token_count = 0usize;
let mut finish_reason: &'static str = "length";
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
let first_text = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut is_eos_first = qwen.eos_token_ids.contains(&next_token);
if !is_eos_first && qwen35_grammar_terminal_token(grammar_runtime.as_ref(), params, next_token)
{
is_eos_first = true;
finish_reason = "stop";
}
if !is_eos_first && !first_text.is_empty() {
accumulated_text.push_str(&first_text);
if !emit_fragment_qwen35(
&mut reasoning_splitter,
&mut tool_splitter,
&mut tool_call_body,
&mut tool_call_index,
&mut saw_tool_call,
params.tool_call_policy,
registration,
params.tool_argument_wire_kinds.as_deref(),
events,
&first_text,
) {
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
}
completion_tokens += 1;
if !is_eos_first {
generated_tokens.push(next_token);
}
if reasoning_splitter
.as_ref()
.map(|s| s.in_reasoning())
.unwrap_or(false)
{
reasoning_token_count += 1;
}
if is_eos_first {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&accumulated_text, ¶ms.stop_strings) {
finish_reason = "stop";
is_eos_first = true;
}
if !is_eos_first {
for step in 1..max_tokens {
let pos = t_post + (step as i32 - 1);
if pos as u32 >= kv_cache.max_seq_len {
break;
}
let decode_positions = vec![pos; 4];
let dec_result = if is_greedy {
supervised_gpu_call(supervisor, "qwen35_serial_stream_decode", || {
qwen.model
.forward_gpu_greedy(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
})
} else {
match supervised_gpu_call(supervisor, "qwen35_serial_stream_decode", || {
qwen.model.forward_gpu_last_logits(
&[next_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
}) {
Ok(logits) => {
let mut tmp = logits;
let token = sample_logits_qwen35_constrained(
&mut tmp,
params,
&generated_tokens,
grammar_runtime.as_ref(),
false,
)?
.0;
Ok(token)
}
Err(e) => Err(e),
}
};
next_token = match dec_result {
Ok(t) => t,
Err(e) => {
return Err(e)
.with_context(|| format!("Qwen35 SerialFifo stream decode step {step}"));
}
};
advance_qwen35_grammar(&mut grammar_runtime, params, next_token);
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
if qwen35_grammar_terminal_token(grammar_runtime.as_ref(), params, next_token) {
finish_reason = "stop";
break;
}
completion_tokens += 1;
generated_tokens.push(next_token);
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
accumulated_text.push_str(&fragment);
if !emit_fragment_qwen35(
&mut reasoning_splitter,
&mut tool_splitter,
&mut tool_call_body,
&mut tool_call_index,
&mut saw_tool_call,
params.tool_call_policy,
registration,
params.tool_argument_wire_kinds.as_deref(),
events,
&fragment,
) {
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
if reasoning_splitter
.as_ref()
.map(|s| s.in_reasoning())
.unwrap_or(false)
{
reasoning_token_count += 1;
}
if qwen35_hit_stop_string(&accumulated_text, ¶ms.stop_strings) {
finish_reason = "stop";
break;
}
}
}
if let Some(rs) = reasoning_splitter.as_mut() {
if let Some((slot, tail)) = rs.finish() {
match slot {
SplitSlot::Reasoning => {
if !tail.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Reasoning,
text: tail,
})
.is_err()
{
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
}
SplitSlot::Content => {
if !route_content_qwen35(
&mut tool_splitter,
&mut tool_call_body,
&mut tool_call_index,
&mut saw_tool_call,
params.tool_call_policy,
registration,
params.tool_argument_wire_kinds.as_deref(),
events,
&tail,
) {
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
}
}
}
}
if let Some(tcs) = tool_splitter.as_mut() {
if let Some(ev) = tcs.finish() {
match ev {
ToolCallEvent::Content(t) => {
if !t.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: t,
})
.is_err()
{
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
}
ToolCallEvent::ToolCallText(t) => {
if params.tool_call_policy.enforces_body_grammar() {
send!(GenerationEvent::Error(
"tool_call_truncated_under_constrained".to_string()
));
return Ok(SerialStreamEnd::TerminalSent);
}
let prefix = registration.and_then(|r| r.tool_open).unwrap_or("");
let fallback = format!("{prefix}{t}");
if !fallback.is_empty()
&& events
.blocking_send(GenerationEvent::Delta {
kind: DeltaKind::Content,
text: fallback,
})
.is_err()
{
if let Some(c) = cancellation_counter {
c.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
return Ok(SerialStreamEnd::ClientClosed);
}
}
ToolCallEvent::ToolCallOpen | ToolCallEvent::ToolCallClose => {
}
}
}
}
if matches!(params.tool_call_policy, ToolCallPolicy::Constrained) && !saw_tool_call {
send!(GenerationEvent::Error(
"tool_call_no_call_under_constrained".to_string()
));
return Ok(SerialStreamEnd::TerminalSent);
}
if saw_tool_call {
finish_reason = "tool_calls";
}
let decode_duration = decode_start.elapsed();
let cached_tokens = reported_cached_tokens;
let stats = StreamStats {
prefill_time_secs: Some(prefill_duration.as_secs_f64()),
decode_time_secs: Some(decode_duration.as_secs_f64()),
total_time_secs: Some((prefill_duration + decode_duration).as_secs_f64()),
time_to_first_token_ms: Some(prefill_duration.as_secs_f64() * 1000.0),
prefill_tokens_per_sec: Some(if prefill_duration.as_secs_f64() > 0.0 {
prompt_len.saturating_sub(cached_tokens) as f64 / prefill_duration.as_secs_f64()
} else {
0.0
}),
decode_tokens_per_sec: Some(if decode_duration.as_secs_f64() > 0.0 {
completion_tokens as f64 / decode_duration.as_secs_f64()
} else {
0.0
}),
gpu_sync_count: Some(mlx_native::sync_count().saturating_sub(pre_syncs)),
gpu_dispatch_count: Some(mlx_native::dispatch_count().saturating_sub(pre_dispatches)),
cached_prompt_tokens: Some(cached_tokens),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
};
qwen.persistent_kv_cache = Some(kv_cache);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "stream",
generated_tokens = completion_tokens,
elapsed_ms = decode_duration.as_secs_f64() * 1000.0,
tokens_per_second = if decode_duration.is_zero() {
0.0
} else {
completion_tokens as f64 / decode_duration.as_secs_f64()
},
"Qwen35 decode complete"
);
tracing::info!(
target: "hf2q::serve::api::engine_qwen35::progress",
mode = "stream",
prompt_tokens = prompt_len,
cached_tokens,
completion_tokens,
total_ms = request_start.elapsed().as_secs_f64() * 1000.0,
"Qwen35 request complete"
);
send!(GenerationEvent::Done {
finish_reason,
prompt_tokens: prompt_len,
completion_tokens,
stats,
});
Ok(SerialStreamEnd::TerminalSent)
}
pub(super) fn embed_qwen35(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
supervisor: &EngineSupervisor,
) -> Result<Vec<f32>> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"embed_qwen35: empty prompt_tokens"
);
let device =
MlxDevice::new().map_err(|e| anyhow::anyhow!("MlxDevice::new (qwen35 embed): {e}"))?;
let mut kv_cache = alloc_kv_cache_for_request(qwen, &device, prompt_tokens.len(), 0)?;
let positions = prefill_positions_for(prompt_tokens.len());
supervised_gpu_call(supervisor, "qwen35_serial_embed", || {
qwen.model
.forward_embed_last(prompt_tokens, &positions, &mut kv_cache, SlotId(0))
.context("Qwen35Model::forward_embed_last")
})
}
pub fn embed_qwen35_slot_aware(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
) -> Result<Vec<f32>> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"embed_qwen35_slot_aware: empty prompt_tokens"
);
anyhow::ensure!(
slot_id.0 < kv_cache.n_seqs,
"embed_qwen35_slot_aware: SlotOutOfRange slot={} max_slots={} \
(ADR-040 iter-C2d-cont-kernel iter-3)",
slot_id.0,
kv_cache.n_seqs,
);
let prompt_len = prompt_tokens.len();
let need_seq = prompt_len + 64;
if need_seq > kv_cache.max_seq_len as usize {
return Err(anyhow::anyhow!(
"embed_qwen35_slot_aware: per-request need_seq={} exceeds \
persistent cache max_seq_len={} (slot={} prompt_len={}). \
ADR-040 iter-C2d-cont-kernel iter-3 sizes the persistent \
cache to cfg.max_position_embeddings; use a shorter prompt.",
need_seq,
kv_cache.max_seq_len,
slot_id.0,
prompt_len
));
}
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-3: reset_for_slot at entry")?;
let positions = prefill_positions_for(prompt_len);
let embed_result = qwen
.model
.forward_embed_last(prompt_tokens, &positions, kv_cache, slot_id)
.context(
"Qwen35Model::forward_embed_last (ADR-040 iter-C2d-cont-kernel iter-3 slot-aware)",
);
finish_qwen35_slot_embed(embed_result, || {
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-3: reset_for_slot at exit")
})
}
fn finish_qwen35_slot_embed<T>(
forward: Result<T>,
reset_after_success: impl FnOnce() -> Result<()>,
) -> Result<T> {
let output = forward?;
reset_after_success()?;
Ok(output)
}
pub fn generate_qwen35_once_with_soft_tokens_slot_aware(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[crate::serve::forward_prefill::SoftTokenInjection<'_>],
params: &SamplingParams,
registration: Option<&ModelRegistration>,
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
) -> Result<GenerationResult> {
if soft_tokens.is_empty() {
return generate_qwen35_once_slot_aware(
qwen,
prompt_tokens,
params,
registration,
kv_cache,
slot_id,
);
}
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen35_once_with_soft_tokens_slot_aware: empty prompt_tokens"
);
anyhow::ensure!(
slot_id.0 < kv_cache.n_seqs,
"generate_qwen35_once_with_soft_tokens_slot_aware: SlotOutOfRange slot={} \
max_slots={} (ADR-040 iter-C2d-cont-kernel iter-4)",
slot_id.0,
kv_cache.n_seqs,
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let need_seq = prompt_len + max_tokens + 64;
if need_seq > kv_cache.max_seq_len as usize {
return Err(anyhow::anyhow!(
"generate_qwen35_once_with_soft_tokens_slot_aware: per-request \
need_seq={} exceeds persistent cache max_seq_len={} (slot={} \
prompt_len={} max_tokens={}). ADR-040 iter-C2d-cont-kernel iter-4 \
sizes the persistent cache to cfg.max_position_embeddings; reduce \
max_tokens or use a shorter prompt.",
need_seq,
kv_cache.max_seq_len,
slot_id.0,
prompt_len,
max_tokens
));
}
let is_greedy = is_greedy_eligible(params);
let want_logprobs = params.logprobs;
let mut logprobs_vec: Option<Vec<f32>> = if want_logprobs {
Some(Vec::with_capacity(max_tokens))
} else {
None
};
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-4: reset_for_slot at entry")?;
let prefill_start = Instant::now();
let positions = prefill_positions_for(prompt_len);
let prefill_logits = qwen
.model
.forward_gpu_last_logits_with_soft_tokens(
prompt_tokens,
&positions,
soft_tokens,
kv_cache,
slot_id,
)
.context(
"Qwen35Model::forward_gpu_last_logits_with_soft_tokens \
(slot-aware prefill, ADR-040 iter-C2d-cont-kernel iter-4)",
)?;
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 slot-aware soft-tokens prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
let mut next_token: u32 = if want_logprobs {
let mut logits = prefill_logits.clone();
let (tok, lp) = sample_logits_qwen35_with_logprob(&mut logits, params, &[]);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32)
} else {
let mut logits = prefill_logits.clone();
sample_logits_qwen35(&mut logits, params, &[])
};
let prefill_duration = prefill_start.elapsed();
let decode_start = Instant::now();
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let first_fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut decoded_text = first_fragment.clone();
let mut finish_reason: &'static str = "length";
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
} else {
for step in 1..max_tokens {
let pos = (prompt_len + step - 1) as i32;
if pos as u32 >= kv_cache.max_seq_len {
tracing::warn!(
pos,
max_seq = kv_cache.max_seq_len,
"qwen35 slot-aware decode (soft tokens): hit kv-cache bound; \
stopping with finish=length",
);
break;
}
let decode_positions = vec![pos; 4];
next_token = if want_logprobs {
let logits_full = qwen
.model
.forward_gpu_last_logits(&[next_token], &decode_positions, kv_cache, slot_id)
.with_context(|| {
format!(
"forward_gpu_last_logits slot-aware decode step {step} \
(soft tokens, logprobs)"
)
})?;
let mut logits = logits_full;
let (tok, lp) =
sample_logits_qwen35_with_logprob(&mut logits, params, &generated_tokens);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
qwen.model
.forward_gpu_greedy(&[next_token], &decode_positions, kv_cache, slot_id)
.with_context(|| {
format!(
"forward_gpu_greedy slot-aware decode step {step} \
(soft tokens; ADR-040 §6.1.50 iter-G)"
)
})?
} else {
let logits_full = qwen
.model
.forward_gpu_last_logits(&[next_token], &decode_positions, kv_cache, slot_id)
.with_context(|| {
format!(
"forward_gpu_last_logits slot-aware decode step {step} \
(soft tokens)"
)
})?;
let mut logits = logits_full;
sample_logits_qwen35(&mut logits, params, &generated_tokens)
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
generated_tokens.push(next_token);
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
decoded_text.push_str(&fragment);
if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
break;
}
}
}
let decode_duration = decode_start.elapsed();
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-4: reset_for_slot at exit")?;
let (content, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (decoded_text, None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &tok in &generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
Ok(GenerationResult {
text: content,
reasoning_text,
prompt_tokens: prompt_len,
completion_tokens: generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: 0,
logprobs: logprobs_vec,
})
}
#[allow(clippy::too_many_arguments)]
pub fn generate_qwen35_once_with_soft_tokens_and_deepstack_slot_aware(
qwen: &mut Qwen35LoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[crate::serve::forward_prefill::SoftTokenInjection<'_>],
deepstack: Option<&crate::serve::forward_prefill::DeepstackInjection<'_>>,
positions_flat: Option<&[i32]>,
params: &SamplingParams,
registration: Option<&ModelRegistration>,
kv_cache: &mut HybridKvCache,
slot_id: SlotId,
) -> Result<GenerationResult> {
if soft_tokens.is_empty() && deepstack.is_none() && positions_flat.is_none() {
return generate_qwen35_once_slot_aware(
qwen,
prompt_tokens,
params,
registration,
kv_cache,
slot_id,
);
}
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen35_once_with_soft_tokens_and_deepstack_slot_aware: \
empty prompt_tokens"
);
anyhow::ensure!(
slot_id.0 < kv_cache.n_seqs,
"generate_qwen35_once_with_soft_tokens_and_deepstack_slot_aware: \
SlotOutOfRange slot={} max_slots={} (ADR-040 iter-C2d-cont-kernel iter-4)",
slot_id.0,
kv_cache.n_seqs,
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let need_seq = prompt_len + max_tokens + 64;
if need_seq > kv_cache.max_seq_len as usize {
return Err(anyhow::anyhow!(
"generate_qwen35_once_with_soft_tokens_and_deepstack_slot_aware: \
per-request need_seq={} exceeds persistent cache max_seq_len={} \
(slot={} prompt_len={} max_tokens={}). ADR-040 iter-C2d-cont-\
kernel iter-4 sizes the persistent cache to \
cfg.max_position_embeddings; reduce max_tokens or use a shorter \
prompt.",
need_seq,
kv_cache.max_seq_len,
slot_id.0,
prompt_len,
max_tokens
));
}
let is_greedy = is_greedy_eligible(params);
let want_logprobs = params.logprobs;
let mut logprobs_vec: Option<Vec<f32>> = if want_logprobs {
Some(Vec::with_capacity(max_tokens))
} else {
None
};
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-4: reset_for_slot at entry (deepstack)")?;
let prefill_start = Instant::now();
let positions_owned: Vec<i32>;
let positions: &[i32] = match positions_flat {
Some(p) => {
anyhow::ensure!(
p.len() == 4 * prompt_len,
"generate_qwen35_once_with_soft_tokens_and_deepstack_slot_aware: \
positions_flat.len() = {} != 4 * prompt_len = {}",
p.len(),
4 * prompt_len
);
p
}
None => {
positions_owned = prefill_positions_for(prompt_len);
&positions_owned
}
};
let prefill_logits = qwen
.model
.forward_gpu_last_logits_with_soft_tokens_and_deepstack(
prompt_tokens,
positions,
soft_tokens,
deepstack,
kv_cache,
slot_id,
)
.context(
"Qwen35Model::forward_gpu_last_logits_with_soft_tokens_and_deepstack \
(slot-aware prefill, ADR-040 iter-C2d-cont-kernel iter-4)",
)?;
anyhow::ensure!(
prefill_logits.len() == qwen.vocab_size,
"qwen35 slot-aware deepstack prefill logits len {} != vocab_size {}",
prefill_logits.len(),
qwen.vocab_size
);
let mut next_token: u32 = if want_logprobs {
let mut logits = prefill_logits.clone();
let (tok, lp) = sample_logits_qwen35_with_logprob(&mut logits, params, &[]);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
greedy_argmax_last_token(&prefill_logits, qwen.vocab_size as u32)
} else {
let mut logits = prefill_logits.clone();
sample_logits_qwen35(&mut logits, params, &[])
};
let prefill_duration = prefill_start.elapsed();
let t_post: i32 = match positions_flat {
Some(p) => {
let mut max_t = 0i32;
for i in 0..prompt_len {
let v = p[i]; if v > max_t {
max_t = v;
}
}
max_t.saturating_add(1)
}
None => prompt_len as i32,
};
let decode_start = Instant::now();
let mut generated_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
generated_tokens.push(next_token);
let first_fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
let mut decoded_text = first_fragment.clone();
let mut finish_reason: &'static str = "length";
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
} else if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
} else {
for step in 1..max_tokens {
let pos = t_post + (step as i32 - 1);
if pos as u32 >= kv_cache.max_seq_len {
tracing::warn!(
pos,
max_seq = kv_cache.max_seq_len,
"qwen35 slot-aware decode (deepstack): hit kv-cache bound; \
stopping with finish=length",
);
break;
}
let decode_positions = vec![pos; 4];
next_token = if want_logprobs {
let logits_full = qwen
.model
.forward_gpu_last_logits(&[next_token], &decode_positions, kv_cache, slot_id)
.with_context(|| {
format!(
"forward_gpu_last_logits slot-aware decode step {step} \
(deepstack, logprobs)"
)
})?;
let mut logits = logits_full;
let (tok, lp) =
sample_logits_qwen35_with_logprob(&mut logits, params, &generated_tokens);
if let Some(v) = logprobs_vec.as_mut() {
v.push(lp);
}
tok
} else if is_greedy {
qwen.model
.forward_gpu_greedy(&[next_token], &decode_positions, kv_cache, slot_id)
.with_context(|| {
format!(
"forward_gpu_greedy slot-aware decode step {step} \
(deepstack; ADR-040 §6.1.50 iter-G)"
)
})?
} else {
let logits_full = qwen
.model
.forward_gpu_last_logits(&[next_token], &decode_positions, kv_cache, slot_id)
.with_context(|| {
format!(
"forward_gpu_last_logits slot-aware decode step {step} \
(deepstack)"
)
})?;
let mut logits = logits_full;
sample_logits_qwen35(&mut logits, params, &generated_tokens)
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
generated_tokens.push(next_token);
let fragment = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
decoded_text.push_str(&fragment);
if qwen35_hit_stop_string(&decoded_text, ¶ms.stop_strings) {
finish_reason = "stop";
qwen35_strip_trailing_stop(&mut decoded_text, ¶ms.stop_strings);
break;
}
}
}
let decode_duration = decode_start.elapsed();
kv_cache
.reset_for_slot(slot_id)
.context("ADR-040 iter-C2d-cont-kernel iter-4: reset_for_slot at exit (deepstack)")?;
let (content, reasoning_text) = match registration {
Some(reg) if reg.has_reasoning() => super::registry::split_full_output_forced(
reg,
&decoded_text,
params.reasoning_forced_open,
),
_ => (decoded_text, None),
};
let reasoning_token_count = match registration {
Some(reg) if reg.has_reasoning() => {
let mut sp =
super::registry::make_reasoning_splitter(reg, params.reasoning_forced_open);
let mut count = 0usize;
for &tok in &generated_tokens {
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
if let Some(splitter) = sp.as_mut() {
let _ = splitter.feed(&frag);
if splitter.in_reasoning() {
count += 1;
}
}
}
count
}
_ => 0,
};
Ok(GenerationResult {
text: content,
reasoning_text,
prompt_tokens: prompt_len,
completion_tokens: generated_tokens.len(),
reasoning_tokens: if reasoning_token_count > 0 {
Some(reasoning_token_count)
} else {
None
},
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: 0,
logprobs: logprobs_vec,
})
}
fn params_tool_call_policy_for_qwen35_stream() -> super::engine::ToolCallPolicy {
super::engine::ToolCallPolicy::Auto
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::models::qwen35::kv_cache::HybridKvCache;
use crate::inference::models::qwen35::{
default_layer_types, Qwen35Config, Qwen35MoeConfig, Qwen35Variant,
};
use mlx_native::{MlxBuffer, MlxDevice};
use std::cell::Cell;
fn thinking_budget_params(limit: usize) -> SamplingParams {
SamplingParams {
reasoning_forced_open: true,
thinking_token_budget: Some(limit),
reasoning_end_tokens: Some(Arc::new(vec![90, 91])),
reasoning_close_tokens: Some(Arc::new(vec![91])),
..SamplingParams::default()
}
}
#[test]
fn thinking_budget_forces_close_then_stops_overriding_answer_tokens() {
let mut budget =
Qwen35ThinkingBudgetState::from_params(&thinking_budget_params(2)).unwrap();
let mut generated = Vec::new();
for token in [11, 12] {
generated.push(token);
budget.observe_generated(&generated, false);
}
assert_eq!(budget.next_forced_token(), Some((90, true)));
assert!(!budget.was_forced_closed());
generated.push(90);
budget.observe_generated(&generated, false);
assert_eq!(budget.next_forced_token(), Some((91, false)));
generated.push(91);
budget.observe_generated(&generated, false);
assert_eq!(budget.next_forced_token(), None);
assert!(budget.closed);
assert!(budget.was_forced_closed());
}
#[test]
fn thinking_budget_honors_natural_reasoning_and_tool_boundaries() {
let mut natural =
Qwen35ThinkingBudgetState::from_params(&thinking_budget_params(4)).unwrap();
natural.observe_generated(&[11, 91], false);
assert!(natural.closed);
assert_eq!(natural.next_forced_token(), None);
let mut tool = Qwen35ThinkingBudgetState::from_params(&thinking_budget_params(4)).unwrap();
tool.observe_generated(&[11], true);
assert!(tool.closed);
assert_eq!(tool.next_forced_token(), None);
}
#[test]
fn bounded_prefill_splits_exactly_at_the_rewriteable_prompt_boundary() {
assert_eq!(qwen35_next_prefill_end(0, 9_000, 4_096, Some(3_000)), 3_000);
assert_eq!(
qwen35_next_prefill_end(3_000, 9_000, 4_096, Some(3_000)),
7_096
);
assert_eq!(
qwen35_next_prefill_end(7_096, 9_000, 4_096, Some(3_000)),
9_000
);
assert_eq!(qwen35_next_prefill_end(0, 9_000, 4_096, None), 4_096);
}
fn f32_buffer(device: &MlxDevice, rows: usize, hidden: usize, values: &[f32]) -> MlxBuffer {
assert_eq!(values.len(), rows * hidden);
let mut buffer = device
.alloc_buffer(
values.len() * std::mem::size_of::<f32>(),
mlx_native::DType::F32,
vec![rows, hidden],
)
.expect("allocate vision test buffer");
buffer
.as_mut_slice::<f32>()
.expect("vision test buffer f32 view")
.copy_from_slice(values);
buffer
}
#[test]
fn bounded_vision_chunk_rebases_soft_tokens_deepstack_and_mrope() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let hidden = 2;
let soft = f32_buffer(
&device,
6,
hidden,
&[
0.0, 1.0, 10.0, 11.0, 20.0, 21.0, 30.0, 31.0, 40.0, 41.0, 50.0, 51.0,
],
);
let deep = f32_buffer(
&device,
4,
hidden,
&[100.0, 101.0, 110.0, 111.0, 120.0, 121.0, 130.0, 131.0],
);
let prompt_len = 10;
let mut positions = Vec::with_capacity(4 * prompt_len);
for axis in 0..4_i32 {
positions.extend((0..prompt_len as i32).map(|position| axis * 100 + position));
}
let vision = Qwen35VisionPrefillData::new(
vec![SoftTokenData {
range: 2..8,
embeddings: soft,
}],
Some(DeepstackData {
image_token_positions: vec![3, 4, 5, 6],
chunks: vec![deep],
}),
Some(positions),
);
vision.validate(prompt_len, hidden).expect("valid vision");
let chunk = vision.chunk(4, 6, hidden).expect("middle chunk");
assert_eq!(chunk.soft_tokens.len(), 1);
assert_eq!(chunk.soft_tokens[0].0, 0..2);
assert_eq!(
chunk.soft_tokens[0].1.as_slice::<f32>().expect("soft view"),
&[20.0, 21.0, 30.0, 31.0]
);
let (deep_positions, deep_chunks) = chunk.deepstack.expect("deepstack chunk");
assert_eq!(deep_positions, vec![0, 1]);
assert_eq!(
deep_chunks[0].as_slice::<f32>().expect("deepstack view"),
&[110.0, 111.0, 120.0, 121.0]
);
assert_eq!(
chunk.positions_flat.expect("mRoPE chunk"),
vec![4, 5, 104, 105, 204, 205, 304, 305]
);
}
#[test]
fn bounded_vision_decode_base_uses_text_axis_not_expanded_prompt_length() {
let prompt_len = 8;
let mut positions = vec![0; 4 * prompt_len];
positions[..prompt_len].copy_from_slice(&[0, 1, 2, 3, 4, 4, 5, 6]);
let vision = Qwen35VisionPrefillData::new(Vec::new(), None, Some(positions));
assert_eq!(vision.decode_position_base(prompt_len), 7);
}
#[test]
fn failed_embed_forward_never_resets_a_potentially_poisoned_slot() {
let reset_called = Cell::new(false);
let forward: Result<Vec<f32>> = Err(anyhow::Error::new(
mlx_native::MlxError::CommandBufferError(
"Caused GPU Timeout Error (00000002)".to_string(),
),
))
.context("Qwen35 slot-aware embed forward");
let error = finish_qwen35_slot_embed(forward, || {
reset_called.set(true);
Ok(())
})
.expect_err("typed Metal failure must escape before reset");
assert!(!reset_called.get());
assert!(error.chain().any(|source| {
matches!(
source.downcast_ref::<mlx_native::MlxError>(),
Some(mlx_native::MlxError::CommandBufferError(_))
)
}));
}
#[test]
fn serial_kv_cache_capacity_grows_geometrically() {
let maximum = 262_144;
assert_eq!(serial_kv_cache_capacity(5_366, 0, maximum), 8_192);
assert_eq!(serial_kv_cache_capacity(5_579, 8_192, maximum), 16_384);
assert_eq!(serial_kv_cache_capacity(9_000, 8_192, maximum), 16_384);
assert_eq!(serial_kv_cache_capacity(200_000, 131_072, maximum), maximum);
}
#[derive(Debug)]
struct ResumeTestPayload(u64);
impl crate::serve::kv_persist::lcp_registry::ByteSized for ResumeTestPayload {
fn byte_len(&self) -> u64 {
self.0
}
}
#[test]
fn latest_turn_checkpoint_covers_short_prompts_and_longer_stride_wins() {
use crate::serve::kv_persist::format::ModelFingerprint;
use crate::serve::kv_persist::lcp_registry::{LcpKey, LcpRegistry};
use std::sync::Arc;
let base_key = LcpKey {
model_fingerprint: ModelFingerprint([9; 32]),
tenant_id: String::new(),
params_hash: 0,
};
let mut registry = LcpRegistry::new(8);
registry
.store(
base_key.clone(),
vec![1, 2, 3, 4, 5],
vec![Arc::new(ResumeTestPayload(1))],
0,
0,
)
.unwrap();
let short =
lookup_qwen35_resume_checkpoint(&mut registry, &base_key, &[1, 2, 3, 4, 5, 6], 8)
.expect("latest-turn checkpoint must cover prompt shorter than stride");
assert_eq!((short.0.k, short.1), (5, 0));
let mut chunk_key = base_key.clone();
chunk_key.tenant_id = "qwen35:lcp_chunk:8".into();
registry
.store(
chunk_key,
vec![1, 2, 3, 4, 5, 6, 7, 8],
vec![Arc::new(ResumeTestPayload(1))],
0,
0,
)
.unwrap();
let longest = lookup_qwen35_resume_checkpoint(
&mut registry,
&base_key,
&[1, 2, 3, 4, 5, 6, 7, 8, 9],
8,
)
.expect("stride checkpoint");
assert_eq!((longest.0.k, longest.1), (8, 8));
assert!(
lookup_qwen35_resume_checkpoint(&mut registry, &base_key, &[1, 2, 99, 4, 5, 6], 8,)
.is_none()
);
}
#[test]
fn recovery_tail_uses_only_a_verified_prompt_suffix() {
let prompt = [10, 11, 12, 20, 21, 22, 23];
assert_eq!(recovery_tail_for_suffix(&prompt, &[20, 21, 22, 23], 64), 4);
assert_eq!(
recovery_tail_for_suffix(&prompt, &[21, 22, 99], 64),
64,
"a custom or drifted template must retain the conservative fallback"
);
assert_eq!(recovery_tail_for_suffix(&prompt, &[], 64), 64);
}
#[test]
fn recovery_anchor_replaces_exact_or_immediately_preceding_stride_checkpoint() {
assert!(stride_checkpoint_superseded_by_recovery_anchor(
true, 4096, 4096, 5042
));
assert!(stride_checkpoint_superseded_by_recovery_anchor(
true, 4096, 4096, 4096
));
assert!(!stride_checkpoint_superseded_by_recovery_anchor(
true, 4096, 4096, 8192
));
assert!(!stride_checkpoint_superseded_by_recovery_anchor(
false, 4096, 4096, 5042
));
}
#[test]
fn recovery_capture_plan_is_limited_to_short_non_chunked_suffixes() {
assert_eq!(
qwen35_recovery_capture_plan(5_239, 5_255, 5_259, true, false),
Some((20, 15))
);
assert_eq!(
qwen35_recovery_capture_plan(5_042, 5_239, 5_243, true, false),
None,
"201 changed tokens stay on the normal prefill kernel"
);
assert_eq!(
qwen35_recovery_capture_plan(5_239, 5_255, 5_259, true, true),
None,
"chunked prefill owns its checkpoint schedule"
);
assert_eq!(
qwen35_recovery_capture_plan(5_239, 5_255, 5_259, false, false),
None
);
}
fn moe_cfg_40layer_for_cache_test() -> Qwen35Config {
Qwen35Config {
variant: Qwen35Variant::Moe,
hidden_size: 64,
num_hidden_layers: 4,
num_attention_heads: 4,
num_key_value_heads: 2,
head_dim: 16,
linear_num_key_heads: 4,
linear_num_value_heads: 8,
linear_key_head_dim: 16,
linear_value_head_dim: 16,
linear_conv_kernel_dim: 4,
full_attention_interval: 4,
layer_types: default_layer_types(4, 4),
partial_rotary_factor: 0.25,
rope_theta: 1e7,
rotary_dim: 4,
mrope_section: [1, 1, 0, 0],
mrope_interleaved: true,
rms_norm_eps: 1e-6,
max_position_embeddings: 1024,
vocab_size: 256,
attn_output_gate: true,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: true,
intermediate_size: None,
moe: Some(Qwen35MoeConfig {
moe_intermediate_size: 16,
num_experts: 4,
num_experts_per_tok: 2,
shared_expert_intermediate_size: 16,
}),
}
}
fn greedy_params() -> SamplingParams {
SamplingParams {
max_tokens: 16,
..SamplingParams::default()
}
}
#[test]
fn mtp_server_gate_rejects_non_default_server_semantics() {
let base = greedy_params();
assert!(is_qwen_server_speculation_exact_eligible(&base));
let mut frequency = base.clone();
frequency.frequency_penalty = 0.25;
assert!(!is_qwen_server_speculation_exact_eligible(&frequency));
let mut repetition = base.clone();
repetition.repetition_penalty = 1.05;
assert!(is_qwen_server_speculation_exact_eligible(&repetition));
let mut forced_thinking = base.clone();
forced_thinking.reasoning_forced_open = true;
forced_thinking.thinking_token_budget = Some(64);
forced_thinking.reasoning_end_tokens = Some(Arc::new(vec![90, 91]));
forced_thinking.reasoning_close_tokens = Some(Arc::new(vec![91]));
assert!(is_qwen_server_speculation_exact_eligible(&forced_thinking));
assert!(!is_serial_mtp_exact_eligible(&forced_thinking));
let mut lazy_tool = base.clone();
lazy_tool.tool_call_policy = ToolCallPolicy::AutoLazyGrammar;
assert!(!is_qwen_server_speculation_exact_eligible(&lazy_tool));
lazy_tool.grammar =
Some(crate::serve::api::grammar::parse("root ::= \"x\"\n").expect("test grammar"));
assert!(is_qwen_server_speculation_exact_eligible(&lazy_tool));
let mut stop = base;
stop.stop_strings.push("END".to_string());
assert!(!is_qwen_server_speculation_exact_eligible(&stop));
}
#[test]
fn repetition_window_includes_prompt_tail_and_tracks_committed_tokens() {
let prompt: Vec<u32> = (0..80).collect();
let mut history = qwen35_prompt_sampling_history(&prompt);
assert_eq!(history, (16..80).collect::<Vec<_>>());
qwen35_observe_sampling_history(&mut history, 80);
assert_eq!(history, (17..=80).collect::<Vec<_>>());
let mut params = greedy_params();
params.repetition_penalty = 1.05;
let mut logits = vec![0.0; 96];
logits[17] = 10.0;
logits[3] = 9.9;
let (token, _) =
sample_logits_qwen35_constrained(&mut logits, ¶ms, &history, None, false)
.expect("prompt-aware repetition sample");
assert_eq!(
token, 3,
"a token inside the last-64 prompt/generation window must be penalized"
);
}
#[test]
fn mtp_k3_full_accept_queues_three_drafts_and_bonus() {
let canonical = [
Qwen35CanonicalDecision {
token: 11,
terminal: false,
},
Qwen35CanonicalDecision {
token: 12,
terminal: false,
},
Qwen35CanonicalDecision {
token: 13,
terminal: false,
},
Qwen35CanonicalDecision {
token: 14,
terminal: false,
},
];
let plan = plan_qwen35_verified_block(&[11, 12, 13], |row| Ok(canonical[row]))
.expect("full accept plan");
assert_eq!(
plan.output.into_iter().collect::<Vec<_>>(),
vec![11, 12, 13, 14]
);
assert_eq!(plan.matched_drafts, 3);
assert_eq!(plan.rejected_drafts, 0);
assert_eq!(plan.valid_input_tokens, 4);
assert_eq!(plan.carry_hidden_row, 3);
assert!(!plan.terminal_after_pending);
}
#[test]
fn mtp_k3_partial_reject_keeps_only_target_valid_prefix() {
let canonical = [
Qwen35CanonicalDecision {
token: 11,
terminal: false,
},
Qwen35CanonicalDecision {
token: 99,
terminal: false,
},
];
let plan = plan_qwen35_verified_block(&[11, 12, 13], |row| Ok(canonical[row]))
.expect("partial reject plan");
assert_eq!(plan.output.into_iter().collect::<Vec<_>>(), vec![11, 99]);
assert_eq!(plan.matched_drafts, 1);
assert_eq!(plan.rejected_drafts, 1);
assert_eq!(plan.valid_input_tokens, 2);
assert_eq!(plan.carry_hidden_row, 1);
assert!(!plan.terminal_after_pending);
}
#[test]
fn mtp_k3_terminal_draft_is_neither_streamed_nor_retained() {
let canonical = [
Qwen35CanonicalDecision {
token: 11,
terminal: false,
},
Qwen35CanonicalDecision {
token: 12,
terminal: true,
},
];
let plan = plan_qwen35_verified_block(&[11, 12, 13], |row| Ok(canonical[row]))
.expect("terminal plan");
assert_eq!(plan.output.into_iter().collect::<Vec<_>>(), vec![11]);
assert_eq!(plan.matched_drafts, 2);
assert_eq!(plan.rejected_drafts, 0);
assert_eq!(plan.valid_input_tokens, 2);
assert_eq!(plan.carry_hidden_row, 1);
assert!(plan.terminal_after_pending);
}
#[test]
fn pending_mtp_bonus_is_emitted_before_history_can_select_a_new_round() {
let mut pending = VecDeque::from([41, 42]);
assert_eq!(take_pending_speculation_output(&mut pending), Some(41));
assert_eq!(take_pending_speculation_output(&mut pending), Some(42));
assert!(take_pending_speculation_output(&mut pending).is_none());
}
#[test]
fn history_miss_routes_to_mtp_until_two_negative_cost_windows() {
let mut cost = SpeculationCostController::new();
assert!(!may_route_history_miss_to_mtp(true, &cost));
cost.observe_ordinary_target(Duration::from_millis(10));
assert!(may_route_history_miss_to_mtp(true, &cost));
for _ in 0..4 {
cost.observe_speculative_round(1, Duration::from_millis(10));
}
assert!(may_route_history_miss_to_mtp(true, &cost));
for _ in 0..4 {
cost.observe_speculative_round(1, Duration::from_millis(10));
}
assert!(!may_route_history_miss_to_mtp(true, &cost));
assert!(!may_route_history_miss_to_mtp(false, &cost));
}
#[test]
fn mtp_transaction_rollback_restores_target_and_drafter_cursors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut cfg = moe_cfg_40layer_for_cache_test();
cfg.mtp_num_hidden_layers = 1;
let mut kv = HybridKvCache::new(&cfg, &device, 16, 1).expect("kv");
for slot in &mut kv.full_attn {
slot.current_len[0] = 7;
}
kv.mtp_slot.as_mut().expect("MTP slot").current_len[0] = 7;
rollback_slot_mtp_transaction(&mut kv, SlotId(0), 4, 4).expect("rollback");
assert!(kv.full_attn.iter().all(|slot| slot.current_len[0] == 4));
assert_eq!(kv.mtp_slot.as_ref().expect("MTP slot").current_len[0], 4);
}
fn initialized_prompt_cache_snapshot(
cfg: &Qwen35Config,
device: &MlxDevice,
prompt_len: usize,
) -> HybridKvCacheSnapshot {
let mut kv = HybridKvCache::new(cfg, device, 16, 1).expect("kv");
for slot in &mut kv.full_attn {
for buf in [slot.k.as_mut(), slot.v.as_mut()].into_iter().flatten() {
buf.as_mut_slice::<u8>()
.expect("full-attention test buffer")
.fill(0x5a);
}
slot.current_len[0] = prompt_len as u32;
}
if let Some(slot) = kv.mtp_slot.as_mut() {
for buf in [slot.k.as_mut(), slot.v.as_mut()].into_iter().flatten() {
buf.as_mut_slice::<u8>()
.expect("MTP test buffer")
.fill(0xa5);
}
slot.current_len[0] = prompt_len as u32;
}
kv.snapshot(device).expect("cursor-bounded snapshot")
}
#[test]
fn gpu_greedy_gate_rejects_cpu_sampling_features() {
assert!(is_greedy_eligible(&greedy_params()));
let mut repetition = greedy_params();
repetition.repetition_penalty = 1.05;
assert!(!is_greedy_eligible(&repetition));
let mut biased = greedy_params();
biased.logit_bias.insert(7, 2.0);
assert!(!is_greedy_eligible(&biased));
let mut grammar = greedy_params();
grammar.grammar =
Some(crate::serve::api::grammar::parse("root ::= \"a\"\n").expect("test grammar"));
assert!(!is_greedy_eligible(&grammar));
}
#[test]
fn constrained_sampler_rejects_higher_invalid_logit() {
let grammar = crate::serve::api::grammar::parse("root ::= \"a\"\n").expect("test grammar");
let mut params = greedy_params();
params.grammar = Some(grammar);
params.token_bytes = Some(std::sync::Arc::new(vec![b"x".to_vec(), b"a".to_vec()]));
let runtime = grammar_runtime_for_request(¶ms, None)
.expect("runtime build")
.expect("grammar runtime");
let mut logits = vec![100.0, 1.0];
let (token, logprob) =
sample_logits_qwen35_constrained(&mut logits, ¶ms, &[], Some(&runtime), false)
.expect("valid grammar/table invariant");
assert_eq!(token, 1);
assert_eq!(logprob, None);
}
#[test]
fn agentic_grammar_contract_sampler_rejects_missing_short_and_long_token_tables() {
let grammar = crate::serve::api::grammar::parse("root ::= \"a\"\n").expect("test grammar");
let mut params = greedy_params();
params.grammar = Some(grammar.clone());
let missing = grammar_runtime_for_request(¶ms, None)
.expect_err("grammar without token table must fail before sampling");
assert!(missing
.to_string()
.contains("without its authoritative token byte table"));
params.token_bytes = Some(std::sync::Arc::new(vec![b"a".to_vec()]));
let runtime = grammar_runtime_for_request(¶ms, None)
.expect("runtime build")
.expect("grammar runtime");
let sampler = SamplerPureParams {
temperature: 0.55,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 8,
seed: None,
};
let mut logits = vec![1.0, 0.0];
let short = sample_logits_with_grammar(
&mut logits,
&sampler,
&[],
Some(&runtime),
params.token_bytes.as_deref().map(Vec::as_slice),
false,
)
.expect_err("short token table must fail closed");
assert!(short
.to_string()
.contains("length 1 != logits vocabulary 2"));
params.token_bytes = Some(std::sync::Arc::new(vec![
b"a".to_vec(),
b"b".to_vec(),
b"c".to_vec(),
]));
let long = sample_logits_with_grammar(
&mut logits,
&sampler,
&[],
Some(&runtime),
params.token_bytes.as_deref().map(Vec::as_slice),
false,
)
.expect_err("long token table must fail closed");
assert!(long.to_string().contains("length 3 != logits vocabulary 2"));
}
#[test]
fn bounded_prefill_rejects_invalid_grammar_before_mutating_cache() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let cfg = moe_cfg_40layer_for_cache_test();
let mut kv = HybridKvCache::new(&cfg, &device, 128, 1).expect("kv");
let mut params = greedy_params();
params.grammar = Some(
crate::serve::api::grammar::parse("not_root ::= \"x\"\n")
.expect("syntactically valid grammar without root"),
);
let error = match Qwen35PrefillState::begin(
vec![7],
params,
None,
&mut kv,
SlotId(0),
0,
None,
None,
None,
cfg.hidden_size as usize,
) {
Ok(_) => panic!("missing-root grammar must fail before bounded prefill"),
Err(error) => error,
};
assert!(format!("{error:#}").contains("grammar has no root rule"));
assert_eq!(
kv.sequence_len_for_slot(SlotId(0)).expect("cursor"),
0,
"grammar validation must precede any cache mutation or Metal work"
);
}
#[test]
fn accepted_empty_grammar_token_terminates_the_prefill_seed() {
let grammar =
crate::serve::api::grammar::parse("root ::= \"\"\n").expect("empty terminal grammar");
let mut params = greedy_params();
params.grammar = Some(grammar);
params.token_bytes = Some(std::sync::Arc::new(vec![Vec::new()]));
let runtime = grammar_runtime_for_request(¶ms, None)
.expect("runtime build")
.expect("grammar runtime");
assert!(qwen35_grammar_terminal_token(Some(&runtime), ¶ms, 0));
let src = include_str!("engine_qwen35.rs");
assert!(
src.contains(
"|| qwen35_grammar_terminal_token(grammar_runtime.as_ref(), ¶ms, next_token)"
),
"bounded prefill seed must stop on the same grammar terminal as legacy SSE"
);
}
#[test]
fn cached_token_reporting_includes_partial_lcp_resume() {
assert_eq!(qwen35_reported_cached_tokens(17_132, false, 16_384), 16_384);
assert_eq!(qwen35_reported_cached_tokens(17_132, true, 16_384), 17_132);
assert_eq!(qwen35_reported_cached_tokens(17_132, false, 0), 0);
}
#[test]
fn cached_token_reporting_clamps_defensively_to_prompt_length() {
assert_eq!(qwen35_reported_cached_tokens(128, false, 256), 128);
}
fn sampling_params_with_temperature() -> SamplingParams {
SamplingParams {
temperature: 0.7,
max_tokens: 16,
..SamplingParams::default()
}
}
#[test]
fn hybrid_prompt_cache_new_is_empty() {
let cache = HybridPromptCache::new();
assert!(!cache.has_entry());
assert!(cache.snapshot().is_none());
assert!(cache.try_match(&[1, 2, 3], &greedy_params()).is_none());
}
#[test]
fn hybrid_prompt_cache_invalidates_on_prompt_divergence() {
let cfg = moe_cfg_40layer_for_cache_test();
let device = MlxDevice::new().expect("device");
let prompt = vec![10u32, 20, 30, 40];
let snap = initialized_prompt_cache_snapshot(&cfg, &device, prompt.len());
let mut cache = HybridPromptCache::new();
cache.update(prompt.clone(), snap, 99u32, &greedy_params());
assert!(cache.has_entry());
assert_eq!(
cache.try_match(&prompt, &greedy_params()),
Some(prompt.len()),
"exact-match prompt should hit"
);
let mut diverged = prompt.clone();
diverged[2] = 999;
assert!(
cache.try_match(&diverged, &greedy_params()).is_none(),
"divergent prompt should miss"
);
let mut shorter = prompt.clone();
shorter.pop();
assert!(
cache.try_match(&shorter, &greedy_params()).is_none(),
"shorter prompt should miss"
);
let mut longer = prompt.clone();
longer.push(50);
assert!(
cache.try_match(&longer, &greedy_params()).is_none(),
"longer prompt should miss"
);
}
#[test]
fn hybrid_prompt_cache_keys_vision_identity() {
let cfg = moe_cfg_40layer_for_cache_test();
let device = MlxDevice::new().expect("device");
let prompt = vec![10u32, 20, 30, 40];
let snapshot = initialized_prompt_cache_snapshot(&cfg, &device, prompt.len());
let mut cached_params = greedy_params();
cached_params.vision_fingerprint = Some([0x11; 32]);
let mut cache = HybridPromptCache::new();
cache.update(prompt.clone(), snapshot, 99, &cached_params);
assert_eq!(cache.try_match(&prompt, &cached_params), Some(prompt.len()));
let mut other_image = cached_params.clone();
other_image.vision_fingerprint = Some([0x22; 32]);
assert!(cache.try_match(&prompt, &other_image).is_none());
let mut text_only = cached_params;
text_only.vision_fingerprint = None;
assert!(cache.try_match(&prompt, &text_only).is_none());
}
#[test]
fn hybrid_prompt_cache_invalidates_on_genparams_mismatch() {
let cfg = moe_cfg_40layer_for_cache_test();
let device = MlxDevice::new().expect("device");
let prompt = vec![1u32, 2, 3];
let snap = initialized_prompt_cache_snapshot(&cfg, &device, prompt.len());
let mut cache = HybridPromptCache::new();
let stored_params = SamplingParams {
max_tokens: 32,
stop_strings: vec!["</done>".into()],
..SamplingParams::default()
};
cache.update(prompt.clone(), snap, 7u32, &stored_params);
assert_eq!(
cache.try_match(&prompt, &stored_params),
Some(prompt.len()),
"matching key should hit"
);
let diff_max = SamplingParams {
max_tokens: 64,
stop_strings: vec!["</done>".into()],
..SamplingParams::default()
};
assert!(
cache.try_match(&prompt, &diff_max).is_none(),
"max_tokens mismatch must miss"
);
let diff_stop = SamplingParams {
max_tokens: 32,
stop_strings: vec!["</STOP>".into()],
..SamplingParams::default()
};
assert!(
cache.try_match(&prompt, &diff_stop).is_none(),
"stop_strings mismatch must miss"
);
}
#[test]
fn hybrid_prompt_cache_sampling_mode_bypasses_lookup_and_store() {
let cfg = moe_cfg_40layer_for_cache_test();
let device = MlxDevice::new().expect("device");
let prompt = vec![1u32, 2, 3];
let snap = initialized_prompt_cache_snapshot(&cfg, &device, prompt.len());
let mut cache = HybridPromptCache::new();
cache.update(prompt.clone(), snap, 1u32, &greedy_params());
assert!(cache.has_entry());
assert!(
cache
.try_match(&prompt, &sampling_params_with_temperature())
.is_none(),
"sampling-mode lookup must miss"
);
let snap2 = initialized_prompt_cache_snapshot(&cfg, &device, prompt.len());
let mut cache2 = HybridPromptCache::new();
cache2.update(
prompt.clone(),
snap2,
5u32,
&sampling_params_with_temperature(),
);
assert!(!cache2.has_entry(), "sampling-mode update must be a no-op");
}
#[test]
fn splitter_helper_extracts_reasoning_from_qwen35_thinkblocks() {
let reg = super::super::registry::QWEN35;
let raw = "Sure! <think>Let me solve this step by step.</think>The answer is 42.";
let (content, reasoning) = super::super::registry::split_full_output(®, raw);
assert_eq!(
content, "Sure! The answer is 42.",
"content must exclude the <think>...</think> span"
);
assert_eq!(
reasoning.as_deref(),
Some("Let me solve this step by step."),
"reasoning must contain the inner span"
);
}
#[test]
fn splitter_helper_extracts_tool_calls_from_qwen35_toolblocks() {
let reg = super::super::registry::QWEN35;
let mut sp = super::super::registry::ToolCallSplitter::from_registration(®)
.expect("QWEN35 has tool markers");
let raw =
"Let me search.<tool_call><function=search><parameter=q>weather</parameter></function></tool_call> Done.";
let mut events = Vec::new();
events.extend(sp.feed(raw));
if let Some(tail) = sp.finish() {
events.push(tail);
}
let mut saw_open = false;
let mut saw_text = false;
let mut saw_close = false;
let mut content_runs: Vec<String> = Vec::new();
for ev in events {
use super::super::registry::ToolCallEvent::*;
match ev {
Content(t) => content_runs.push(t),
ToolCallOpen => saw_open = true,
ToolCallText(_) => saw_text = true,
ToolCallClose => saw_close = true,
}
}
assert!(saw_open, "must observe ToolCallOpen for QWEN35 marker");
assert!(saw_text, "must observe ToolCallText body");
assert!(saw_close, "must observe ToolCallClose");
let joined: String = content_runs.join("");
assert!(
joined.contains("Let me search."),
"preamble content must round-trip"
);
assert!(
joined.contains(" Done."),
"post-close content must round-trip"
);
}
#[test]
fn qwen35_loaded_model_has_initialized_prompt_cache() {
let cache = HybridPromptCache::default();
assert!(!cache.has_entry());
assert!(cache.snapshot().is_none());
assert!(cache.try_match(&[1, 2], &greedy_params()).is_none());
}
#[test]
fn qwen35_loaded_model_load_errors_when_path_missing() {
let opts = LoadOptions {
model_path: std::path::PathBuf::from("/tmp/iter-215-does-not-exist.gguf"),
tokenizer_path: None,
config_path: None,
dwq_overlay_path: None,
kv_persist_dir: None,
};
let res = Qwen35LoadedModel::load(&opts);
assert!(res.is_err());
let msg = format!("{:#}", res.err().unwrap());
assert!(
msg.contains("Model not found"),
"expected 'Model not found' in error; got: {msg}"
);
}
#[test]
fn qwen35_serving_tokenizer_is_sidecar_free_and_detects_vision_markers() {
use crate::backends::gguf::types::MetaValue;
use crate::backends::gguf::writer::GgufWriter;
let dir = tempfile::tempdir().expect("temporary sidecar-free directory");
let path = dir.path().join("qwen35-tokenizer.gguf");
let file = std::fs::File::create(&path).expect("create tokenizer GGUF fixture");
let metadata = [
("tokenizer.ggml.pre", MetaValue::String("qwen35".into())),
(
"tokenizer.ggml.tokens",
MetaValue::ArrayString(vec![
"<|vision_start|>".into(),
"<|image_pad|>".into(),
"<|vision_end|>".into(),
"a".into(),
]),
),
("tokenizer.ggml.merges", MetaValue::ArrayString(Vec::new())),
(
"tokenizer.ggml.token_type",
MetaValue::ArrayI32(vec![4, 4, 4, 1]),
),
];
let mut writer = GgufWriter::new(file);
writer
.write_header(0, metadata.len() as u64)
.expect("write fixture header");
for (key, value) in &metadata {
writer
.write_metadata_kv(key, value)
.expect("write tokenizer metadata");
}
writer.pad_to_alignment().expect("align fixture");
writer.finalize().expect("finalize fixture");
assert!(!dir.path().join("tokenizer.json").exists());
let gguf = mlx_native::gguf::GgufFile::open(&path).expect("open tokenizer fixture");
let (tokenizer, has_vision_markers) =
build_qwen35_serving_tokenizer(&gguf).expect("build embedded tokenizer");
assert!(has_vision_markers);
assert_eq!(tokenizer.token_to_id("<|image_pad|>"), Some(1));
}
use super::super::registry::{
SplitSlot as _SplitSlot, ToolCallEvent as _ToolCallEvent,
ToolCallSplitter as _ToolCallSplitter, QWEN35,
};
#[test]
fn wedge4e_reasoning_splitter_is_mode_invariant() {
let mut sp = super::super::registry::make_reasoning_splitter(&QWEN35, false)
.expect("Qwen35 has reasoning markers");
let mut all_pairs: Vec<(_SplitSlot, String)> = Vec::new();
for frag in ["<thi", "nk>let me reason", " more</thin", "k>final answer"] {
for pair in sp.feed(frag) {
all_pairs.push(pair);
}
}
if let Some(tail) = sp.finish() {
all_pairs.push(tail);
}
let mut reasoning = String::new();
let mut content = String::new();
for (slot, text) in &all_pairs {
match slot {
_SplitSlot::Reasoning => reasoning.push_str(text),
_SplitSlot::Content => content.push_str(text),
}
}
assert_eq!(
reasoning, "let me reason more",
"Wedge-4e: reasoning text must be cleanly extracted"
);
assert_eq!(
content, "final answer",
"Wedge-4e: content text must NOT contain reasoning brackets"
);
}
#[test]
fn wedge4e_tool_call_splitter_is_mode_invariant() {
let mut tcs =
_ToolCallSplitter::from_registration(&QWEN35).expect("Qwen35 has tool markers");
let mut events: Vec<_ToolCallEvent> = Vec::new();
for frag in [
"let me search.<tool_",
"call>{\"name\":\"search\",\"arguments\":{\"q\":\"x\"}}</tool_",
"call> done.",
] {
for ev in tcs.feed(frag) {
events.push(ev);
}
}
if let Some(tail) = tcs.finish() {
events.push(tail);
}
let mut saw_open = false;
let mut saw_close = false;
let mut body = String::new();
let mut content_runs: Vec<String> = Vec::new();
for ev in events {
match ev {
_ToolCallEvent::Content(t) => content_runs.push(t),
_ToolCallEvent::ToolCallOpen => saw_open = true,
_ToolCallEvent::ToolCallText(t) => body.push_str(&t),
_ToolCallEvent::ToolCallClose => saw_close = true,
}
}
assert!(saw_open, "Wedge-4e: must observe ToolCallOpen");
assert!(saw_close, "Wedge-4e: must observe ToolCallClose");
assert!(
body.contains("\"name\":\"search\""),
"Wedge-4e: tool-call body must round-trip; got {body:?}"
);
let joined: String = content_runs.join("");
assert!(
joined.contains("let me search."),
"Wedge-4e: pre-tool-call content must round-trip"
);
assert!(
joined.contains(" done."),
"Wedge-4e: post-tool-call content must round-trip"
);
}
#[test]
fn wedge4e_legacy_stream_entry_is_thin_wrapper() {
let src = include_str!("engine_qwen35.rs");
assert!(
src.contains("generate_stream_qwen35_once_extended(\n qwen,\n prompt_tokens,\n &[],\n None,\n None,"),
"Wedge-4e: generate_stream_qwen35_once must delegate to \
generate_stream_qwen35_once_extended with empty extensions \
— the byte-identical text-only regression contract \
requires this exact shape"
);
}
#[test]
fn wedge4e_extended_stream_validates_positions_len() {
let prompt_len = 5usize;
let bad_positions = vec![0i32; 17]; let expected_err = format!(
"qwen35 stream (wedge-4e): positions_flat.len() = {} != 4 * prompt_len = {}",
bad_positions.len(),
4 * prompt_len
);
let src = include_str!("engine_qwen35.rs");
assert!(
src.contains("qwen35 stream (wedge-4e): positions_flat.len() = "),
"Wedge-4e: positions_flat length validator must surface \
the actionable diagnostic byte string"
);
assert!(
expected_err.contains("17 != 4 * prompt_len = 20"),
"expected_err format check"
);
}
#[test]
fn wedge4e_t_post_advance_rule_matches_non_streaming_sibling() {
let src = include_str!("engine_qwen35.rs");
assert!(
src.contains("max_t.saturating_add(1)"),
"Wedge-4e: t_post must use saturating_add(1) over axis-0 max"
);
assert!(
src.contains("None => prompt_len as i32,"),
"Wedge-4e: t_post must default to prompt_len when no 3D \
positions supplied (text-only byte-identity)"
);
assert!(
src.contains("let pos = t_post + (step as i32 - 1);"),
"Wedge-4e: streaming decode position formula must use t_post \
advance — not the legacy (prompt_len + step - 1) form, \
which would silently misalign on multi-image prefill"
);
}
#[test]
fn wedge4e_extended_stream_bypasses_prompt_cache_on_extension() {
let src = include_str!("engine_qwen35.rs");
assert!(
src.contains("let prompt_cache_hit = !has_extension"),
"Wedge-4e: streaming prompt-cache MUST be bypassed when \
any extension is present (cache key is prompt_tokens \
only — same placeholder ids + different image ⇒ false \
hit)"
);
assert!(
src.contains("if is_greedy && !has_extension {"),
"Wedge-4e: streaming prompt-cache write MUST be skipped on \
extension paths to avoid poisoning subsequent text-only \
requests with a soft-token-tainted KV snapshot"
);
}
#[test]
fn lcp_prefix_stores_are_not_greedy_gated() {
let src = include_str!("engine_qwen35.rs");
let forbidden = concat!("lcp_resume_enabled", " && ", "is_greedy");
assert!(
!src.contains(forbidden),
"LCP prefix-KV stores must NOT be greedy-gated — KV state is \
sampling-independent; greedy gating belongs only on \
HybridPromptCache's decoded-token replay"
);
let recovery_guard = concat!("&& !superseded_by_", "recovery_anchor");
let kill_switch = concat!("&& !mid_store_", "disabled");
assert_eq!(
src.matches(recovery_guard).count(),
4,
"non-stream + stream mid-prefill notification and store gates \
must both be suppressed when a near-tip recovery anchor \
replaces the final stride checkpoint"
);
assert_eq!(
src.matches(kill_switch).count(),
2,
"non-stream + stream mid-prefill store gates must both retain \
the symmetric mid-store kill-switch"
);
assert!(
src.contains("if is_greedy && !has_extension {"),
"HybridPromptCache write must REMAIN greedy-gated (first \
decoded token replay is sampling-dependent)"
);
}
#[test]
fn wedge4e_multimodal_streaming_is_explicit_per_scheduler_mode() {
let src = include_str!("engine.rs");
let qwen = include_str!("engine_qwen35.rs");
assert!(
src.contains("Qwen35VisionPrefillData::new("),
"SlotAware Qwen must retain every multimodal stream field at admission"
);
assert!(
qwen.contains("vision.validate(prompt_len, hidden_size)?"),
"SlotAware Qwen must validate multimodal state before GPU prefill"
);
assert!(
qwen.contains("forward_gpu_last_logits_with_soft_tokens_and_deepstack("),
"SlotAware Qwen must route multimodal chunks through the vision-aware graph"
);
assert!(
src.contains("generate_stream_qwen35_once_extended"),
"SerialFifo must retain the extended multimodal stream primitive"
);
}
#[test]
fn wedge4e_handler_streaming_501_reject_is_removed() {
let src = include_str!("handlers.rs");
assert!(
!src.contains("streaming chat with Qwen3-VL DeepStack injection is not yet"),
"Wedge-4e: handler-side streaming 501 reject MUST be \
removed — streaming Qwen3-VL chat now flows through \
generate_stream_with_deepstack"
);
assert!(
src.contains("generate_stream_with_deepstack"),
"Wedge-4e: handler must call generate_stream_with_deepstack \
so soft_tokens + deepstack + positions reach the worker"
);
}
fn synth_loaded_model_for_alloc_test(
cfg: Qwen35Config,
tq_kv_active: bool,
) -> Qwen35LoadedModel {
Qwen35LoadedModel {
model: super::Qwen35Model::empty_from_cfg(cfg),
tokenizer: tokenizers::Tokenizer::new(tokenizers::models::bpe::BPE::default()),
chat_template: "{{ messages }}".to_string(),
model_id: "iter-12-test".to_string(),
model_path: std::path::PathBuf::from("/tmp/iter-12-test.gguf"),
eos_token_ids: vec![151_645],
hidden_size: 64,
vocab_size: 256,
context_length: Some(1024),
quant_type: Some("Q4_K".to_string()),
load_duration: std::time::Duration::from_millis(1),
provenance: crate::core::provenance::Provenance::External,
vision_projector_profile: None,
vision_deepstack_output_count: None,
vision_special_tokens_present: false,
prompt_cache: HybridPromptCache::new(),
lcp_registry: crate::serve::kv_persist::lcp_registry::LcpRegistry::new(1),
kv_metrics_sink: None,
disk_persistor: None,
lcp_hydrated_for_cfg: std::collections::HashSet::new(),
tq_kv_active,
persistent_kv_cache: None,
speculation: crate::serve::api::qwen35_speculation::QwenSpeculationController::new(
crate::serve::api::qwen35_speculation::QwenSpeculationPolicy::Off,
),
}
}
#[test]
fn alloc_kv_cache_for_request_tq_off_keeps_full_attn_tq_none() {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(e) => {
eprintln!("skipping: no Metal device: {e}");
return;
}
};
let cfg = moe_cfg_40layer_for_cache_test();
let qwen = synth_loaded_model_for_alloc_test(cfg, false);
let cache =
alloc_kv_cache_for_request(&qwen, &device, 32, 16).expect("alloc_kv_cache_for_request");
assert!(!cache.full_attn.is_empty(), "fixture has full-attn layers");
for (i, slot) in cache.full_attn.iter().enumerate() {
assert!(
slot.tq.is_none(),
"tq_kv_active=false: full_attn[{i}].tq must be None \
(legacy F32 path preserved)"
);
}
}
#[test]
fn alloc_kv_cache_for_request_tq_on_populates_tq_per_full_attn_slot() {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(e) => {
eprintln!("skipping: no Metal device: {e}");
return;
}
};
let mut cfg = moe_cfg_40layer_for_cache_test();
cfg.head_dim = 256;
cfg.num_attention_heads = 8;
cfg.num_key_value_heads = 2;
let qwen = synth_loaded_model_for_alloc_test(cfg, true);
let cache =
alloc_kv_cache_for_request(&qwen, &device, 32, 16).expect("alloc_kv_cache_for_request");
assert!(!cache.full_attn.is_empty());
for (i, slot) in cache.full_attn.iter().enumerate() {
assert!(
slot.tq.is_some(),
"tq_kv_active=true: full_attn[{i}].tq must be populated"
);
let tq = slot.tq.as_ref().unwrap();
assert_eq!(tq.norms_per_pos, 1, "head_dim=256 → norms_per_pos=1");
}
}
}