use std::path::PathBuf;
use std::time::{Duration, Instant};
use anyhow::{Context, Result};
use tokenizers::Tokenizer;
use crate::inference::models::qwen35::kv_cache::{
HybridKvCache, HybridKvCacheSnapshot, HybridKvSlotAnchor,
};
use crate::inference::models::qwen35::model::Qwen35Model;
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, LoadOptions, SamplingParams, SerialStreamEnd, SerialStreamResult,
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|>"];
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 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>,
}
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 tokenizer_path =
crate::serve::find_tokenizer(model_path, opts.tokenizer_path.as_deref())?;
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_path = tokenizer_path;
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 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,
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,
};
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>,
}
impl HybridPromptCacheKey {
pub fn from_params(params: &SamplingParams) -> Self {
Self {
max_tokens: params.max_tokens,
stop_strings: params.stop_strings.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())
}
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,
};
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,
};
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,
) -> (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,
};
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;
}
}
}
pub(super) fn generate_qwen35_once(
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 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(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>,
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 {
#[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>>,
) -> 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"
);
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 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,
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, 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, 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 positions = 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 forward =
qwen.model
.forward_gpu_last_logits(chunk, &positions, kv_cache, self.slot_id);
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) {
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(),
})
} else {
None
};
(logits, 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 state = Qwen35DecodeState::from_prefill_logits(
qwen,
self.prompt_tokens,
self.params,
registration,
self.slot_id,
self.cached_tokens,
&prefill_logits,
prefill_duration,
)?;
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,
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>,
decoded_text: String,
stop_strings: Vec<String>,
finish_reason: &'static str,
step: usize,
semantic_fragment_reported: bool,
prefill_duration: Duration,
decode_start: Instant,
}
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,
)?;
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,
) -> 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 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,
&[],
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);
let mut decoded_text = qwen
.tokenizer
.decode(&[next_token], false)
.unwrap_or_default();
if let Some(splitter) = tool_splitter.as_mut() {
let marker_events = splitter.feed(&decoded_text);
if marker_events
.iter()
.any(|event| matches!(event, ToolCallEvent::ToolCallOpen))
{
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";
}
Ok(Self {
slot_id,
prompt_tokens,
prompt_len,
max_tokens,
is_greedy,
want_logprobs,
logprobs_vec,
cached_tokens,
params,
grammar_runtime,
tool_splitter,
next_token,
generated_tokens,
decoded_text,
stop_strings,
finish_reason,
step: 1,
semantic_fragment_reported: false,
prefill_duration,
decode_start: Instant::now(),
})
}
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 prompt_cache_identity(&self) -> (&[u32], &SamplingParams) {
(&self.prompt_tokens, &self.params)
}
pub(crate) fn mark_first_semantic_fragment(&mut self) -> bool {
if self.semantic_fragment_reported {
false
} else {
self.semantic_fragment_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_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,
});
}
let pos = self.prompt_len + 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 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()?;
forward.with_context(|| {
format!(
"Qwen35Model::forward_gpu_greedy (slot-aware decode step {}; \
ADR-040 Phase F M1)",
self.step
)
})?
} 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 mut logits = logits;
let (token, logprob) = sample_logits_qwen35_constrained(
&mut logits,
&self.params,
&self.generated_tokens,
self.grammar_runtime.as_ref(),
self.want_logprobs,
);
if let (Some(values), Some(logprob)) = (self.logprobs_vec.as_mut(), logprob) {
values.push(logprob);
}
token
};
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);
let frag = qwen.tokenizer.decode(&[tok], false).unwrap_or_default();
self.decoded_text.push_str(&frag);
if let Some(splitter) = self.tool_splitter.as_mut() {
let marker_events = splitter.feed(&frag);
if marker_events
.iter()
.any(|event| matches!(event, ToolCallEvent::ToolCallOpen))
{
if let Some(runtime) = self.grammar_runtime.as_mut() {
runtime.trigger();
}
}
}
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,
})
}
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>,
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(r, body));
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>,
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,
events,
&text,
) {
return false;
}
}
}
}
true
} else {
route_content_qwen35_slot_aware(
tool_splitter,
body,
tc_index,
saw_tc,
registration,
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 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,
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, &[next_token]))
}
}
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;
}
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,
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,
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: if prompt_cache_hit {
Some(prompt_len)
} else {
None
},
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 wedge-4d generate): {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_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 prefill (wedge-4d) 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 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();
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_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_hit =
!has_extension && 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 && !has_extension {
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>,
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(r, body));
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>,
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,
events,
&text,
) {
return false;
}
}
}
}
true
} else {
route_content_qwen35(
tool_splitter,
body,
tc_index,
saw_tc,
tool_call_policy,
registration,
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,
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,
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,
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: (cached_tokens > 0).then_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::MlxDevice;
use std::cell::Cell;
#[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);
}
#[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()
}
}
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);
assert_eq!(token, 1);
assert_eq!(logprob, None);
}
#[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) {
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_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}"
);
}
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");
assert!(
src.contains("validate_qwen35_slot_stream_payload("),
"SlotAware Qwen must validate every multimodal stream field before admission"
);
assert!(
src.contains("reject_stream_before_sse(events, admission, error)"),
"unsupported SlotAware multimodal streams must fail before SSE"
);
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,
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,
}
}
#[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");
}
}
}