pub mod api;
#[allow(dead_code)]
pub mod auto_pipeline;
#[allow(dead_code)]
pub mod cache;
pub mod config;
mod deepseek4_cli;
pub mod encoder_worker_singleton;
pub mod forward_mlx_shared;
pub mod forward_prefill;
pub mod forward_prefill_batched;
pub mod gpu;
pub mod header;
#[allow(dead_code)]
pub mod kv_persist;
pub mod layer_ctx;
#[allow(dead_code)]
pub mod load_info;
#[allow(dead_code)]
pub mod multi_model;
#[allow(dead_code)]
pub mod multi_seq_kv;
pub mod parity_quality;
#[allow(dead_code)]
#[allow(dead_code)]
pub mod quant_select;
#[allow(dead_code)]
pub mod sampler_pure;
#[allow(dead_code)]
pub mod scheduler;
pub mod spec_decode_cli;
use anyhow::{Context, Result};
use std::path::Path;
use crate::cli;
use crate::debug::INVESTIGATION_ENV;
fn build_warmed_embedding_registry(
em: &api::state::EmbeddingModel,
) -> Result<mlx_native::KernelRegistry> {
use crate::inference::models::bert::bert_gpu::{
apply_bert_full_forward_gpu, register_bert_custom_shaders,
};
use crate::inference::models::nomic_bert::{
apply_nomic_bert_full_forward_gpu, register_nomic_bert_kernels,
};
use api::state::EmbeddingArch;
use mlx_native::{DType, KernelRegistry, MlxDevice};
let arch = em
.arch
.as_ref()
.ok_or_else(|| anyhow::anyhow!("registry warmup: EmbeddingModel has no arch"))?;
let device =
MlxDevice::new().map_err(|e| anyhow::anyhow!("registry warmup: MlxDevice::new: {e}"))?;
let mut registry = KernelRegistry::new();
let seq_len: u32 = 32;
let pad_id = em.tokenizer.specials().pad;
let ids: Vec<u32> = vec![pad_id; seq_len as usize];
let ids_buf = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.map_err(|e| anyhow::anyhow!("registry warmup: alloc ids: {e}"))?;
unsafe {
let s: &mut [u32] =
std::slice::from_raw_parts_mut(ids_buf.contents_ptr() as *mut u32, seq_len as usize);
s.copy_from_slice(&ids);
}
let mut encoder = device
.command_encoder()
.map_err(|e| anyhow::anyhow!("registry warmup: command_encoder: {e}"))?;
let valid_token_count: u32 = 1;
let _out = match arch {
EmbeddingArch::Bert { config, weights } => {
register_bert_custom_shaders(&mut registry);
apply_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&ids_buf,
None,
weights,
config,
seq_len,
valid_token_count,
)?
}
EmbeddingArch::NomicBert { config, weights } => {
register_nomic_bert_kernels(&mut registry);
apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&ids_buf,
None,
weights,
config,
seq_len,
valid_token_count,
)?
}
};
encoder
.commit_and_wait()
.map_err(|e| anyhow::anyhow!("registry warmup: commit_and_wait: {e}"))?;
tracing::info!(
arch = arch.arch_name(),
cached_pipelines = registry.cached_count(),
"Warmed embedding kernel registry"
);
Ok(registry)
}
pub(crate) fn find_tokenizer(
model_path: &Path,
explicit: Option<&Path>,
) -> Result<std::path::PathBuf> {
if let Some(p) = explicit {
return Ok(p.to_path_buf());
}
let dir = model_path.parent().unwrap_or(Path::new("."));
let candidate = dir.join("tokenizer.json");
if candidate.exists() {
return Ok(candidate);
}
let _stem = model_path.file_stem().unwrap_or_default().to_string_lossy();
for subdir in &["gemma4", "gemma-4"] {
let candidate = Path::new("models").join(subdir).join("tokenizer.json");
if candidate.exists() {
return Ok(candidate);
}
}
let models_dir = Path::new("models");
if models_dir.is_dir() {
for entry in std::fs::read_dir(models_dir)? {
let entry = entry?;
if entry.path().is_dir() {
let tok = entry.path().join("tokenizer.json");
if tok.exists() {
return Ok(tok);
}
}
}
}
anyhow::bail!(
"Cannot find tokenizer.json. Tried next to GGUF and in models/. \
Use --tokenizer to specify the path explicitly."
)
}
fn resolve_prompt(args: &cli::GenerateArgs) -> Result<String> {
match (&args.prompt, &args.prompt_file) {
(Some(text), _) => Ok(text.clone()),
(None, Some(path)) => {
let content = std::fs::read_to_string(path)
.with_context(|| format!("Failed to read prompt file: {}", path.display()))?;
let trimmed = content.trim().to_string();
anyhow::ensure!(
!trimmed.is_empty(),
"Prompt file is empty: {}",
path.display()
);
Ok(trimmed)
}
(None, None) => anyhow::bail!("Either --prompt or --prompt-file must be specified"),
}
}
fn detect_hardware_info() -> (String, u64) {
use crate::core::hardware::HardwareProfiler;
match HardwareProfiler::detect() {
Ok(profile) => {
let mem_gb = profile.total_memory_bytes / (1024 * 1024 * 1024);
(profile.chip_model, mem_gb)
}
Err(_) => ("Unknown".to_string(), 0),
}
}
const BENCH_NUM_RUNS: usize = 5;
const BENCH_GEN_LENS: &[usize] = &[200, 1000, 2500];
fn median_f64(values: &[f64]) -> f64 {
if values.is_empty() {
return 0.0;
}
let mut sorted: Vec<f64> = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = sorted.len();
if n % 2 == 1 {
sorted[n / 2]
} else {
(sorted[n / 2 - 1] + sorted[n / 2]) / 2.0
}
}
fn p95_f64(values: &[f64]) -> f64 {
if values.is_empty() {
return 0.0;
}
let mut sorted: Vec<f64> = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = sorted.len();
let rank = ((0.95 * n as f64).ceil() as usize)
.saturating_sub(1)
.min(n - 1);
sorted[rank]
}
fn print_benchmark_summary(
model_path: &Path,
prompt_tokens: usize,
generated_per_run: &[usize],
prefill_tps: Option<&[f64]>,
decode_tps: &[f64],
extras: &[(String, String)],
) {
let (chip, mem_gb) = detect_hardware_info();
let model_filename = model_path
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| "unknown".to_string());
let generated_str = if generated_per_run
.iter()
.all(|&g| Some(g) == generated_per_run.first().copied())
{
format!("{}", generated_per_run.first().copied().unwrap_or(0))
} else {
let parts: Vec<String> = generated_per_run.iter().map(|g| g.to_string()).collect();
parts.join(",")
};
println!();
println!("=== Benchmark Results ===");
println!("Hardware: {}, {} GB", chip, mem_gb);
println!("Model: {}", model_filename);
println!("Prompt tokens: {}", prompt_tokens);
println!("Generated tokens: {}", generated_str);
println!("Runs: {}", decode_tps.len());
if let Some(pp) = prefill_tps {
for (i, (&pp_tps, &dec_tps)) in pp.iter().zip(decode_tps.iter()).enumerate() {
println!(
"Run {}: prefill {:.1} tok/s, decode {:.1} tok/s",
i + 1,
pp_tps,
dec_tps
);
}
} else {
for (i, &dec_tps) in decode_tps.iter().enumerate() {
println!("Run {}: decode {:.1} tok/s", i + 1, dec_tps);
}
}
if let Some(pp) = prefill_tps {
println!("Prefill tok/s: {:.1}", median_f64(pp));
}
println!("Decode tok/s: {:.1}", median_f64(decode_tps));
println!("Median: {:.1} tok/s", median_f64(decode_tps));
println!("P95: {:.1} tok/s", p95_f64(decode_tps));
for (label, value) in extras {
println!("{}: {}", label, value);
}
}
pub(crate) const FALLBACK_GEMMA4_CHAT_TEMPLATE: &str =
"<bos><|turn>user\n{{PROMPT}}<turn|>\n<|turn>model\n<|channel>thought\n<channel|>";
pub(crate) const FALLBACK_GEMMA4_API_CHAT_TEMPLATE: &str = concat!(
"<bos>",
"{%- for m in messages -%}",
"<|turn>{{ m.role }}\n{{ m.content }}<turn|>\n",
"{%- endfor -%}",
"<|turn>model\n",
"<|channel>thought\n<channel|>",
);
fn template_supports_enable_thinking(template_str: &str) -> bool {
let render_true = render_jinja_template(template_str, "x", Some(true));
let render_false = render_jinja_template(template_str, "x", Some(false));
match (render_true, render_false) {
(Ok(enabled), Ok(disabled)) => rendered_prompt_opens_thinking(&enabled, &disabled),
_ => false,
}
}
fn rendered_prompt_opens_thinking(enabled: &str, disabled: &str) -> bool {
let e = enabled.trim_end();
let d = disabled.trim_end();
if e == d {
return false;
}
fn unclosed_thinking_count(s: &str) -> i32 {
let opens = s.matches("<think>").count() as i32
+ s.matches("<reasoning>").count() as i32
+ s.matches("<thinking>").count() as i32;
let closes = s.matches("</think>").count() as i32
+ s.matches("</reasoning>").count() as i32
+ s.matches("</thinking>").count() as i32;
opens - closes
}
unclosed_thinking_count(e) > unclosed_thinking_count(d)
}
fn resolve_enable_thinking(args: &cli::GenerateArgs, template_str: Option<&str>) -> Option<bool> {
if args.enable_thinking {
return Some(true);
}
if args.no_thinking {
return Some(false);
}
let auto = template_str
.map(template_supports_enable_thinking)
.unwrap_or(false);
Some(auto)
}
fn render_chat_template(
gguf: &mlx_native::gguf::GgufFile,
args: &cli::GenerateArgs,
tokenizer: Option<&tokenizers::Tokenizer>,
user_prompt: &str,
) -> Result<String> {
let template_str: String = if let Some(tmpl) = args.chat_template.as_deref() {
tracing::info!("Chat template: using CLI --chat-template override");
tmpl.to_string()
} else if let Some(path) = args.chat_template_file.as_deref() {
tracing::info!(
"Chat template: loading from --chat-template-file {}",
path.display()
);
std::fs::read_to_string(path)
.with_context(|| format!("Failed to read --chat-template-file {}", path.display()))?
} else if let Some(tmpl) = gguf.metadata_string("tokenizer.chat_template") {
tracing::info!(
"Chat template: using GGUF metadata tokenizer.chat_template ({} chars)",
tmpl.len()
);
tmpl.to_string()
} else {
tracing::warn!(
"Chat template: no GGUF metadata tokenizer.chat_template and no \
CLI override; falling back to hardcoded Gemma4 template"
);
return Ok(FALLBACK_GEMMA4_CHAT_TEMPLATE.replace("{{PROMPT}}", user_prompt));
};
let enable_thinking = resolve_enable_thinking(args, Some(template_str.as_str()));
let bos_token = resolve_token_text(gguf, tokenizer, "tokenizer.ggml.bos_token_id", "<bos>");
let eos_token = resolve_token_text(gguf, tokenizer, "tokenizer.ggml.eos_token_id", "<eos>");
render_jinja_template_with_specials(
&template_str,
user_prompt,
enable_thinking,
&bos_token,
&eos_token,
)
}
fn render_jinja_template(
template_str: &str,
user_prompt: &str,
enable_thinking: Option<bool>,
) -> Result<String> {
render_jinja_template_with_specials(
template_str,
user_prompt,
enable_thinking,
"<bos>",
"<eos>",
)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum RaisePolicy {
Lenient,
Strict,
}
pub(crate) fn build_chat_template_env<'s>(raise_policy: RaisePolicy) -> minijinja::Environment<'s> {
let mut env = minijinja::Environment::new();
env.add_filter("tojson", |v: minijinja::Value| {
serde_json::to_string(&v).unwrap_or_else(|_| "null".to_string())
});
match raise_policy {
RaisePolicy::Lenient => {
env.add_function("raise_exception", |msg: String| -> minijinja::Value {
tracing::warn!("chat template raise_exception: {}", msg);
minijinja::Value::UNDEFINED
});
}
RaisePolicy::Strict => {
env.add_function(
"raise_exception",
|msg: String| -> std::result::Result<minijinja::Value, minijinja::Error> {
Err(minijinja::Error::new(
minijinja::ErrorKind::InvalidOperation,
format!("chat template raise_exception: {msg}"),
))
},
);
}
}
env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
env
}
fn render_jinja_template_with_specials(
template_str: &str,
user_prompt: &str,
enable_thinking: Option<bool>,
bos_token: &str,
eos_token: &str,
) -> Result<String> {
let mut env = build_chat_template_env(RaisePolicy::Lenient);
env.add_template("chat", template_str)
.context("Failed to parse chat template as Jinja2")?;
let tmpl = env
.get_template("chat")
.context("Failed to load parsed chat template")?;
let rendered = tmpl
.render(minijinja::context! {
messages => vec![
minijinja::context! { role => "user", content => user_prompt }
],
add_generation_prompt => true,
bos_token => bos_token,
eos_token => eos_token,
enable_thinking => enable_thinking,
})
.context("Failed to render chat template")?;
Ok(rendered)
}
fn resolve_token_id(
gguf: &mlx_native::gguf::GgufFile,
tokenizer: &tokenizers::Tokenizer,
metadata_key: &str,
) -> Option<u32> {
llama_cpp_special_token_id(gguf, metadata_key)
.and_then(|id| tokenizer.id_to_token(id).map(|_| id))
}
fn llama_cpp_special_token_id(
gguf: &mlx_native::gguf::GgufFile,
metadata_key: &str,
) -> Option<u32> {
if let Some(id) = gguf.metadata_u32(metadata_key) {
return Some(id);
}
let tokenizer_model = gguf.metadata_string("tokenizer.ggml.model")?;
llama_cpp_special_token_id_for_model(tokenizer_model, metadata_key)
}
fn llama_cpp_special_token_id_for_model(tokenizer_model: &str, metadata_key: &str) -> Option<u32> {
match (tokenizer_model, metadata_key) {
("gpt2", "tokenizer.ggml.bos_token_id") | ("gpt2", "tokenizer.ggml.eos_token_id") => {
Some(11)
}
_ => None,
}
}
fn resolve_token_text(
gguf: &mlx_native::gguf::GgufFile,
tokenizer: Option<&tokenizers::Tokenizer>,
metadata_key: &str,
template_literal_when_unavailable: &str,
) -> String {
let Some(tokenizer) = tokenizer else {
return template_literal_when_unavailable.to_string();
};
resolve_token_id(gguf, tokenizer, metadata_key)
.and_then(|id| tokenizer.id_to_token(id))
.unwrap_or_else(|| template_literal_when_unavailable.to_string())
}
fn tokenize_rendered_prompt_llama_style(
gguf: &mlx_native::gguf::GgufFile,
tokenizer: &tokenizers::Tokenizer,
prompt_text: &str,
) -> Result<Vec<u32>> {
crate::core::tokenizer_adapter::tokenize_with_bos_eos_from_gguf(gguf, tokenizer, prompt_text)
}
pub fn cmd_generate(args: cli::GenerateArgs) -> Result<()> {
let model_path = &args.model;
anyhow::ensure!(
model_path.exists(),
"Model not found: {}",
model_path.display()
);
if let Some(bits) = &args.kv_bits {
std::env::set_var("HF2Q_TQ_CODEBOOK_BITS", bits);
tracing::info!("ADR-007 F-6.1: KV codebook bits set via --kv-bits {bits} (overrides HF2Q_TQ_CODEBOOK_BITS env)");
}
{
let gguf_peek = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF open (arch peek): {e}"))?;
if let Some(arch) = gguf_peek.metadata_string("general.architecture") {
use crate::inference::models::qwen35::{
is_qwen36_gguf, is_qwen3_vl_arch, is_qwen3_vl_moe_arch, ARCH_QWEN35, ARCH_QWEN35MOE,
};
if is_qwen3_vl_arch(arch) {
let variant_label = if is_qwen3_vl_moe_arch(arch) {
"MoE"
} else {
"dense"
};
anyhow::bail!(
"Qwen3-VL ({variant_label}, general.architecture = {arch:?}) GGUFs are recognized \
but the CLI `hf2q generate` path requires a working LM forward (iter-228b \
scope). iter-228a (this commit) lands the SERVE-side load path: \
`hf2q serve --model <gguf>` will load this GGUF and surface it on \
`/v1/models`, but chat completion still returns HTTP 501 with the \
`qwen3vl_text_forward_pending` sentinel until iter-228b wires the \
dense transformer forward. For text-only chat today, use a \
Qwen3.5 / Qwen3.6 GGUF (full chat path). Model: {}",
model_path.display(),
);
}
if arch == ARCH_QWEN35 || arch == ARCH_QWEN35MOE {
if is_qwen36_gguf(&gguf_peek) && !INVESTIGATION_ENV.qwen36_autoreg {
anyhow::bail!(
"Qwen3.6 GGUF detected (general.name contains 'qwen3.6'), but \
autoregressive forward-path support is opt-in. Set \
HF2Q_QWEN36_AUTOREG=1 to dispatch through the existing \
autoregressive Qwen3.5 path (correct at short prefill; long-prefill \
SOTA via chunk-scan kernel deferred to Wave 5b). Model: {}",
model_path.display(),
);
}
tracing::info!("Detected architecture '{}' → routing to Qwen3.5 path", arch);
return cmd_generate_qwen35(args, gguf_peek);
}
if arch == "deepseek4" {
tracing::info!("Detected architecture 'deepseek4' → routing to native path");
return deepseek4_cli::cmd_generate(args, gguf_peek);
}
if arch != "gemma4" {
anyhow::bail!(
"unsupported GGUF general.architecture={arch:?}; `hf2q generate` supports \
gemma4 and the explicitly dispatched Qwen families in this build. Model: {}",
model_path.display(),
);
}
} else {
anyhow::bail!(
"GGUF is missing required `general.architecture`; refusing to guess Gemma. \
Model: {}",
model_path.display(),
);
}
}
let stderr_is_tty = std::io::IsTerminal::is_terminal(&std::io::stderr());
let stdout_is_tty = std::io::IsTerminal::is_terminal(&std::io::stdout());
let load_opts = api::engine::LoadOptions {
model_path: model_path.clone(),
tokenizer_path: args.tokenizer.clone(),
config_path: args.config.clone(),
dwq_overlay_path: None,
kv_persist_dir: std::env::var("HF2Q_KV_PERSIST")
.ok()
.filter(|s| !s.is_empty())
.map(std::path::PathBuf::from),
};
let load_start = std::time::Instant::now();
let loaded =
api::engine::GemmaLoadedModel::load(&load_opts).context("GemmaLoadedModel::load")?;
let load_elapsed = load_start.elapsed();
let gguf = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF re-open (post-load, banner+prompt): {e}"))?;
let mut info = <api::engine::GemmaLoadedModel as load_info::LoadInfoBuilder>::build_load_info(
&loaded,
&gguf,
load_elapsed,
None,
false,
);
anyhow::ensure!(
args.image.is_none() || args.mmproj.is_some(),
"--image requires --mmproj <path>; pass the projector GGUF that pairs \
with this base model (e.g. mmproj-gemma4-f16.gguf alongside the main \
GGUF on the same HF repo)"
);
let loaded_mmproj: Option<api::state::LoadedMmproj> = if let Some(mmp_path) =
args.mmproj.as_ref()
{
anyhow::ensure!(
mmp_path.exists(),
"mmproj not found: {}",
mmp_path.display()
);
let mmp_gguf = mlx_native::gguf::GgufFile::open(mmp_path)
.map_err(|e| anyhow::anyhow!("mmproj GGUF header parse failed: {e}"))?;
let mmp_config = crate::inference::vision::mmproj::MmprojConfig::from_gguf(&mmp_gguf)
.map_err(|e| anyhow::anyhow!("mmproj GGUF config parse failed: {e}"))?;
let actual_names: Vec<&str> = mmp_gguf.tensor_names();
crate::inference::vision::mmproj::validate_tensor_set(&mmp_config, &actual_names)
.map_err(|e| anyhow::anyhow!("mmproj GGUF tensor-set validation: {e}"))?;
let arch_profile = crate::inference::vision::mmproj::detect_arch_profile_with_projector(
&mmp_config.projector,
&actual_names,
);
anyhow::ensure!(
arch_profile.is_supported(),
"mmproj arch profile is Unknown — neither Gemma 4 SigLIP nor \
classic CLIP nor Qwen3-VL SigLIP markers found. hf2q's ViT \
forward path cannot dispatch on this file."
);
let mmproj_sha256 = mmp_gguf
.metadata_string("hf2q.mmproj_sha256")
.map(|s| s.to_string());
info.vision_projector = Some(load_info::VisionProjector {
mmproj_path: mmp_path.clone(),
mmproj_sha256,
});
if args.image.is_some() {
let device = mlx_native::MlxDevice::new()
.map_err(|e| anyhow::anyhow!("MlxDevice for mmproj load: {e}"))?;
let weights = crate::inference::vision::mmproj_weights::LoadedMmprojWeights::load(
&mmp_gguf,
&mmp_config,
device,
)
.map_err(|e| anyhow::anyhow!("mmproj weight load: {e}"))?;
let model_id = mmp_path
.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| "mmproj".into());
tracing::info!(
path = %mmp_path.display(),
image_size = mmp_config.image_size,
patch_size = mmp_config.patch_size,
hidden = mmp_config.hidden_size,
layers = mmp_config.num_hidden_layers,
projector = mmp_config.projector.as_str(),
arch = arch_profile.as_str(),
tensors_loaded = weights.len(),
"Loaded mmproj GGUF header + tensor set + weights"
);
Some(api::state::LoadedMmproj {
gguf_path: mmp_path.clone(),
config: mmp_config,
arch: arch_profile,
weights: std::sync::Arc::new(weights),
model_id,
})
} else {
tracing::info!(
path = %mmp_path.display(),
image_size = mmp_config.image_size,
patch_size = mmp_config.patch_size,
arch = arch_profile.as_str(),
"Loaded mmproj GGUF header (no --image; weight load skipped)"
);
None
}
} else {
None
};
load_info::emit_tracing(&info);
let mut stdout = std::io::stdout();
load_info::print_banner(&info, &mut stdout, stdout_is_tty).context("print load banner")?;
let mut ctx = loaded.ctx;
let mut mlx_w = loaded.weights;
let tokenizer = loaded.tokenizer;
let prompt_text_raw = resolve_prompt(&args)?;
let prompt_text_with_image_marker = if args.image.is_some() {
if let Some(mmproj) = loaded_mmproj.as_ref() {
let family = mmproj.arch.vision_family();
let placeholder = family.placeholder_token_literal().ok_or_else(|| {
anyhow::anyhow!(
"mmproj arch profile {:?} has no placeholder token literal",
mmproj.arch
)
})?;
let (open, close) = family.marker_pair();
format!("{open}{placeholder}{close}\n{prompt_text_raw}")
} else {
prompt_text_raw
}
} else {
prompt_text_raw
};
let prompt_text = render_chat_template(
&gguf,
&args,
Some(&tokenizer),
&prompt_text_with_image_marker,
)?;
if let Some(dump_path) = INVESTIGATION_ENV.dump_rendered_prompt.as_deref() {
std::fs::write(dump_path, prompt_text.as_bytes())
.with_context(|| format!("HF2Q_DUMP_RENDERED_PROMPT: failed to write {dump_path}"))?;
eprintln!(
"HF2Q_DUMP_RENDERED_PROMPT: wrote {} bytes to {}",
prompt_text.len(),
dump_path
);
return Ok(());
}
let prompt_tokens = tokenize_rendered_prompt_llama_style(&gguf, &tokenizer, &prompt_text)?;
let mut soft_tokens_owned: Vec<api::engine::SoftTokenData> = Vec::new();
let prompt_tokens = if let Some(image_path) = args.image.as_ref() {
let mmproj = loaded_mmproj
.as_ref()
.expect("--mmproj checked above when --image is set");
let image_input =
crate::inference::vision::parse_image_url(image_path.to_string_lossy().as_ref())
.with_context(|| format!("--image: parse {}", image_path.display()))?;
let bytes = crate::inference::vision::load_image_bytes(&image_input)
.with_context(|| format!("--image: load {}", image_path.display()))?;
let preprocessed_input = match mmproj.arch {
crate::inference::vision::mmproj::ArchProfile::Gemma4Siglip => {
let cfg = &crate::inference::vision::preprocess::GEMMA4V_PREPROCESS_DEFAULT;
let pp = crate::inference::vision::preprocess::preprocess_gemma4v(&bytes, cfg)
.with_context(|| {
format!("--image: gemma4v preprocess {}", image_path.display())
})?;
let source_label = image_path
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("image")
.to_string();
crate::inference::vision::vit_gpu::VisionInput::Gemma4v(
crate::inference::vision::vit_gpu::Gemma4vPreprocessedImage {
patches: pp.patches,
pos_x: pp.pos_x,
pos_y: pp.pos_y,
n_x: pp.n_x,
n_y: pp.n_y,
source_label,
},
)
}
crate::inference::vision::mmproj::ArchProfile::ClipClassic => {
let preprocess_cfg = mmproj.config.preprocess_config();
let pixel_values =
crate::inference::vision::preprocess_rgb_chw(&bytes, &preprocess_cfg)
.with_context(|| {
format!("--image: clip preprocess {}", image_path.display())
})?;
let source_label = image_path
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("image")
.to_string();
crate::inference::vision::vit_gpu::VisionInput::Siglip49(
crate::inference::vision::PreprocessedImage {
pixel_values,
target_size: preprocess_cfg.target_size,
pixel_w: None,
pixel_h: None,
source_label,
},
)
}
crate::inference::vision::mmproj::ArchProfile::Qwen3VlSiglip => {
anyhow::bail!(
"--image with Qwen3-VL mmproj on `hf2q generate` is not \
supported: the Qwen3-VL text LM forward path lives behind \
iter-228b (cmd_generate routes via `qwen3vl_text_forward_pending` \
501 today). Use `hf2q serve --mmproj <qwen3vl-mmproj>` for \
Qwen3-VL vision; this CLI works for Gemma4 + classic-CLIP."
);
}
crate::inference::vision::mmproj::ArchProfile::Unknown => {
unreachable!("Unknown arch rejected at mmproj load above")
}
};
let pipeline_out = crate::inference::vision::pipeline::run_vit_forward(
std::slice::from_ref(&preprocessed_input),
mmproj,
mlx_w.hidden_size,
)
.context("--image: ViT forward")?;
let (expanded_tokens, soft_tokens, _image_token_positions) =
crate::inference::vision::pipeline::expand_image_placeholders(
&tokenizer,
&prompt_tokens,
&pipeline_out.embeddings,
pipeline_out.family,
pipeline_out.per_row_floats,
mlx_w.hidden_size,
)
.context("--image: soft-token expansion")?;
tracing::info!(
image = %image_path.display(),
arch = mmproj.arch.as_str(),
n_image_tokens = pipeline_out.total_image_tokens(),
forward_ms = pipeline_out.forward_ms,
prompt_tokens_pre = prompt_tokens.len(),
prompt_tokens_post = expanded_tokens.len(),
"vision pipeline complete; prompt expanded with soft tokens"
);
soft_tokens_owned = soft_tokens;
expanded_tokens
} else {
prompt_tokens
};
tracing::info!("Prompt: {} tokens", prompt_tokens.len());
if INVESTIGATION_ENV.dump_prompt_tokens {
eprintln!(
"HF2Q_DUMP_PROMPT_TOKENS: first10={:?} last10={:?} total={}",
&prompt_tokens[..prompt_tokens.len().min(10)],
&prompt_tokens[prompt_tokens.len().saturating_sub(10)..],
prompt_tokens.len()
);
eprintln!("HF2Q_DUMP_PROMPT_TOKENS: full={:?}", prompt_tokens);
}
let params = sampler_pure::SamplingParams {
temperature: args.temperature,
top_p: args.top_p,
top_k: args.top_k,
min_p: args.min_p,
repetition_penalty: args.repetition_penalty,
max_tokens: args.max_tokens,
};
use std::io::Write;
tracing::info!("Running mlx-native forward pass");
let eos_token_ids: Vec<u32> = vec![1, 106];
if let Some(()) = crate::serve::spec_decode_cli::try_dispatch_dflash_spec_decode(
&mut mlx_w,
&prompt_tokens,
args.max_tokens,
&eos_token_ids,
args.ignore_eos,
&tokenizer,
&mut ctx,
)? {
return Ok(());
}
if let Some(()) = crate::serve::spec_decode_cli::try_dispatch_gemma4_eagle3_spec_decode(
&mut mlx_w,
&prompt_tokens,
args.max_tokens,
&eos_token_ids,
args.ignore_eos,
&tokenizer,
&mut ctx,
)? {
return Ok(());
}
if let Some(()) = crate::serve::spec_decode_cli::try_dispatch_ngram_spec_decode(
&mut mlx_w,
&prompt_tokens,
args.max_tokens,
&eos_token_ids,
args.ignore_eos,
&tokenizer,
&mut ctx,
)? {
return Ok(());
}
let mut profiler = crate::inference::models::gemma4::ProfileAccumulator::new(2);
let kernel_profile_mode = INVESTIGATION_ENV.mlx_kernel_profile;
let use_batched = INVESTIGATION_ENV.batched_prefill;
if args.benchmark {
let mut regime_results: Vec<(usize, Vec<f64>, Vec<usize>)> = Vec::new();
for ®ime_target in BENCH_GEN_LENS.iter() {
let regime_cap = regime_target.min(args.max_tokens);
if regime_cap == 0 {
continue;
}
let mut decode_tps_runs: Vec<f64> = Vec::with_capacity(BENCH_NUM_RUNS);
let mut generated_per_run: Vec<usize> = Vec::with_capacity(BENCH_NUM_RUNS);
eprintln!(
"\n=== Bench regime: gen_target={} (cap={}) ===",
regime_target, regime_cap
);
for run_idx in 0..BENCH_NUM_RUNS {
let last_token = if !soft_tokens_owned.is_empty() {
let borrowed: Vec<crate::serve::forward_prefill::SoftTokenInjection<'_>> =
soft_tokens_owned
.iter()
.map(|s| crate::serve::forward_prefill::SoftTokenInjection {
range: s.range.clone(),
embeddings: &s.embeddings,
})
.collect();
mlx_w.forward_prefill_with_soft_tokens(
&prompt_tokens,
&borrowed,
regime_cap,
&mut ctx,
)?
} else if use_batched {
mlx_w.forward_prefill_batched(&prompt_tokens, regime_cap, 0, &mut ctx)?
} else {
mlx_w.forward_prefill(&prompt_tokens, regime_cap, &mut ctx)?
};
let mut all_tokens = prompt_tokens.to_vec();
let mut next_token = last_token;
all_tokens.push(next_token);
let mut decoded_tokens: Vec<u32> = vec![next_token];
let decode_start = std::time::Instant::now();
let mut generated = 1usize;
let mut p: Option<crate::inference::models::gemma4::TokenProfile> = None;
for _ in 1..regime_cap {
if !args.ignore_eos && eos_token_ids.contains(&next_token) {
break;
}
let pos = all_tokens.len() - 1;
next_token = mlx_w.forward_decode(next_token, pos, &mut ctx, &mut p)?;
all_tokens.push(next_token);
generated += 1;
decoded_tokens.push(next_token);
if detect_greedy_repetition_loop(&decoded_tokens).is_some() {
break;
}
}
let decode_elapsed = decode_start.elapsed();
let tps = if decode_elapsed.as_secs_f64() > 0.0 {
generated as f64 / decode_elapsed.as_secs_f64()
} else {
0.0
};
eprintln!(
" [gen={}] Run {}/{}: {} tokens in {:.2}s ({:.1} tok/s)",
regime_target,
run_idx + 1,
BENCH_NUM_RUNS,
generated,
decode_elapsed.as_secs_f64(),
tps,
);
decode_tps_runs.push(tps);
generated_per_run.push(generated);
}
regime_results.push((regime_target, decode_tps_runs, generated_per_run));
}
println!();
println!("=== Benchmark Results (multi-regime) ===");
let (chip, mem_gb) = detect_hardware_info();
let model_filename = model_path
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| "unknown".to_string());
println!("Hardware: {}, {} GB", chip, mem_gb);
println!("Model: {}", model_filename);
println!("Prompt tokens: {}", prompt_tokens.len());
println!("Runs per regime: {}", BENCH_NUM_RUNS);
println!();
println!(
"{:>10} {:>10} {:>10} {:>10}",
"regime", "median", "p95", "min/max"
);
for (target, tps_runs, _gen_per_run) in ®ime_results {
let med = median_f64(tps_runs);
let p95 = p95_f64(tps_runs);
let min = tps_runs.iter().copied().fold(f64::INFINITY, f64::min);
let max = tps_runs.iter().copied().fold(f64::NEG_INFINITY, f64::max);
println!(
" gen={:>4} {:>10.1} {:>10.1} {:>5.1}/{:.1}",
target, med, p95, min, max
);
}
if let Some((_, tps_runs, _)) = regime_results.last() {
println!();
println!(
"Decode tok/s: {:.1} (longest-regime median; full table above)",
median_f64(tps_runs)
);
println!("Median: {:.1} tok/s", median_f64(tps_runs));
println!("P95: {:.1} tok/s", p95_f64(tps_runs));
}
return Ok(());
}
let prefill_start = std::time::Instant::now();
let last_token = if !soft_tokens_owned.is_empty() {
let borrowed: Vec<crate::serve::forward_prefill::SoftTokenInjection<'_>> =
soft_tokens_owned
.iter()
.map(|s| crate::serve::forward_prefill::SoftTokenInjection {
range: s.range.clone(),
embeddings: &s.embeddings,
})
.collect();
mlx_w.forward_prefill_with_soft_tokens(
&prompt_tokens,
&borrowed,
args.max_tokens,
&mut ctx,
)?
} else if use_batched {
mlx_w.forward_prefill_batched(&prompt_tokens, args.max_tokens, 0, &mut ctx)?
} else {
mlx_w.forward_prefill(&prompt_tokens, args.max_tokens, &mut ctx)?
};
let prefill_elapsed = prefill_start.elapsed();
let prefill_n = prompt_tokens.len();
let prefill_ms = prefill_elapsed.as_secs_f64() * 1000.0;
let prefill_tok_s = if prefill_elapsed.as_secs_f64() > 0.0 {
prefill_n as f64 / prefill_elapsed.as_secs_f64()
} else {
0.0
};
header::print_header_prefill(
&mut stdout,
&header::HeaderInfoPrefill {
prefill_n,
prefill_ms,
prefill_tok_s,
},
stdout_is_tty,
)
.context("print header prefill")?;
let mut all_tokens = prompt_tokens.to_vec();
let mut next_token = last_token;
all_tokens.push(next_token);
let mut decoded_tokens: Vec<u32> = vec![next_token];
let mut printed_text = tokenizer.decode(&decoded_tokens, false).unwrap_or_default();
print!("{}", printed_text);
std::io::stdout().flush()?;
if std::env::var("HF2Q_DUMP_COUNTERS").ok().as_deref() == Some("1") {
mlx_native::reset_counters();
}
let decode_start = std::time::Instant::now();
let mut generated = 1usize;
let mut kernel_profiles: Vec<crate::inference::models::gemma4::KernelTypeProfile> = Vec::new();
let kernel_profile_warmup = 2usize;
let kernel_profile_measure = 3usize;
for _ in 1..params.max_tokens {
if !args.ignore_eos && eos_token_ids.contains(&next_token) {
break;
}
let pos = all_tokens.len() - 1;
let kernel_profile_break = if kernel_profile_mode {
let (tok, kp) = mlx_w.forward_decode_kernel_profile(next_token, pos, &mut ctx)?;
next_token = tok;
if generated > kernel_profile_warmup {
kernel_profiles.push(kp);
}
kernel_profiles.len() >= kernel_profile_measure
} else {
let mut p = profiler.start_token();
let argmax_token = mlx_w.forward_decode(next_token, pos, &mut ctx, &mut p)?;
next_token = if qwen35_generate_uses_sampling(&args) {
let mut logits: Vec<f32> = mlx_w
.logits_view()
.context("Gemma sampling: logits_view")?
.to_vec();
sample_qwen35_logits_for_generate(&mut logits, &args, &decoded_tokens)
} else {
argmax_token
};
profiler.finish_token(p);
false
};
all_tokens.push(next_token);
generated += 1;
decoded_tokens.push(next_token);
let new_full = tokenizer.decode(&decoded_tokens, false).unwrap_or_default();
if new_full.len() > printed_text.len() && new_full.starts_with(&printed_text) {
print!("{}", &new_full[printed_text.len()..]);
std::io::stdout().flush()?;
}
printed_text = new_full;
if kernel_profile_break {
break;
}
if !args.ignore_eos && !qwen35_generate_uses_sampling(&args) {
let detect = detect_greedy_repetition_loop_with_text(&decoded_tokens, |cycle| {
tokenizer
.decode(cycle, false)
.unwrap_or_default()
});
if let Some((ngram, repeats)) = detect {
tracing::info!(
"Gemma decode: greedy n-gram repetition detected (last {} tokens \
repeated {} times); stopping. Pass --temperature 0.8 or \
--repetition-penalty 1.1 (or use the chat-completion API) to \
opt into sampling.",
ngram,
repeats
);
eprintln!(
"\n[hf2q] Gemma greedy decode entered a {}-token repetition loop \
— stopping. Pass --temperature 0.8 or --repetition-penalty 1.1 \
to escape via sampling.",
ngram
);
break;
}
}
}
let decode_elapsed = decode_start.elapsed();
let tok_per_sec = generated as f64 / decode_elapsed.as_secs_f64();
let (td, tr) = if stderr_is_tty {
("\x1b[2m", "\x1b[0m")
} else {
("", "")
};
eprintln!(
"\n\n{td}--- mlx-native: [ Prompt: {:.1} t/s | Generation: {:.1} t/s ] \
({} gen tokens in {:.2}s) ---{tr}",
prefill_tok_s,
tok_per_sec,
generated,
decode_elapsed.as_secs_f64(),
);
profiler.print_summary();
if kernel_profile_mode && !kernel_profiles.is_empty() {
crate::inference::models::gemma4::MlxModelWeights::print_kernel_profile_report(
&kernel_profiles,
);
}
if std::env::var("HF2Q_DUMP_COUNTERS").ok().as_deref() == Some("1") {
let dispatches = mlx_native::dispatch_count();
let syncs = mlx_native::sync_count();
let cmd_bufs = mlx_native::cmd_buf_count();
let barriers = mlx_native::barrier_count();
let prompt_n = prompt_tokens.len() as u64;
let decode_n = generated as u64;
let dispatches_per_decode_tok = if decode_n > 0 {
dispatches as f64 / decode_n as f64
} else {
0.0
};
let syncs_per_decode_tok = if decode_n > 0 {
syncs as f64 / decode_n as f64
} else {
0.0
};
let cb_per_decode_tok = if decode_n > 0 {
cmd_bufs as f64 / decode_n as f64
} else {
0.0
};
let barriers_per_decode_tok = if decode_n > 0 {
barriers as f64 / decode_n as f64
} else {
0.0
};
eprintln!(
"[MLX_COUNTERS] dispatches={} syncs={} cmd_bufs={} barriers={} \
prompt_tokens={} decode_tokens={} \
dispatches/decode_tok={:.2} syncs/decode_tok={:.2} \
cmd_bufs/decode_tok={:.4} barriers/decode_tok={:.2}",
dispatches,
syncs,
cmd_bufs,
barriers,
prompt_n,
decode_n,
dispatches_per_decode_tok,
syncs_per_decode_tok,
cb_per_decode_tok,
barriers_per_decode_tok,
);
let buckets = mlx_native::pipeline_dispatch_buckets();
if !buckets.is_empty() {
let total: u64 = buckets.iter().map(|(_, c)| *c).sum();
eprintln!(
"[MLX_DISP_BUCKET] Per-pipeline breakdown ({} unique pipelines, total={}):",
buckets.len(),
total,
);
for (label, count) in &buckets {
let pct = if total > 0 {
100.0 * (*count as f64) / (total as f64)
} else {
0.0
};
eprintln!(
"[MLX_DISP_BUCKET] {:>8} ({:5.2}%) {}",
count, pct, label,
);
}
}
}
Ok(())
}
const QWEN35_PREFILL_SWEEP_TEXT: &str =
"The benchmark prompt describes a careful engineering investigation. \
It repeats neutral facts about measuring model speed, preserving token \
probabilities, checking coherence, and comparing current code against \
peer implementations. ";
fn qwen35_sweep_token_count(tokenizer: &tokenizers::Tokenizer, text: &str) -> Result<usize> {
let encoding = tokenizer
.encode(text, false)
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep tokenization failed: {e}"))?;
Ok(encoding.get_ids().len())
}
fn qwen35_sweep_prompt(
tokenizer: &tokenizers::Tokenizer,
target_tokens: usize,
) -> Result<(String, Vec<u32>)> {
let repeated = QWEN35_PREFILL_SWEEP_TEXT.repeat((target_tokens / 20) + 500);
let ids = tokenizer
.encode(repeated.as_str(), false)
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep base tokenization failed: {e}"))?
.get_ids()
.to_vec();
anyhow::ensure!(
ids.len() >= target_tokens,
"qwen35 prefill sweep base prompt produced only {} tokens for target {}",
ids.len(),
target_tokens,
);
let mut best_text = tokenizer
.decode(&ids[..target_tokens], false)
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep decode failed: {e}"))?;
let mut best_n = qwen35_sweep_token_count(tokenizer, &best_text)?;
if best_n == target_tokens {
let enc = tokenizer
.encode(best_text.as_str(), false)
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep final tokenization failed: {e}"))?;
return Ok((best_text, enc.get_ids().to_vec()));
}
let lo = target_tokens.saturating_sub(256).max(1);
let hi = (target_tokens + 256).min(ids.len());
for n_ids in lo..=hi {
let text = tokenizer
.decode(&ids[..n_ids], false)
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep decode failed: {e}"))?;
let n = qwen35_sweep_token_count(tokenizer, &text)?;
if n == target_tokens {
let enc = tokenizer.encode(text.as_str(), false).map_err(|e| {
anyhow::anyhow!("qwen35 prefill sweep final tokenization failed: {e}")
})?;
return Ok((text, enc.get_ids().to_vec()));
}
if n.abs_diff(target_tokens) < best_n.abs_diff(target_tokens) {
best_text = text;
best_n = n;
}
}
let enc = tokenizer
.encode(best_text.as_str(), false)
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep final tokenization failed: {e}"))?;
Ok((best_text, enc.get_ids().to_vec()))
}
fn qwen35_positions(seq_len: usize) -> Vec<i32> {
let mut flat = vec![0i32; 4 * seq_len];
for axis in 0..4 {
for t in 0..seq_len {
flat[axis * seq_len + t] = t as i32;
}
}
flat
}
fn maybe_run_qwen35_prefill_sweep(
model: &crate::inference::models::qwen35::model::Qwen35Model,
tokenizer: &tokenizers::Tokenizer,
) -> Result<bool> {
let Ok(lengths_raw) = std::env::var("HF2Q_QWEN35_PREFILL_SWEEP") else {
return Ok(false);
};
let lengths: Vec<usize> = lengths_raw
.split(',')
.filter(|s| !s.trim().is_empty())
.map(|s| {
s.trim().parse::<usize>().with_context(|| {
format!("HF2Q_QWEN35_PREFILL_SWEEP contains non-integer length {s:?}")
})
})
.collect::<Result<_>>()?;
anyhow::ensure!(
!lengths.is_empty(),
"HF2Q_QWEN35_PREFILL_SWEEP must contain at least one length"
);
let trials = std::env::var("HF2Q_QWEN35_PREFILL_SWEEP_TRIALS")
.ok()
.map(|s| {
s.parse::<usize>()
.context("parse HF2Q_QWEN35_PREFILL_SWEEP_TRIALS")
})
.transpose()?
.unwrap_or(3)
.max(1);
let warmups = std::env::var("HF2Q_QWEN35_PREFILL_SWEEP_WARMUPS")
.ok()
.map(|s| {
s.parse::<usize>()
.context("parse HF2Q_QWEN35_PREFILL_SWEEP_WARMUPS")
})
.transpose()?
.unwrap_or(1);
use crate::inference::models::qwen35::io_heads::greedy_argmax_last_token;
use crate::inference::models::qwen35::kv_cache::HybridKvCache;
use crate::serve::multi_seq_kv::SlotId;
use mlx_native::MlxDevice;
fn top_n_indices(values: &[f32], n: usize) -> Vec<usize> {
let mut indexed: Vec<(usize, f32)> = values.iter().copied().enumerate().collect();
indexed.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
indexed.into_iter().take(n).map(|(idx, _)| idx).collect()
}
fn logit_row_metrics(a: &[f32], b: &[f32]) -> (f32, f64, usize) {
let mut max_abs = 0.0f32;
let mut dot = 0.0f64;
let mut a2 = 0.0f64;
let mut b2 = 0.0f64;
for (&av, &bv) in a.iter().zip(b) {
max_abs = max_abs.max((av - bv).abs());
let af = av as f64;
let bf = bv as f64;
dot += af * bf;
a2 += af * af;
b2 += bf * bf;
}
let cosine = if a2 > 0.0 && b2 > 0.0 {
dot / (a2.sqrt() * b2.sqrt())
} else {
f64::NAN
};
let top_a = top_n_indices(a, 10);
let top_b = top_n_indices(b, 10);
let overlap = top_a.iter().filter(|idx| top_b.contains(idx)).count();
(max_abs, cosine, overlap)
}
let device = MlxDevice::new()
.map_err(|e| anyhow::anyhow!("qwen35 prefill sweep MlxDevice::new: {e}"))?;
println!(
"{{\"event\":\"qwen35_prefill_sweep_start\",\"lengths\":{:?},\"warmups\":{},\"trials\":{}}}",
lengths, warmups, trials
);
for target in lengths {
let (_prompt, prompt_tokens) = qwen35_sweep_prompt(tokenizer, target)?;
let prompt_len = prompt_tokens.len();
let positions = qwen35_positions(prompt_len);
let max_seq = (prompt_len + 65)
.max(128)
.min(model.cfg.max_position_embeddings as usize);
for iteration in 0..(warmups + trials) {
let phase = if iteration < warmups {
"warmup"
} else {
"measure"
};
let trial = if phase == "warmup" {
iteration
} else {
iteration - warmups
};
let mut kv_cache = HybridKvCache::new(&model.cfg, &device, max_seq as u32, 1)
.context("qwen35 prefill sweep HybridKvCache::new")?;
let t0 = std::time::Instant::now();
let full_logits =
std::env::var("HF2Q_QWEN35_PREFILL_SWEEP_FULL_LOGITS").as_deref() == Ok("1");
let compare_full_last =
std::env::var("HF2Q_QWEN35_PREFILL_SWEEP_COMPARE_FULL_LAST").as_deref() == Ok("1");
let logits = if compare_full_last {
let mut full_kv = HybridKvCache::new(&model.cfg, &device, max_seq as u32, 1)
.context("qwen35 prefill sweep compare full HybridKvCache::new")?;
let full = model
.forward_gpu(&prompt_tokens, &positions, &mut full_kv, SlotId(0))
.context("qwen35 prefill sweep compare forward_gpu")?;
let mut last_kv = HybridKvCache::new(&model.cfg, &device, max_seq as u32, 1)
.context("qwen35 prefill sweep compare last HybridKvCache::new")?;
let last = model
.forward_gpu_last_logits(&prompt_tokens, &positions, &mut last_kv, SlotId(0))
.context("qwen35 prefill sweep compare forward_gpu_last_logits")?;
let vocab = model.cfg.vocab_size as usize;
anyhow::ensure!(
full.len() == prompt_len * vocab,
"qwen35 prefill sweep compare full logits.len()={} != expected {}",
full.len(),
prompt_len * vocab,
);
anyhow::ensure!(
last.len() == vocab,
"qwen35 prefill sweep compare last logits.len()={} != expected {}",
last.len(),
vocab,
);
let full_last = &full[full.len() - vocab..];
let full_token = greedy_argmax_last_token(full_last, model.cfg.vocab_size);
let last_token = greedy_argmax_last_token(&last, model.cfg.vocab_size);
let (max_abs, cosine, top10_overlap) = logit_row_metrics(full_last, &last);
println!(
"{{\"event\":\"qwen35_prefill_sweep_compare\",\"target_tokens\":{},\"actual_tokens\":{},\"phase\":\"{}\",\"iteration\":{},\"trial\":{},\"full_token\":{},\"last_token\":{},\"max_abs\":{:.9},\"cosine\":{:.12},\"top10_overlap\":{}}}",
target,
prompt_len,
phase,
iteration,
trial,
full_token,
last_token,
max_abs,
cosine,
top10_overlap,
);
last
} else if full_logits {
model
.forward_gpu(&prompt_tokens, &positions, &mut kv_cache, SlotId(0))
.context("qwen35 prefill sweep forward_gpu")?
} else {
model
.forward_gpu_last_logits(&prompt_tokens, &positions, &mut kv_cache, SlotId(0))
.context("qwen35 prefill sweep forward_gpu_last_logits")?
};
let elapsed = t0.elapsed();
let vocab = model.cfg.vocab_size as usize;
let expected_logits = if full_logits && !compare_full_last {
prompt_len * vocab
} else {
vocab
};
anyhow::ensure!(
logits.len() == expected_logits,
"qwen35 prefill sweep logits.len()={} != expected {} (full_logits={})",
logits.len(),
expected_logits,
full_logits,
);
let last_logits = &logits[logits.len() - vocab..];
let first_token = greedy_argmax_last_token(last_logits, model.cfg.vocab_size);
let ms = elapsed.as_secs_f64() * 1000.0;
let tps = prompt_len as f64 / elapsed.as_secs_f64();
println!(
"{{\"event\":\"qwen35_prefill_sweep\",\"target_tokens\":{},\"actual_tokens\":{},\"phase\":\"{}\",\"iteration\":{},\"trial\":{},\"prefill_ms\":{:.3},\"prefill_tps\":{:.3},\"first_token\":{},\"output_head\":\"{}\"}}",
target,
prompt_len,
phase,
iteration,
trial,
ms,
tps,
first_token,
if full_logits { "all" } else { "last" },
);
}
}
println!("{{\"event\":\"qwen35_prefill_sweep_end\"}}");
Ok(true)
}
fn detect_greedy_repetition_loop(decoded_tokens: &[u32]) -> Option<(usize, usize)> {
detect_greedy_repetition_loop_with_text(decoded_tokens, |_| String::new())
}
fn count_consecutive_cycle_copies(
decoded_tokens: &[u32],
key: &[u32],
min_confirmed: usize,
) -> usize {
let ngram = key.len();
let n = decoded_tokens.len();
let mut actual = min_confirmed;
let mut probe_start = match n.checked_sub((min_confirmed + 1) * ngram) {
Some(p) => p,
None => return actual,
};
loop {
if &decoded_tokens[probe_start..probe_start + ngram] != key {
break;
}
actual += 1;
probe_start = match probe_start.checked_sub(ngram) {
Some(p) => p,
None => break,
};
}
actual
}
fn detect_greedy_repetition_loop_with_text<F>(
decoded_tokens: &[u32],
mut decode_cycle: F,
) -> Option<(usize, usize)>
where
F: FnMut(&[u32]) -> String,
{
const MIN_NGRAM: usize = 2;
const MAX_NGRAM: usize = 64;
const MIN_OCCURRENCES_CONTENT: usize = 3;
const MIN_OCCURRENCES_STRUCTURAL: usize = 8;
const STRUCTURAL_ALPHA_THRESHOLD: f32 = 0.30;
let n = decoded_tokens.len();
for ngram in MIN_NGRAM..=MAX_NGRAM {
let need_content = ngram * MIN_OCCURRENCES_CONTENT;
if n < need_content {
continue;
}
let tail = &decoded_tokens[n - need_content..];
let key = &tail[need_content - ngram..];
let mut all_match = true;
for i in 0..MIN_OCCURRENCES_CONTENT {
let start = i * ngram;
if &tail[start..start + ngram] != key {
all_match = false;
break;
}
}
if !all_match {
continue;
}
let cycle_text = decode_cycle(key);
let total_chars = cycle_text.chars().count() as f32;
if total_chars == 0.0 {
let actual =
count_consecutive_cycle_copies(decoded_tokens, key, MIN_OCCURRENCES_CONTENT);
return Some((ngram, actual));
}
let alpha = cycle_text.chars().filter(|c| c.is_alphabetic()).count() as f32;
let alpha_ratio = alpha / total_chars;
if alpha_ratio >= STRUCTURAL_ALPHA_THRESHOLD {
let actual =
count_consecutive_cycle_copies(decoded_tokens, key, MIN_OCCURRENCES_CONTENT);
return Some((ngram, actual));
}
let need_structural = ngram * MIN_OCCURRENCES_STRUCTURAL;
if n < need_structural {
continue;
}
let tail2 = &decoded_tokens[n - need_structural..];
let mut all_match2 = true;
for i in 0..MIN_OCCURRENCES_STRUCTURAL {
let start = i * ngram;
if &tail2[start..start + ngram] != key {
all_match2 = false;
break;
}
}
if all_match2 {
let actual =
count_consecutive_cycle_copies(decoded_tokens, key, MIN_OCCURRENCES_STRUCTURAL);
return Some((ngram, actual));
}
}
None
}
const SPECIAL_TOKEN_STOPS: &[&str] = &[
"<|im_start|>user",
"<|im_start|>system",
"<|endoftext|>",
"<|end|>",
];
#[allow(dead_code)] fn find_special_token_stop(generated_text: &str) -> Option<&'static str> {
find_special_token_stop_pos(generated_text).map(|(_, marker)| marker)
}
fn find_special_token_stop_pos(generated_text: &str) -> Option<(usize, &'static str)> {
SPECIAL_TOKEN_STOPS
.iter()
.copied()
.filter_map(|m| generated_text.find(m).map(|pos| (pos, m)))
.min_by_key(|(pos, _)| *pos)
}
fn qwen35_generate_uses_sampling(args: &cli::GenerateArgs) -> bool {
args.temperature > crate::serve::sampler_pure::SAMPLING_EPS || args.repetition_penalty != 1.0
}
fn sample_qwen35_logits_for_generate(
logits: &mut [f32],
args: &cli::GenerateArgs,
previous_tokens: &[u32],
) -> u32 {
let params = crate::serve::sampler_pure::SamplingParams {
temperature: args.temperature,
top_p: args.top_p,
top_k: args.top_k,
min_p: args.min_p,
repetition_penalty: args.repetition_penalty,
max_tokens: args.max_tokens,
};
crate::serve::sampler_pure::sample_token(logits, ¶ms, previous_tokens)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum DecodeStopReason {
EosTokenId(u32),
SpecialTokenLeak(&'static str),
RepetitionLoop { ngram: usize, repeats: usize },
MaxSeq,
MaxTokensReached,
}
#[derive(Debug)]
pub(crate) struct DecodeLoopOutcome {
pub generated: usize,
pub stop_reason: DecodeStopReason,
}
#[allow(dead_code)]
pub(crate) struct DecodeStepEvent<'a> {
pub step: usize,
pub token: u32,
pub cumulative_text: &'a str,
pub delta: &'a str,
}
pub(crate) fn run_decode_loop<F, P, E>(
initial_token: u32,
prompt_len: usize,
max_seq: usize,
max_tokens: usize,
eos_token_ids: &[u32],
mut step_model: F,
mut decode_text: P,
mut on_event: E,
) -> Result<DecodeLoopOutcome>
where
F: FnMut(u32, i32, &[u32]) -> Result<u32>,
P: FnMut(&[u32]) -> String,
E: FnMut(DecodeStepEvent<'_>),
{
let mut decoded_tokens: Vec<u32> = vec![initial_token];
let mut emitted_text: String = String::new();
let mut next_token = initial_token;
let mut generated = 1usize;
let seed_full = decode_text(&decoded_tokens);
on_event(DecodeStepEvent {
step: 0,
token: initial_token,
cumulative_text: &seed_full,
delta: &seed_full,
});
emitted_text.push_str(&seed_full);
for step in 1..max_tokens {
if eos_token_ids.contains(&next_token) {
return Ok(DecodeLoopOutcome {
generated,
stop_reason: DecodeStopReason::EosTokenId(next_token),
});
}
let pos = (prompt_len + step - 1) as i32;
if (pos as usize) >= max_seq {
return Ok(DecodeLoopOutcome {
generated,
stop_reason: DecodeStopReason::MaxSeq,
});
}
next_token = step_model(next_token, pos, &decoded_tokens)?;
generated += 1;
decoded_tokens.push(next_token);
let new_full = decode_text(&decoded_tokens);
let delta = if new_full.len() > emitted_text.len() && new_full.starts_with(&emitted_text) {
&new_full[emitted_text.len()..]
} else {
""
};
on_event(DecodeStepEvent {
step,
token: next_token,
cumulative_text: &new_full,
delta,
});
if let Some(marker) = find_special_token_stop(&new_full) {
return Ok(DecodeLoopOutcome {
generated,
stop_reason: DecodeStopReason::SpecialTokenLeak(marker),
});
}
if let Some((ngram, repeats)) =
detect_greedy_repetition_loop_with_text(&decoded_tokens, |cycle| decode_text(cycle))
{
return Ok(DecodeLoopOutcome {
generated,
stop_reason: DecodeStopReason::RepetitionLoop { ngram, repeats },
});
}
emitted_text = new_full;
}
Ok(DecodeLoopOutcome {
generated,
stop_reason: DecodeStopReason::MaxTokensReached,
})
}
struct DisplayScrubber {
buf: String,
emitted: usize,
}
const SCRUB_TAGS: &[&str] = &[
"<|im_end|>",
"<|im_start|>assistant",
"<|im_start|>user",
"<|im_start|>system",
"<|endoftext|>",
"<|end|>",
"</think>",
"<think>",
];
impl DisplayScrubber {
fn new() -> Self {
Self {
buf: String::new(),
emitted: 0,
}
}
fn push(&mut self, delta: &str) -> String {
if delta.is_empty() {
return String::new();
}
self.buf.push_str(delta);
let trail_hold = max_trailing_tag_prefix_len(&self.buf);
let safe_to = self.buf.len().saturating_sub(trail_hold);
if safe_to <= self.emitted {
return String::new();
}
let raw_segment = &self.buf[self.emitted..safe_to];
let mut out = raw_segment.to_string();
for tag in SCRUB_TAGS {
if out.contains(tag) {
out = out.replace(tag, "");
}
}
self.emitted = safe_to;
out
}
fn flush(&self) -> String {
if self.emitted >= self.buf.len() {
return String::new();
}
let mut out = self.buf[self.emitted..].to_string();
for tag in SCRUB_TAGS {
if out.contains(tag) {
out = out.replace(tag, "");
}
}
out
}
}
fn max_trailing_tag_prefix_len(text: &str) -> usize {
let mut max = 0;
for tag in SCRUB_TAGS {
for len in 1..tag.len() {
let prefix = &tag[..len];
if text.ends_with(prefix) && len > max {
max = len;
}
}
}
max
}
fn cmd_generate_qwen35(args: cli::GenerateArgs, gguf: mlx_native::gguf::GgufFile) -> Result<()> {
use crate::inference::models::qwen35::io_heads::greedy_argmax_last_token;
use crate::inference::models::qwen35::kv_cache::HybridKvCache;
use crate::serve::api::engine_qwen35::Qwen35LoadedModel;
use crate::serve::multi_seq_kv::SlotId;
use mlx_native::MlxDevice;
use std::io::Write;
let model_path = &args.model;
let stdout_is_tty = std::io::IsTerminal::is_terminal(&std::io::stdout());
if INVESTIGATION_ENV.dump_rendered_prompt.is_some()
|| std::env::var("HF2Q_DEBUG_TOKENIZE_ONLY").as_deref() == Ok("1")
{
let tokenizer =
crate::inference::models::qwen35::tokenizer::build_tokenizer_from_gguf(&gguf)
.context("Qwen35 tokenizer build for prompt diagnostics")?;
let prompt_text_raw = resolve_prompt(&args)?;
let prompt_text = render_chat_template(&gguf, &args, Some(&tokenizer), &prompt_text_raw)?;
if let Some(dump_path) = INVESTIGATION_ENV.dump_rendered_prompt.as_deref() {
std::fs::write(dump_path, prompt_text.as_bytes()).with_context(|| {
format!("HF2Q_DUMP_RENDERED_PROMPT: failed to write {dump_path}")
})?;
eprintln!(
"HF2Q_DUMP_RENDERED_PROMPT: wrote {} bytes to {}",
prompt_text.len(),
dump_path
);
}
if std::env::var("HF2Q_DEBUG_TOKENIZE_ONLY").as_deref() == Ok("1") {
let prompt_tokens =
tokenize_rendered_prompt_llama_style(&gguf, &tokenizer, &prompt_text)?;
let id_str: Vec<String> = prompt_tokens.iter().map(|i| i.to_string()).collect();
println!("TOKENIZE_DEBUG_IDS: {}", id_str.join(" "));
}
return Ok(());
}
let load_opts = api::engine::LoadOptions {
model_path: model_path.clone(),
tokenizer_path: args.tokenizer.clone(),
config_path: args.config.clone(),
dwq_overlay_path: None,
kv_persist_dir: std::env::var("HF2Q_KV_PERSIST")
.ok()
.filter(|s| !s.is_empty())
.map(std::path::PathBuf::from),
};
let load_start = std::time::Instant::now();
let loaded = Qwen35LoadedModel::load(&load_opts).context("Qwen35LoadedModel::load")?;
let load_elapsed = load_start.elapsed();
let info = <Qwen35LoadedModel as load_info::LoadInfoBuilder>::build_load_info(
&loaded,
&gguf,
load_elapsed,
None,
false,
);
load_info::emit_tracing(&info);
let mut stdout = std::io::stdout();
load_info::print_banner(&info, &mut stdout, stdout_is_tty).context("print load banner")?;
let model = loaded.model;
let tokenizer = loaded.tokenizer;
let eos_token_ids: Vec<u32> = if loaded.eos_token_ids.is_empty() {
vec![151_645]
} else {
loaded.eos_token_ids.clone()
};
if maybe_run_qwen35_prefill_sweep(&model, &tokenizer)? {
return Ok(());
}
let prompt_text_raw = resolve_prompt(&args)?;
let prompt_text = render_chat_template(&gguf, &args, Some(&tokenizer), &prompt_text_raw)?;
if let Some(dump_path) = INVESTIGATION_ENV.dump_rendered_prompt.as_deref() {
std::fs::write(dump_path, prompt_text.as_bytes())
.with_context(|| format!("HF2Q_DUMP_RENDERED_PROMPT: failed to write {dump_path}"))?;
eprintln!(
"HF2Q_DUMP_RENDERED_PROMPT: wrote {} bytes to {}",
prompt_text.len(),
dump_path
);
return Ok(());
}
let prompt_tokens = tokenize_rendered_prompt_llama_style(&gguf, &tokenizer, &prompt_text)?;
let prompt_len = prompt_tokens.len();
tracing::info!("Qwen3.5: {} prompt tokens", prompt_len);
if std::env::var("HF2Q_DEBUG_TOKENIZE_ONLY").as_deref() == Ok("1") {
let id_str: Vec<String> = prompt_tokens.iter().map(|i| i.to_string()).collect();
println!("TOKENIZE_DEBUG_IDS: {}", id_str.join(" "));
return Ok(());
}
let max_seq = (prompt_len + args.max_tokens + 64)
.max(128)
.min(model.cfg.max_position_embeddings as usize);
let sample_logits = qwen35_generate_uses_sampling(&args);
let mut model = model;
if let Some(()) = crate::serve::spec_decode_cli::try_dispatch_qwen35_dflash_spec_decode(
&mut model,
&prompt_tokens,
args.max_tokens,
&eos_token_ids,
args.ignore_eos,
&tokenizer,
)? {
return Ok(());
}
if let Some(()) = crate::serve::spec_decode_cli::try_dispatch_qwen35_eagle3_spec_decode(
&mut model,
&prompt_tokens,
args.max_tokens,
&eos_token_ids,
args.ignore_eos,
&tokenizer,
)? {
return Ok(());
}
let spec_env = std::env::var("HF2Q_SPEC_DECODE").ok();
let mut use_spec_decode = match spec_env.as_deref() {
Some("0") => false,
Some("1") => true,
_ => !sample_logits && (args.speculative || model.mtp.is_some()),
};
if sample_logits && args.speculative {
tracing::warn!(
"Speculative decoding requested with sampling parameters; using sampler path"
);
}
if use_spec_decode && model.mtp.is_none() {
tracing::warn!(
"Speculative decoding requested but this GGUF has no MTP weights; using greedy decode"
);
use_spec_decode = false;
}
if use_spec_decode {
use crate::inference::models::qwen35::spec_decode::SpecDecode;
tracing::info!("Qwen3.5 speculative decode enabled");
model
.ensure_gpu_cache_primed()
.context("Qwen35Model::ensure_gpu_cache_primed (P19 H12 spec-decode warmup)")?;
if args.benchmark {
let mut prefill_tps_runs: Vec<f64> = Vec::with_capacity(BENCH_NUM_RUNS);
let mut decode_tps_runs: Vec<f64> = Vec::with_capacity(BENCH_NUM_RUNS);
let mut generated_per_run: Vec<usize> = Vec::with_capacity(BENCH_NUM_RUNS);
let mut accept_pct_runs: Vec<f64> = Vec::with_capacity(BENCH_NUM_RUNS);
let bench_sampler = crate::inference::models::qwen35::spec_decode::SpecSampler::new(
args.temperature as f32,
0,
);
for run_idx in 0..BENCH_NUM_RUNS {
let result = SpecDecode::run_with_sampler_eos_set(
&model,
&prompt_tokens,
args.max_tokens,
eos_token_ids.clone(),
max_seq as u32,
bench_sampler,
)
.context("qwen35 SpecDecode::run_with_sampler_eos_set (benchmark)")?;
let prefill_tps = if result.stats.prefill_elapsed.as_secs_f64() > 0.0 {
prompt_len as f64 / result.stats.prefill_elapsed.as_secs_f64()
} else {
0.0
};
let gen = result.tokens.len();
let decode_tps = if result.stats.decode_elapsed.as_secs_f64() > 0.0 {
gen as f64 / result.stats.decode_elapsed.as_secs_f64()
} else {
0.0
};
let accept_pct = result.stats.acceptance_rate_pct();
eprintln!(
" Run {}/{}: prefill {} tok in {:.2}s ({:.1} tok/s); decode {} tok in {:.2}s ({:.1} tok/s; accept {:.1}%)",
run_idx + 1,
BENCH_NUM_RUNS,
prompt_len,
result.stats.prefill_elapsed.as_secs_f64(),
prefill_tps,
gen,
result.stats.decode_elapsed.as_secs_f64(),
decode_tps,
accept_pct,
);
prefill_tps_runs.push(prefill_tps);
decode_tps_runs.push(decode_tps);
generated_per_run.push(gen);
accept_pct_runs.push(accept_pct);
}
let median_accept = median_f64(&accept_pct_runs);
print_benchmark_summary(
model_path,
prompt_len,
&generated_per_run,
Some(&prefill_tps_runs),
&decode_tps_runs,
&[("Spec accept %".to_string(), format!("{:.1}", median_accept))],
);
return Ok(());
}
let sampler = crate::inference::models::qwen35::spec_decode::SpecSampler::new(
args.temperature as f32,
0,
);
let effective_eos_for_spec: Vec<u32> = if args.ignore_eos {
Vec::new()
} else {
eos_token_ids.clone()
};
let result = SpecDecode::run_with_sampler_eos_set(
&model,
&prompt_tokens,
args.max_tokens,
effective_eos_for_spec,
max_seq as u32,
sampler,
)
.context("qwen35 SpecDecode::run_with_sampler_eos_set")?;
let prefill_tok_s = if result.stats.prefill_elapsed.as_secs_f64() > 0.0 {
prompt_len as f64 / result.stats.prefill_elapsed.as_secs_f64()
} else {
0.0
};
header::print_header_prefill(
&mut stdout,
&header::HeaderInfoPrefill {
prefill_n: prompt_len,
prefill_ms: result.stats.prefill_elapsed.as_secs_f64() * 1000.0,
prefill_tok_s,
},
stdout_is_tty,
)
.context("print header prefill")?;
let decoded = tokenizer.decode(&result.tokens, false).unwrap_or_default();
print!("{}", decoded);
stdout.flush()?;
let generated = result.tokens.len();
let tok_per_sec = if result.stats.decode_elapsed.as_secs_f64() > 0.0 {
generated as f64 / result.stats.decode_elapsed.as_secs_f64()
} else {
0.0
};
let (td, tr) = if std::io::IsTerminal::is_terminal(&std::io::stderr()) {
("\x1b[2m", "\x1b[0m")
} else {
("", "")
};
eprintln!(
"\n\n{td}--- mlx-native (qwen35 spec): {} tokens in {:.2}s ({:.1} tok/s, accept {:.1}%) ---{tr}",
generated,
result.stats.decode_elapsed.as_secs_f64(),
tok_per_sec,
result.stats.acceptance_rate_pct(),
);
return Ok(());
}
let device = MlxDevice::new().map_err(|e| anyhow::anyhow!("MlxDevice::new: {e}"))?;
let mut kv_cache =
HybridKvCache::new(&model.cfg, &device, max_seq as u32, 1).context("HybridKvCache::new")?;
tracing::info!(
"Qwen3.5 KV cache allocated: max_seq={}, {} MB",
max_seq,
kv_cache.total_bytes() / (1024 * 1024)
);
let prefill_positions: Vec<i32> = {
let mut flat = vec![0i32; 4 * prompt_len];
for axis in 0..4 {
for t in 0..prompt_len {
flat[axis * prompt_len + t] = t as i32;
}
}
flat
};
let warmup_start = std::time::Instant::now();
model
.ensure_gpu_cache_primed()
.context("Qwen35Model::ensure_gpu_cache_primed (P19 H12 warmup)")?;
tracing::info!(
"Qwen3.5 GPU warmup (P19 H12): {:.2}s",
warmup_start.elapsed().as_secs_f64()
);
if args.benchmark {
let mut prefill_tps_runs: Vec<f64> = Vec::with_capacity(BENCH_NUM_RUNS);
let mut decode_tps_runs: Vec<f64> = Vec::with_capacity(BENCH_NUM_RUNS);
let mut generated_per_run: Vec<usize> = Vec::with_capacity(BENCH_NUM_RUNS);
for run_idx in 0..BENCH_NUM_RUNS {
kv_cache.reset_all_buffers();
let prefill_start = std::time::Instant::now();
let prefill_logits = model
.forward_gpu_last_logits(
&prompt_tokens,
&prefill_positions,
&mut kv_cache,
SlotId(0),
)
.context("Qwen35Model::forward_gpu_last_logits (benchmark prefill)")?;
let prefill_elapsed = prefill_start.elapsed();
let prefill_tps = if prefill_elapsed.as_secs_f64() > 0.0 {
prompt_len as f64 / prefill_elapsed.as_secs_f64()
} else {
0.0
};
let next_token = if sample_logits {
let mut logits = prefill_logits.to_vec();
sample_qwen35_logits_for_generate(&mut logits, &args, &[])
} else {
greedy_argmax_last_token(&prefill_logits, model.cfg.vocab_size)
};
let decode_start = std::time::Instant::now();
let outcome = if sample_logits {
run_decode_loop(
next_token,
prompt_len,
max_seq,
args.max_tokens,
&eos_token_ids,
|prev_token, pos, generated_tokens| -> Result<u32> {
let decode_positions = vec![pos; 4];
let mut logits = model
.forward_gpu_last_logits(
&[prev_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_last_logits decode at pos {pos} (benchmark)")
})?;
Ok(sample_qwen35_logits_for_generate(
&mut logits,
&args,
generated_tokens,
))
},
|_toks| String::new(),
|_event| {},
)?
} else {
run_decode_loop(
next_token,
prompt_len,
max_seq,
args.max_tokens,
&eos_token_ids,
|prev_token, pos, _generated_tokens| -> Result<u32> {
let decode_positions = vec![pos; 4];
model
.forward_gpu_greedy(
&[prev_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| {
format!("forward_gpu_greedy decode at pos {pos} (benchmark)")
})
},
|_toks| String::new(),
|_event| {},
)?
};
let decode_elapsed = decode_start.elapsed();
let gen = outcome.generated;
let decode_tps = if decode_elapsed.as_secs_f64() > 0.0 {
gen as f64 / decode_elapsed.as_secs_f64()
} else {
0.0
};
eprintln!(
" Run {}/{}: prefill {} tok in {:.2}s ({:.1} tok/s); decode {} tok in {:.2}s ({:.1} tok/s)",
run_idx + 1,
BENCH_NUM_RUNS,
prompt_len,
prefill_elapsed.as_secs_f64(),
prefill_tps,
gen,
decode_elapsed.as_secs_f64(),
decode_tps,
);
prefill_tps_runs.push(prefill_tps);
decode_tps_runs.push(decode_tps);
generated_per_run.push(gen);
}
print_benchmark_summary(
model_path,
prompt_len,
&generated_per_run,
Some(&prefill_tps_runs),
&decode_tps_runs,
&[],
);
return Ok(());
}
tracing::info!("Qwen3.5 prefill: seq_len={}", prompt_len);
let profile_sync = std::env::var("HF2Q_PROFILE_SYNC").is_ok();
if profile_sync {
mlx_native::reset_counters();
}
let prefill_start = std::time::Instant::now();
let prefill_logits = model
.forward_gpu_last_logits(&prompt_tokens, &prefill_positions, &mut kv_cache, SlotId(0))
.context("Qwen35Model::forward_gpu_last_logits (prefill)")?;
let prefill_elapsed = prefill_start.elapsed();
if profile_sync {
eprintln!(
"[P19 H9] prefill seq_len={} elapsed_ms={:.1} sync_count={} dispatch_count={} barrier_count={} cmd_buf_count={}",
prompt_len,
prefill_elapsed.as_secs_f64() * 1000.0,
mlx_native::sync_count(),
mlx_native::dispatch_count(),
mlx_native::barrier_count(),
mlx_native::cmd_buf_count(),
);
}
let vocab_size = model.cfg.vocab_size;
anyhow::ensure!(
prefill_logits.len() == vocab_size as usize,
"forward_gpu_last_logits (prefill) returned logits.len()={} != vocab({})",
prefill_logits.len(),
vocab_size,
);
let prefill_tok_s = prompt_len as f64 / prefill_elapsed.as_secs_f64();
header::print_header_prefill(
&mut stdout,
&header::HeaderInfoPrefill {
prefill_n: prompt_len,
prefill_ms: prefill_elapsed.as_secs_f64() * 1000.0,
prefill_tok_s,
},
stdout_is_tty,
)
.context("print header prefill")?;
if std::env::var("HF2Q_DUMP_LOGITS").as_deref() == Ok("1") {
let last_logits = &prefill_logits;
let bytes: &[u8] = unsafe {
std::slice::from_raw_parts(last_logits.as_ptr() as *const u8, last_logits.len() * 4)
};
std::fs::write("/tmp/hf2q_logits_t0.bin", bytes)
.context("HF2Q_DUMP_LOGITS: write /tmp/hf2q_logits_t0.bin")?;
let mut indexed: Vec<(usize, f32)> = last_logits.iter().copied().enumerate().collect();
indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
eprintln!(
"HF2Q_DUMP_LOGITS: wrote {} f32 values to /tmp/hf2q_logits_t0.bin",
last_logits.len()
);
eprintln!(" top-3: {:?}", &indexed[..3.min(indexed.len())]);
return Ok(());
}
let last_prefill_logits = &prefill_logits;
let next_token = if sample_logits {
let mut logits = last_prefill_logits.to_vec();
sample_qwen35_logits_for_generate(&mut logits, &args, &[])
} else {
greedy_argmax_last_token(last_prefill_logits, vocab_size)
};
tracing::info!("Qwen3.5 first decoded token: {}", next_token);
let decode_start = std::time::Instant::now();
let step_profile_enabled = std::env::var("HF2Q_STEP_PROFILE").is_ok();
let mut display = DisplayScrubber::new();
let outcome = if sample_logits {
let mut decode_stream = tokenizer.decode_stream(false);
let mut cumulative_text = String::new();
let mut last_decoded_count = 1usize;
if let Ok(Some(s)) = decode_stream.step(next_token) {
cumulative_text.push_str(&s);
}
run_decode_loop(
next_token,
prompt_len,
max_seq,
args.max_tokens,
&eos_token_ids,
|prev_token, pos, generated_tokens| -> Result<u32> {
let decode_positions = vec![pos; 4];
let _t_step = step_profile_enabled.then(std::time::Instant::now);
let _t_fwd = step_profile_enabled.then(std::time::Instant::now);
let mut logits = model
.forward_gpu_last_logits(
&[prev_token],
&decode_positions,
&mut kv_cache,
SlotId(0),
)
.with_context(|| format!("forward_gpu_last_logits decode at pos {pos}"))?;
let fwd_us = _t_fwd.map(|t| t.elapsed().as_micros());
let _t_smp = step_profile_enabled.then(std::time::Instant::now);
let next = sample_qwen35_logits_for_generate(&mut logits, &args, generated_tokens);
let smp_us = _t_smp.map(|t| t.elapsed().as_micros());
if let Some(t) = _t_step {
eprintln!(
"[STEP_PROFILE] pos={pos} total={:.2}ms fwd={:.2}ms smp={:.2}ms",
t.elapsed().as_micros() as f64 / 1000.0,
fwd_us.unwrap_or(0) as f64 / 1000.0,
smp_us.unwrap_or(0) as f64 / 1000.0,
);
}
Ok(next)
},
|toks| {
while last_decoded_count < toks.len() {
let tok = toks[last_decoded_count];
if let Ok(Some(s)) = decode_stream.step(tok) {
cumulative_text.push_str(&s);
}
last_decoded_count += 1;
}
cumulative_text.clone()
},
|event| {
let cleaned = display.push(event.delta);
if !cleaned.is_empty() {
print!("{}", cleaned);
let _ = stdout.flush();
}
},
)?
} else {
let mut decode_stream = tokenizer.decode_stream(false);
let mut cumulative_text = String::new();
let mut last_decoded_count = 1usize;
if let Ok(Some(s)) = decode_stream.step(next_token) {
cumulative_text.push_str(&s);
}
run_decode_loop(
next_token,
prompt_len,
max_seq,
args.max_tokens,
&eos_token_ids,
|prev_token, pos, _generated_tokens| -> Result<u32> {
let decode_positions = vec![pos; 4];
let _t_step = step_profile_enabled.then(std::time::Instant::now);
let next = model
.forward_gpu_greedy(&[prev_token], &decode_positions, &mut kv_cache, SlotId(0))
.with_context(|| format!("forward_gpu_greedy decode at pos {pos}"))?;
if let Some(t) = _t_step {
eprintln!(
"[STEP_PROFILE] pos={pos} total={:.2}ms",
t.elapsed().as_micros() as f64 / 1000.0
);
}
Ok(next)
},
|toks| {
while last_decoded_count < toks.len() {
let tok = toks[last_decoded_count];
if let Ok(Some(s)) = decode_stream.step(tok) {
cumulative_text.push_str(&s);
}
last_decoded_count += 1;
}
cumulative_text.clone()
},
|event| {
let cleaned = display.push(event.delta);
if !cleaned.is_empty() {
print!("{}", cleaned);
let _ = stdout.flush();
}
},
)?
};
let trailing = display.flush();
if !trailing.is_empty() {
print!("{}", trailing);
let _ = stdout.flush();
}
match &outcome.stop_reason {
DecodeStopReason::EosTokenId(_) | DecodeStopReason::MaxTokensReached => {}
DecodeStopReason::MaxSeq => {
tracing::warn!(
"Qwen3.5 decode: reached max_seq {} after {} tokens; stopping",
max_seq,
outcome.generated
);
}
DecodeStopReason::SpecialTokenLeak(marker) => {
tracing::info!(
"Qwen3.5 decode: special-token string `{marker}` detected in \
cumulative decode after {} tokens; stopping.",
outcome.generated
);
}
DecodeStopReason::RepetitionLoop { ngram, repeats } => {
tracing::info!(
"Qwen3.5 decode: consecutive n-gram repetition detected after {} \
tokens (last {} tokens repeated {} times consecutively); stopping.",
outcome.generated,
ngram,
repeats
);
if args.no_thinking && outcome.generated < 200 {
eprintln!(
"\n[hf2q] decode entered a {}-token consecutive-repetition loop \
after only {} tokens with `--no-thinking`. On this Qwen3.5/3.6 \
thinking-capable checkpoint, the empty `<think></think>` \
suppressor is a known degenerate attractor for some prompts \
(both hf2q AND llama.cpp produce the same loop). \
Recommended fix: drop `--no-thinking` and let the auto-\
detected thinking-mode handle it (the model will emit a \
reasoning trace, then the answer). \
Alternative escape: raise `--repetition-penalty` or use \
`--temperature` >0.",
ngram, outcome.generated
);
} else {
eprintln!(
"\n[hf2q] decode entered a {}-token consecutive-repetition loop \
after {} tokens — stopping. The model emitted the same \
{}-token sequence {} times in a row; raise --repetition-penalty \
or use a higher --temperature / lower --min-p to escape.",
ngram, outcome.generated, ngram, repeats
);
}
}
}
let generated = outcome.generated;
let decode_elapsed = decode_start.elapsed();
let tok_per_sec = generated as f64 / decode_elapsed.as_secs_f64();
let (td, tr) = if std::io::IsTerminal::is_terminal(&std::io::stderr()) {
("\x1b[2m", "\x1b[0m")
} else {
("", "")
};
eprintln!(
"\n\n{td}--- mlx-native (qwen35): {} tokens in {:.2}s ({:.1} tok/s) ---{tr}",
generated,
decode_elapsed.as_secs_f64(),
tok_per_sec,
);
if std::env::var("HF2Q_DUMP_COUNTERS").ok().as_deref() == Some("1") {
let dispatches = mlx_native::dispatch_count();
let syncs = mlx_native::sync_count();
eprintln!(
"[MLX_COUNTERS] dispatches={} syncs={} prompt_tokens={} decode_tokens={}",
dispatches, syncs, prompt_len, generated,
);
let buckets = mlx_native::pipeline_dispatch_buckets();
if !buckets.is_empty() {
let total: u64 = buckets.iter().map(|(_, c)| *c).sum();
eprintln!(
"[MLX_DISP_BUCKET] Per-pipeline breakdown ({} unique pipelines, total={}):",
buckets.len(),
total,
);
for (label, count) in &buckets {
let pct = if total > 0 {
100.0 * (*count as f64) / (total as f64)
} else {
0.0
};
eprintln!(
"[MLX_DISP_BUCKET] {:>8} ({:5.2}%) {}",
count, pct, label,
);
}
}
}
Ok(())
}
pub fn pool_key_for_path(path: &Path) -> String {
path.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| path.to_string_lossy().into_owned())
}
pub fn load_engine(path: &Path, config: &multi_model::EngineConfig) -> Result<api::engine::Engine> {
anyhow::ensure!(path.exists(), "Model not found: {}", path.display());
{
let gguf = mlx_native::gguf::GgufFile::open(path)
.map_err(|e| anyhow::anyhow!("GGUF header parse failed: {e}"))?;
let arch = gguf
.metadata_string("general.architecture")
.map(|s| s.to_string())
.unwrap_or_default();
tracing::info!(
path = %path.display(),
tensors = gguf.tensor_count(),
metadata = gguf.metadata_count(),
arch = %arch,
"Validated GGUF header"
);
}
let load_opts = api::engine::LoadOptions {
model_path: path.to_path_buf(),
tokenizer_path: config.tokenizer_path.clone(),
config_path: config.config_path.clone(),
dwq_overlay_path: config.dwq_overlay_path.clone(),
kv_persist_dir: std::env::var("HF2Q_KV_PERSIST")
.ok()
.filter(|s| !s.is_empty())
.map(std::path::PathBuf::from),
};
let mut loaded = api::engine::LoadedModel::load(&load_opts)?;
if let Some(sink) = config.kv_metrics_sink.as_ref() {
match &mut loaded {
api::engine::LoadedModel::Gemma(g) => {
g.kv_metrics_sink = Some(std::sync::Arc::clone(sink));
}
api::engine::LoadedModel::Qwen35(q) => {
q.kv_metrics_sink = Some(std::sync::Arc::clone(sink));
}
api::engine::LoadedModel::Qwen3VlText(_) => {
}
api::engine::LoadedModel::Deepseek4(_) => {
}
}
}
let engine = api::engine::Engine::spawn_with_mode(
loaded,
config.queue_capacity,
None,
config.engine_mode,
)
.map_err(|e| {
anyhow::anyhow!(
"ADR-040 Phase C iter-4 (C4): Engine::spawn_with_mode rejected the \
requested EngineMode: {e}. Either select `--scheduler fifo_serial` \
(the default; ADR-005 byte-equivalent path), or wait for the iter \
named in the error message to land."
)
})?;
load_info::emit_tracing(engine.info());
if config.warmup_synchronously {
let warmup_rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.context("build tokio runtime for synchronous engine warmup")?;
let warmup_started = std::time::Instant::now();
warmup_rt
.block_on(engine.warmup())
.context("synchronous engine warmup")?;
tracing::info!(
elapsed_ms = warmup_started.elapsed().as_millis() as u64,
"Engine warmed up synchronously (pre-mmproj order — iter-103 fix)"
);
drop(warmup_rt);
}
Ok(engine)
}
pub(crate) fn should_enable_kv_persist(flag: Option<&str>, env: Option<&str>) -> bool {
flag.is_some() && env.map(|v| v.trim()).unwrap_or("") != "0"
}
pub(crate) fn maybe_print_serve_banner<W: std::io::Write>(
info: &load_info::LoadInfo,
w: &mut W,
tty: bool,
quiet: bool,
) -> std::io::Result<()> {
if tty && !quiet {
load_info::print_banner(info, w, tty)?;
}
Ok(())
}
pub(crate) fn parse_scheduler_config(
scheduler_cli: Option<cli::SchedulerArg>,
scheduler_env: Option<&str>,
max_slots_cli: Option<u32>,
max_slots_env: Option<&str>,
) -> std::result::Result<api::engine::EngineMode, String> {
use api::engine::EngineMode;
let scheduler_from_env: Option<cli::SchedulerArg> =
match scheduler_env.map(|s| s.trim()).filter(|s| !s.is_empty()) {
None => None,
Some(raw) => {
let lower = raw.to_ascii_lowercase();
match lower.as_str() {
"fifo_serial" => Some(cli::SchedulerArg::FifoSerial),
"inflight_batched" => Some(cli::SchedulerArg::InflightBatched),
_ => {
return Err(format!(
"ADR-040 C4: HF2Q_SCHEDULER={raw:?} is not a recognized \
scheduler policy. Supported values (case-insensitive): \
`fifo_serial` (default; ADR-005 byte-equivalent path), \
`inflight_batched` (ADR-040 slot-aware path, gated on \
iter-2b/2c worker-arm landing). Unset the env var to \
use the default."
));
}
}
}
};
let scheduler = scheduler_cli.or(scheduler_from_env);
let max_slots_from_env: Option<u32> =
match max_slots_env.map(|s| s.trim()).filter(|s| !s.is_empty()) {
None => None,
Some(raw) => match raw.parse::<u32>() {
Ok(parsed) => Some(parsed),
Err(err) => {
return Err(format!(
"ADR-040 C4: HF2Q_MAX_SLOTS={raw:?} does not parse as a \
non-negative u32: {err}. Supply a positive integer (default \
{}), or unset the env var.",
DEFAULT_MAX_SLOTS_UNDER_INFLIGHT
));
}
},
};
let max_slots_requested = max_slots_cli.or(max_slots_from_env);
if let Some(0) = max_slots_requested {
return Err(format!(
"ADR-040 C4: --max-slots=0 / HF2Q_MAX_SLOTS=0 is rejected per \
ADR-040 iter-2.5 F3a (FifoSchedulerAdapter `.max(1)` discipline \
applied at the CLI layer instead of silently coerced). Either \
omit the flag (defaults to {} when `--scheduler inflight_batched` \
is selected; ignored otherwise), or supply a positive integer.",
DEFAULT_MAX_SLOTS_UNDER_INFLIGHT
));
}
match scheduler {
None | Some(cli::SchedulerArg::FifoSerial) => Ok(EngineMode::SerialFifo),
Some(cli::SchedulerArg::InflightBatched) => {
let max_slots = max_slots_requested.unwrap_or(DEFAULT_MAX_SLOTS_UNDER_INFLIGHT);
Ok(EngineMode::SlotAware {
max_slots: max_slots.max(1),
})
}
}
}
pub(crate) const DEFAULT_MAX_SLOTS_UNDER_INFLIGHT: u32 = 4;
pub fn cmd_serve(args: cli::ServeArgs) -> Result<()> {
use api::schema::OverflowPolicy;
use api::state::ServerConfig;
let auth_token = args.auth_token.clone().or_else(|| {
std::env::var("HF2Q_AUTH_TOKEN")
.ok()
.filter(|s| !s.is_empty())
});
let overflow_policy = match args.overflow_policy {
cli::OverflowPolicyArg::Reject => OverflowPolicy::Reject,
cli::OverflowPolicyArg::TruncateLeft => OverflowPolicy::TruncateLeft,
cli::OverflowPolicyArg::Summarize => OverflowPolicy::Summarize,
};
let cache_dir = args
.cache_dir
.clone()
.or_else(api::state::default_cache_dir);
let scheduler_env = std::env::var("HF2Q_SCHEDULER").ok();
let max_slots_env = std::env::var("HF2Q_MAX_SLOTS").ok();
let engine_mode = parse_scheduler_config(
args.scheduler,
scheduler_env.as_deref(),
args.max_slots,
max_slots_env.as_deref(),
)
.map_err(|msg| anyhow::anyhow!("{msg}"))?;
tracing::info!(
engine_mode = ?engine_mode,
scheduler_cli = ?args.scheduler,
scheduler_env = ?scheduler_env,
max_slots_cli = ?args.max_slots,
max_slots_env = ?max_slots_env,
"ADR-040 C4: resolved scheduler policy"
);
let config = ServerConfig {
host: args.host.clone(),
port: args.port,
auth_token,
cors_allowed_origins: args.cors_origins.clone(),
queue_capacity: args.queue_capacity,
max_concurrent_requests: 0,
request_timeout_seconds: 0,
default_overflow_policy: overflow_policy,
cache_dir,
system_fingerprint: Some(system_fingerprint()),
};
if args.host == "0.0.0.0" {
tracing::warn!(
"Server bound to 0.0.0.0 — exposes API on all interfaces. \
Public-internet exposure is NOT a supported deployment target \
(see ADR-005 Decision #13: reverse-proxy assumption). \
For LAN-only Open WebUI, this is the intended usage."
);
}
let default_model_arg = args
.model
.as_ref()
.map(|p| p.to_string_lossy().into_owned());
let mut state = api::AppState::new_for_serve(
config.clone(),
args.no_integrity,
config.queue_capacity,
default_model_arg.clone(),
)?;
let kv_persist_env = std::env::var("HF2Q_KV_PERSIST").ok();
let kv_persist_flag_path = args.kv_persist_path.as_ref();
let kv_persist_enabled = should_enable_kv_persist(
kv_persist_flag_path.map(|p| p.to_string_lossy()).as_deref(),
kv_persist_env.as_deref(),
);
if !kv_persist_enabled && kv_persist_flag_path.is_some() {
let path_for_log = kv_persist_flag_path
.map(|p| p.display().to_string())
.unwrap_or_default();
tracing::warn!(
kv_persist_path = %path_for_log,
"ADR-017 R-F1 override: HF2Q_KV_PERSIST=0 — disabling kv-persist \
despite --kv-persist={path_for_log}. Operator emergency-disable per \
operating-kv-cache.md §10. Restart without HF2Q_KV_PERSIST=0 to \
re-enable.",
path_for_log = path_for_log,
);
}
let kv_persist_loader_wrapper: Option<
std::sync::Arc<crate::serve::kv_persist::LoaderWrapper<api::engine::Engine>>,
> = if let Some(cache_dir) = args.kv_persist_path.as_ref().filter(|_| kv_persist_enabled) {
use crate::serve::kv_persist::families::gemma4_dense::{
Gemma4DenseConfig, Gemma4DenseSpillFactory,
};
use crate::serve::kv_persist::families::tq_packed::{
flags as tq_flags, TqBitsPerCoord, TqPackedConfig, TqPackedSpillFactory,
};
use crate::serve::kv_persist::registry::FamilyHookFactory;
use crate::serve::kv_persist::{
AsyncWriterHandle, BlockPrefixCacheSpiller, DiskBlockStore, KvPersistRegistry,
LoaderWrapper, StubGemma4Spill, DEFAULT_CHANNEL_CAPACITY,
};
use crate::serve::multi_model::{DefaultModelLoader, HotSwapManager, LoadedPool};
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
std::fs::create_dir_all(cache_dir).with_context(|| {
format!(
"ADR-017 C.1: create kv-persist cache dir at {}",
cache_dir.display()
)
})?;
let metrics_sink: Arc<dyn crate::serve::kv_persist::metrics::KvCacheMetricsSink> =
Arc::clone(&state.kv_spill_counters)
as Arc<dyn crate::serve::kv_persist::metrics::KvCacheMetricsSink>;
let (recovered_index, recovery_report) =
crate::serve::kv_persist::recover_from_disk_with_counters(
cache_dir,
Some(&metrics_sink),
)
.with_context(|| format!("ADR-017 C.1: recover_from_disk({})", cache_dir.display()))?;
tracing::info!(
cache_dir = %cache_dir.display(),
blocks_indexed = recovery_report.blocks_indexed,
blocks_quarantined = recovery_report.blocks_quarantined,
bytes_indexed = recovery_report.bytes_indexed,
elapsed_ms = recovery_report.elapsed_ms,
"ADR-017 C.1: kv-persist recovery scan complete"
);
let kv_persist_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-017 P1-3: HF2Q_KV_PERSIST_BUDGET_BYTES \
parse failed; defaulting to 0 (unlimited)"
);
0
}
},
Err(_) => 0,
};
let store = Arc::new(
DiskBlockStore::new_with_index(cache_dir.clone(), recovered_index, 0).with_context(
|| {
format!(
"ADR-017 C.1: DiskBlockStore::new_with_index({})",
cache_dir.display()
)
},
)?,
);
store.set_budget_bytes(kv_persist_budget_bytes);
store.set_kv_counters(Arc::clone(&metrics_sink));
state.kv_disk_store = Some(Arc::clone(&store));
tracing::info!(
budget_bytes = kv_persist_budget_bytes,
"ADR-017 P1-3: HF2Q_KV_PERSIST_BUDGET_BYTES wired (0 = unlimited)"
);
let writer = Arc::new(AsyncWriterHandle::spawn(
Arc::clone(&store),
DEFAULT_CHANNEL_CAPACITY,
));
let spiller: Arc<BlockPrefixCacheSpiller<api::engine::Engine>> = Arc::new(
BlockPrefixCacheSpiller::new(Arc::clone(&store), Arc::clone(&writer)),
);
let registry = Arc::new(KvPersistRegistry::new());
spiller.set_registry(Arc::clone(®istry));
if let Some(model_arg) = default_model_arg.as_deref() {
let stub = Arc::new(StubGemma4Spill);
let stub_for_spiller: Arc<Mutex<dyn crate::serve::kv_persist::KvCacheSpill>> =
Arc::new(Mutex::new(StubGemma4Spill));
let stub_for_registry: Arc<dyn crate::serve::kv_persist::EngineBindable> = stub.clone();
let pool_repo = pool_key_for_path(&PathBuf::from(model_arg));
let pool_quant = quant_select::QuantType::Q4_K_M;
spiller.register_family(pool_repo.clone(), pool_quant, stub_for_spiller);
registry.register(pool_repo.clone(), pool_quant, stub_for_registry);
tracing::info!(
repo = %pool_repo,
quant = %pool_quant.as_str(),
"ADR-017 C.1: registered StubGemma4Spill for operator --model"
);
let fallback_cfg = Gemma4DenseConfig {
layer_types: vec![
crate::serve::config::LayerType::Sliding,
crate::serve::config::LayerType::Full,
],
nkv_heads: vec![8, 2],
head_dim: vec![256, 512],
kv_dtype: mlx_native::DType::F32,
sliding_window: 4096,
max_decode_tokens: 8192,
};
let factory: Arc<dyn FamilyHookFactory> =
if crate::serve::api::tq_packed_descriptor::is_tq_active_mode() {
let bits = TqBitsPerCoord::new(
crate::serve::api::tq_packed_descriptor::parse_tq_codebook_bits(
std::env::var("HF2Q_TQ_CODEBOOK_BITS").ok().as_deref(),
),
)
.expect("HF2Q_TQ_CODEBOOK_BITS validated to be in {2,3,4,5,6,8}");
let tq_fallback_cfg = TqPackedConfig {
num_layers: 2,
nkv_heads: vec![8, 2],
head_dim: vec![256, 512],
bits_per_coord: bits,
scale: 1.0,
flags: tq_flags::HADAMARD_ROTATED,
block_tokens: crate::serve::kv_persist::format::BLOCK_TOKENS,
};
Arc::new(TqPackedSpillFactory::new(tq_fallback_cfg))
} else {
Arc::new(Gemma4DenseSpillFactory::new(fallback_cfg))
};
let factory_kind = if crate::serve::api::tq_packed_descriptor::is_tq_active_mode() {
"TqPackedSpillFactory"
} else {
"Gemma4DenseSpillFactory"
};
registry.register_factory(pool_repo.clone(), pool_quant, factory);
tracing::info!(
repo = %pool_repo,
quant = %pool_quant.as_str(),
factory = factory_kind,
"ADR-017 B-dense.2 + B-tq.4 iter-2: registered single-mode factory \
(lazy real-hook construction at first engine load); HF2Q_TQ_KV \
selects TQ-active mode at startup"
);
}
let real_loader: Arc<dyn crate::serve::multi_model::ModelLoader<api::engine::Engine>> =
Arc::new(DefaultModelLoader);
let wrapper = Arc::new(LoaderWrapper::new(real_loader, Arc::clone(®istry)));
wrapper.set_spiller(Arc::clone(&spiller));
let loader_for_manager: Arc<
dyn crate::serve::multi_model::ModelLoader<api::engine::Engine>,
> = wrapper.clone();
let pool = LoadedPool::from_hardware(state.hardware.as_ref());
let mut manager: HotSwapManager<api::engine::Engine> = HotSwapManager::new_with_spiller(
pool,
loader_for_manager,
Arc::clone(&spiller)
as Arc<dyn crate::serve::multi_model::KvSpiller<api::engine::Engine>>,
);
manager.set_kv_counters(Arc::clone(&state.kv_spill_counters));
state.pool = Arc::new(std::sync::RwLock::new(manager));
state.kv_spiller = Some(spiller);
tracing::info!(
cache_dir = %cache_dir.display(),
"ADR-017 C.1: kv-persist spiller substrate wired into HotSwapManager + AppState"
);
Some(wrapper)
} else {
None
};
let mut startup_engine_for_banner: Option<api::engine::Engine> = None;
if let Some(model_arg) = default_model_arg.as_ref() {
let mut cache_guard = state
.cache
.lock()
.map_err(|e| anyhow::anyhow!("cache mutex poisoned at startup: {e}"))?;
let resolved = auto_pipeline::resolve_or_prepare_model(
model_arg,
&mut cache_guard,
state.hardware.as_ref(),
state.no_integrity,
)
.context("auto-pipeline: resolve --model into a GGUF path")?;
drop(cache_guard);
if let Some(repo) = resolved.repo_id.as_deref() {
let quant_str: &str = resolved
.quant
.map(quant_select::QuantType::as_str)
.unwrap_or("");
tracing::info!(
repo,
quant = quant_str,
from_cache = resolved.from_cache,
gguf = %resolved.gguf_path.display(),
"auto-pipeline: --model resolved"
);
}
let pool_repo = resolved
.repo_id
.clone()
.unwrap_or_else(|| pool_key_for_path(&resolved.gguf_path));
let pool_quant = resolved.quant.unwrap_or(quant_select::QuantType::Q4_K_M);
let engine_config = multi_model::EngineConfig {
tokenizer_path: args.tokenizer.clone(),
config_path: args.config.clone(),
queue_capacity: config.queue_capacity,
warmup_synchronously: true,
kv_metrics_sink: Some(std::sync::Arc::clone(&state.kv_spill_counters)
as std::sync::Arc<dyn crate::serve::kv_persist::metrics::KvCacheMetricsSink>),
dwq_overlay_path: args.dwq_overlay.clone(),
engine_mode,
};
if let Some(wrapper) = kv_persist_loader_wrapper.as_ref() {
wrapper.set_pending_bind(pool_repo.clone(), pool_quant);
}
let mut pool_guard = state
.pool
.write()
.map_err(|e| anyhow::anyhow!("pool rwlock poisoned at startup: {e}"))?;
let loaded_engine = pool_guard
.load_or_get(&pool_repo, pool_quant, &resolved.gguf_path, &engine_config)
.map_err(|e| anyhow::anyhow!("startup pre-warm: {e}"))?;
startup_engine_for_banner = Some(loaded_engine.engine.clone());
drop(pool_guard);
tracing::info!(
repo = %pool_repo,
quant = %pool_quant.as_str(),
"hf2q startup pre-warm: model admitted to pool"
);
}
let embedding_model = if let Some(emb_path) = args.embedding_model.as_ref() {
anyhow::ensure!(
emb_path.exists(),
"Embedding model not found: {}",
emb_path.display()
);
let gguf = mlx_native::gguf::GgufFile::open(emb_path)
.map_err(|e| anyhow::anyhow!("Embedding GGUF header parse failed: {e}"))?;
let arch_str = gguf
.metadata_string("general.architecture")
.ok_or_else(|| anyhow::anyhow!("Embedding GGUF missing general.architecture"))?
.to_string();
let vocab = crate::inference::models::bert::BertVocab::from_gguf(&gguf)
.map_err(|e| anyhow::anyhow!("Embedding GGUF vocab parse failed: {e}"))?;
let tokenizer = crate::inference::models::bert::BertWpmTokenizer::new(&vocab);
let model_id = emb_path
.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| "embedding-model".into());
let device = mlx_native::MlxDevice::new()
.map_err(|e| anyhow::anyhow!("create MlxDevice for embedding load: {e}"))?;
let arch = match arch_str.as_str() {
"bert" => {
let cfg = crate::inference::models::bert::BertConfig::from_gguf(&gguf)
.map_err(|e| anyhow::anyhow!("BERT GGUF config parse failed: {e}"))?;
crate::inference::models::bert::weights::validate_tensor_set(&gguf, &cfg)
.map_err(|e| anyhow::anyhow!("BERT GGUF tensor validation: {e}"))?;
let weights = crate::inference::models::bert::weights::LoadedBertWeights::load(
&gguf, &cfg, device,
)
.map_err(|e| anyhow::anyhow!("BERT weights load failed: {e}"))?;
tracing::info!(
path = %emb_path.display(),
arch = "bert",
hidden = cfg.hidden_size,
layers = cfg.num_hidden_layers,
pooling = ?cfg.pooling_type,
vocab_size = vocab.len(),
tensor_count = weights.len(),
"Validated embedding GGUF + loaded weights onto device"
);
api::state::EmbeddingArch::Bert {
config: cfg,
weights: std::sync::Arc::new(weights),
}
}
"nomic-bert" => {
let cfg =
crate::inference::models::nomic_bert::NomicBertConfig::from_gguf(&gguf)
.map_err(|e| anyhow::anyhow!("nomic-bert GGUF config parse failed: {e}"))?;
crate::inference::models::nomic_bert::validate_tensor_set(&gguf, &cfg)
.map_err(|e| anyhow::anyhow!("nomic-bert GGUF tensor validation: {e}"))?;
let weights = crate::inference::models::nomic_bert::LoadedNomicBertWeights::load(
&gguf, &cfg, device,
)
.map_err(|e| anyhow::anyhow!("nomic-bert weights load failed: {e}"))?;
tracing::info!(
path = %emb_path.display(),
arch = "nomic-bert",
hidden = cfg.hidden_size,
layers = cfg.num_hidden_layers,
pooling = ?cfg.pooling_type,
rope_freq_base = cfg.rope_freq_base,
vocab_size = vocab.len(),
tensor_count = weights.len(),
"Validated embedding GGUF + loaded weights onto device"
);
api::state::EmbeddingArch::NomicBert {
config: cfg,
weights: std::sync::Arc::new(weights),
}
}
other => {
anyhow::bail!(
"embedding GGUF general.architecture='{other}' is not supported. \
Phase 2b day-one models: 'bert' (bge / mxbai) and 'nomic-bert' \
(nomic-embed-text-v1.5). File: {}",
emb_path.display()
);
}
};
Some(api::state::EmbeddingModel {
gguf_path: emb_path.clone(),
vocab: std::sync::Arc::new(vocab),
tokenizer: std::sync::Arc::new(tokenizer),
model_id,
arch: Some(arch),
})
} else {
None
};
let mmproj = if let Some(mmp_path) = args.mmproj.as_ref() {
anyhow::ensure!(
mmp_path.exists(),
"mmproj not found: {}",
mmp_path.display()
);
let gguf = mlx_native::gguf::GgufFile::open(mmp_path)
.map_err(|e| anyhow::anyhow!("mmproj GGUF header parse failed: {e}"))?;
let mmp_config = crate::inference::vision::mmproj::MmprojConfig::from_gguf(&gguf)
.map_err(|e| anyhow::anyhow!("mmproj GGUF config parse failed: {e}"))?;
let actual_names: Vec<&str> = gguf.tensor_names();
crate::inference::vision::mmproj::validate_tensor_set(&mmp_config, &actual_names)
.map_err(|e| anyhow::anyhow!("mmproj GGUF tensor-set validation: {e}"))?;
let arch = crate::inference::vision::mmproj::detect_arch_profile_with_projector(
&mmp_config.projector,
&actual_names,
);
if !arch.is_supported() {
anyhow::bail!(
"mmproj arch profile is Unknown — neither Gemma 4 \
SigLIP markers (ln1/ln2/post_ffw_norm) nor CLIP marker \
(attn_norm) found in block 0. hf2q's ViT forward pass \
cannot dispatch on this file."
);
}
let skip_mmproj_load = std::env::var("HF2Q_SKIP_MMPROJ_LOAD").as_deref() == Ok("1");
let device = mlx_native::MlxDevice::new()
.map_err(|e| anyhow::anyhow!("create MlxDevice for mmproj load: {e}"))?;
let mmp_weights = if skip_mmproj_load {
tracing::warn!(
"HF2Q_SKIP_MMPROJ_LOAD=1 — using empty mmproj weights; \
vision requests will 500 on first forward attempt"
);
crate::inference::vision::mmproj_weights::LoadedMmprojWeights::empty(device)
} else {
crate::inference::vision::mmproj_weights::LoadedMmprojWeights::load(
&gguf,
&mmp_config,
device,
)
.map_err(|e| anyhow::anyhow!("mmproj weight load: {e}"))?
};
let model_id = mmp_path
.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| "mmproj".into());
tracing::info!(
path = %mmp_path.display(),
image_size = mmp_config.image_size,
patch_size = mmp_config.patch_size,
hidden = mmp_config.hidden_size,
layers = mmp_config.num_hidden_layers,
projector = mmp_config.projector.as_str(),
arch = arch.as_str(),
tensors_loaded = mmp_weights.len(),
"Loaded mmproj GGUF header + tensor set + weights"
);
let skip_vit_warmup = std::env::var("HF2Q_SKIP_VIT_WARMUP").as_deref() == Ok("1");
if skip_vit_warmup {
tracing::warn!(
"HF2Q_SKIP_VIT_WARMUP=1 — skipping ViT GPU warmup; first \
multimodal request will pay kernel-compile cost"
);
} else {
let warmup_t0 = std::time::Instant::now();
match crate::inference::vision::vit_gpu::warmup_vit_gpu(&mmp_weights, &mmp_config) {
Ok(()) => tracing::info!(
elapsed_ms = warmup_t0.elapsed().as_millis() as u64,
"ViT GPU warmup complete"
),
Err(e) => tracing::warn!(
error = %e,
"ViT GPU warmup failed; first multimodal request will pay kernel-compile cost"
),
}
}
Some(api::state::LoadedMmproj {
gguf_path: mmp_path.clone(),
config: mmp_config,
arch,
weights: std::sync::Arc::new(mmp_weights),
model_id,
})
} else {
None
};
if let Some(em) = embedding_model {
let registry = build_warmed_embedding_registry(&em).context("warm embedding registry")?;
state = state
.with_embedding_model(em)
.with_embedding_registry(std::sync::Arc::new(std::sync::Mutex::new(registry)));
}
if let Some(m) = mmproj {
state = state.with_mmproj(m);
}
let state_for_warmup = state.clone();
let router = api::build_router(state);
let rt = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.context("building tokio runtime")?;
let stdout_is_tty = std::io::IsTerminal::is_terminal(&std::io::stdout());
let quiet = args.quiet;
rt.block_on(async move {
let bind = format!("{}:{}", config.host, config.port);
let listener = tokio::net::TcpListener::bind(&bind)
.await
.with_context(|| format!("binding to {bind}"))?;
let local_addr = listener.local_addr().ok();
tracing::info!(
addr = %local_addr.map(|a| a.to_string()).unwrap_or_else(|| bind.clone()),
"hf2q HTTP server listening"
);
if let Some(engine) = startup_engine_for_banner.as_ref() {
let mut stdout = std::io::stdout();
maybe_print_serve_banner(engine.info(), &mut stdout, stdout_is_tty, quiet)
.context("print serve load banner")?;
}
drop(startup_engine_for_banner);
eprintln!("hf2q serving on http://{}", bind);
{
let pool_state_log = state_for_warmup.pool.read().ok().map(|m| m.pool_stats());
if let Some(stats) = pool_state_log {
tracing::info!(
loaded = stats.loaded_count,
capacity = stats.capacity_models,
bytes_resident = stats.total_resident_bytes,
bytes_budget = stats.memory_budget_bytes,
"hf2q ready (pool-backed; pre-warm complete if --model supplied)"
);
}
}
axum::serve(listener, router)
.with_graceful_shutdown(shutdown_signal())
.await
.context("axum::serve")?;
let shutdown_engines: Vec<_> = state_for_warmup
.pool
.read()
.ok()
.map(|mgr| {
mgr.snapshot_engines()
.into_iter()
.map(|le| le.engine.clone())
.collect()
})
.unwrap_or_default();
let state_for_drain = state_for_warmup.clone();
let drain_summary = tokio::task::spawn_blocking(move || {
drain_loaded_models_to_disk(&state_for_drain, std::time::Duration::from_secs(30))
})
.await
.unwrap_or_default();
tracing::info!(
evicted = drain_summary.evicted,
drain_ms = drain_summary.drain_ms,
queue_depth_at_exit = drain_summary.queue_depth_at_exit,
timed_out = drain_summary.timed_out,
"ADR-017 graceful-shutdown KV-cache drain complete"
);
for engine in shutdown_engines {
match engine.shutdown().await {
Ok(()) => tracing::info!("hf2q-engine worker joined"),
Err(e) => tracing::warn!(error = %e, "hf2q-engine worker join failed"),
}
}
tracing::info!("hf2q HTTP server shut down cleanly");
Ok::<(), anyhow::Error>(())
})?;
Ok(())
}
fn system_fingerprint() -> String {
format!("hf2q-{}-mlx-native", env!("CARGO_PKG_VERSION"))
}
#[derive(Debug, Clone, Default)]
pub struct ShutdownDrainSummary {
pub evicted: usize,
pub drain_ms: u64,
pub queue_depth_at_exit: usize,
pub timed_out: bool,
}
pub fn drain_loaded_models_to_disk(
state: &api::state::AppState,
timeout: std::time::Duration,
) -> ShutdownDrainSummary {
let Some(spiller) = state.kv_spiller.as_ref() else {
return ShutdownDrainSummary::default();
};
let start = std::time::Instant::now();
let keys: Vec<(String, crate::serve::quant_select::QuantType)> = state
.pool
.read()
.ok()
.map(|mgr| {
mgr.snapshot_engines()
.into_iter()
.map(|le| (le.repo.clone(), le.quant))
.collect()
})
.unwrap_or_default();
let evicted = if keys.is_empty() {
0
} else if let Ok(mut mgr) = state.pool.write() {
let mut count = 0usize;
for (repo, quant) in &keys {
let _ = mgr.evict(repo, *quant);
count += 1;
}
count
} else {
0
};
let poll_interval = std::time::Duration::from_millis(50);
let mut queue_depth_at_exit = spiller.pending_writer_queue_depth();
let mut timed_out = false;
while queue_depth_at_exit > 0 {
if start.elapsed() >= timeout {
timed_out = true;
break;
}
std::thread::sleep(poll_interval);
queue_depth_at_exit = spiller.pending_writer_queue_depth();
}
let drain_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
ShutdownDrainSummary {
evicted,
drain_ms,
queue_depth_at_exit,
timed_out,
}
}
async fn shutdown_signal() {
let ctrl_c = async {
let _ = tokio::signal::ctrl_c().await;
};
#[cfg(unix)]
let terminate = async {
if let Ok(mut s) = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
{
s.recv().await;
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => tracing::info!("received SIGINT, shutting down"),
_ = terminate => tracing::info!("received SIGTERM, shutting down"),
}
}
pub fn cmd_cache(args: cli::CacheArgs) -> Result<()> {
use cli::CacheAction;
match args.action {
CacheAction::List {
kv_namespace: true,
kv_path,
} => {
return cmd_cache_kv_list(kv_path.as_deref());
}
CacheAction::Size {
kv_namespace: true,
kv_path,
} => {
return cmd_cache_kv_size(kv_path.as_deref());
}
CacheAction::Clear {
kv_namespace: true,
ref model,
ref quant,
all,
yes: _,
ref kv_path,
force,
} => {
return cmd_cache_kv_clear(
kv_path.as_deref(),
model.as_deref(),
quant.as_deref(),
all,
force,
);
}
_ => { }
}
let mut cache = cache::ModelCache::open().context("open model cache")?;
match args.action {
CacheAction::List { .. } => cmd_cache_list(&cache),
CacheAction::Size { .. } => cmd_cache_size(&cache),
CacheAction::Clear {
model,
quant,
all,
yes,
kv_namespace: _,
kv_path: _,
force: _,
} => cmd_cache_clear(&mut cache, model, quant, all, yes),
}
}
fn cmd_cache_kv_list(kv_path: Option<&Path>) -> Result<()> {
use crate::serve::kv_persist::cache_ops;
let kv_root = cache_ops::resolve_kv_root(kv_path)?;
let entries = cache_ops::list_namespaces(&kv_root)?;
if entries.is_empty() {
println!("(kv-cache empty — root: {})", kv_root.display());
return Ok(());
}
println!("hf2q kv-cache @ {}", kv_root.display());
println!(
"{:<20} {:>14} {:>10}",
"FP_SHORT", "BYTES_ON_DISK", "BLOCKS"
);
for e in &entries {
println!(
"{:<20} {:>14} {:>10}",
e.fp_short, e.bytes_on_disk, e.block_count
);
}
Ok(())
}
fn cmd_cache_kv_size(kv_path: Option<&Path>) -> Result<()> {
use crate::serve::kv_persist::cache_ops;
let kv_root = cache_ops::resolve_kv_root(kv_path)?;
let total = cache_ops::total_bytes(&kv_root);
println!(
"hf2q kv-cache @ {} — {} bytes ({:.2} GiB)",
kv_root.display(),
total,
total as f64 / (1u64 << 30) as f64,
);
Ok(())
}
fn cmd_cache_kv_clear(
kv_path: Option<&Path>,
model: Option<&str>,
quant: Option<&str>,
all: bool,
force: bool,
) -> Result<()> {
use crate::serve::kv_persist::cache_ops;
use crate::serve::quant_select::QuantType;
if all {
return Err(anyhow::anyhow!(
"hf2q cache clear --kv-namespace: --all is not supported. \
Per-repo scope only (operator runbook §11 #4). To wipe \
the entire kv-cache, stop `hf2q serve` and `rm -rf \
<kv-persist>/models <kv-persist>/locks` directly."
));
}
let Some(repo) = model else {
return Err(anyhow::anyhow!(
"hf2q cache clear --kv-namespace: --model <repo-id> required \
(no whole-cache wipe via this command — see operator runbook §11 #4)"
));
};
let kv_root = cache_ops::resolve_kv_root(kv_path)?;
if let Some(q_str) = quant {
let q =
QuantType::from_canonical_str(q_str).map_err(|e| anyhow::anyhow!("--quant: {}", e))?;
let outcome = cache_ops::clear_namespace(&kv_root, repo, q, force).map_err(|e| {
anyhow::anyhow!(
"hf2q cache clear --kv-namespace --model {} --quant {}: {}",
repo,
q.as_str(),
e
)
})?;
if outcome.existed {
println!(
"hf2q kv-cache: cleared {}@{} (fp_short={}, {} bytes freed)",
repo,
q.as_str(),
outcome.fp_short,
outcome.bytes_freed
);
} else {
println!(
"hf2q kv-cache: nothing to clear for {}@{} (fp_short={} not present)",
repo,
q.as_str(),
outcome.fp_short
);
}
} else {
let outcomes =
cache_ops::clear_namespace_all_quants(&kv_root, repo, force).map_err(|e| {
anyhow::anyhow!(
"hf2q cache clear --kv-namespace --model {} (all quants): {}",
repo,
e
)
})?;
let total_bytes: u64 = outcomes.iter().map(|o| o.bytes_freed).sum();
let removed: Vec<&str> = outcomes
.iter()
.filter(|o| o.existed)
.map(|o| o.fp_short.as_str())
.collect();
if removed.is_empty() {
println!(
"hf2q kv-cache: nothing to clear for {} (no quant variants present)",
repo
);
} else {
println!(
"hf2q kv-cache: cleared {} (all quants) — {} fp_short dirs, {} bytes freed: [{}]",
repo,
removed.len(),
total_bytes,
removed.join(", ")
);
}
}
Ok(())
}
fn cmd_cache_list(cache: &cache::ModelCache) -> Result<()> {
let entries: Vec<_> = cache.iter_entries().collect();
if entries.is_empty() {
println!("(cache empty — root: {})", cache.root().display());
return Ok(());
}
println!("hf2q cache @ {}", cache.root().display());
println!(
"{:<48} {:<10} {:>12} {:>20}",
"MODEL", "QUANT", "BYTES", "LAST_ACCESSED"
);
for view in &entries {
if view.model.quantizations.is_empty() {
println!(
"{:<48} {:<10} {:>12} {:>20}",
view.repo_id, "(none)", "-", view.model.last_accessed_secs,
);
continue;
}
for (quant, qe) in &view.model.quantizations {
println!(
"{:<48} {:<10} {:>12} {:>20}",
view.repo_id, quant, qe.bytes, view.model.last_accessed_secs,
);
}
}
Ok(())
}
fn cmd_cache_size(cache: &cache::ModelCache) -> Result<()> {
let total = cache.total_bytes_on_disk();
println!(
"hf2q cache @ {} — {} bytes ({:.2} GiB)",
cache.root().display(),
total,
total as f64 / (1u64 << 30) as f64,
);
Ok(())
}
fn cmd_cache_clear(
cache: &mut cache::ModelCache,
model: Option<String>,
quant: Option<String>,
all: bool,
yes: bool,
) -> Result<()> {
use crate::serve::quant_select::QuantType;
if all && (model.is_some() || quant.is_some()) {
return Err(anyhow::anyhow!(
"hf2q cache clear: --all is mutually exclusive with --model / --quant"
));
}
if !all && model.is_none() {
return Err(anyhow::anyhow!(
"hf2q cache clear: must specify --model <repo-id> [--quant <type>] \
OR --all --yes (the latter purges every cached model)"
));
}
if all {
if !yes {
return Err(anyhow::anyhow!(
"hf2q cache clear --all: refused without --yes \
(this would remove every cached model under {})",
cache.root().display()
));
}
let freed = cache.purge().context("purge cache")?;
println!(
"hf2q cache: purged ({} bytes / {:.2} GiB freed)",
freed,
freed as f64 / (1u64 << 30) as f64,
);
return Ok(());
}
let repo = model.expect("validated above");
if let Some(q_str) = quant {
let q =
QuantType::from_canonical_str(&q_str).map_err(|e| anyhow::anyhow!("--quant: {}", e))?;
let freed = cache
.invalidate(&repo, q)
.with_context(|| format!("clear {}@{}", repo, q.as_str()))?;
println!(
"hf2q cache: cleared {}@{} ({} bytes freed)",
repo,
q.as_str(),
freed
);
} else {
let freed = cache
.invalidate_repo(&repo)
.with_context(|| format!("clear {} (all quants)", repo))?;
println!(
"hf2q cache: cleared {} (all quants — {} bytes freed)",
repo, freed
);
}
Ok(())
}
pub fn cmd_parity(args: cli::ParityArgs) -> Result<()> {
use cli::ParityCommand;
match args.command {
ParityCommand::Check {
model,
prompt,
min_prefix,
max_tokens,
self_baseline,
tq_quality,
fixture,
cosine_mean_floor,
cosine_p1_floor,
argmax_max,
ppl_delta_max,
} => {
if tq_quality {
if self_baseline {
anyhow::bail!(
"parity check --tq-quality is incompatible with \
--self-baseline (Gate D vs Gate H — different gates)"
);
}
let fixture = fixture.ok_or_else(|| {
anyhow::anyhow!(
"parity check --tq-quality requires --fixture \
<path/to/<prompt>_tq_quality.json>.\n\
Hint: the fixture is produced by `hf2q parity \
capture --tq-quality --model <gguf> --prompt \
{prompt}` (iter-112)."
)
})?;
parity_quality::cmd_parity_check_tq_quality(
&model,
&prompt,
&fixture,
cosine_mean_floor,
cosine_p1_floor,
argmax_max,
ppl_delta_max,
max_tokens,
)
} else {
let _ = (
fixture,
cosine_mean_floor,
cosine_p1_floor,
argmax_max,
ppl_delta_max,
);
cmd_parity_check(&model, &prompt, min_prefix, max_tokens, self_baseline)
}
}
ParityCommand::Capture {
model,
output,
prompt,
max_tokens,
tq_quality,
} => {
if tq_quality {
parity_quality::cmd_parity_capture_tq_quality(&model, &output, &prompt, max_tokens)
} else {
cmd_parity_capture(&model, &output, &prompt, max_tokens)
}
}
}
}
fn cmd_parity_check(
model_path: &Path,
prompt_name: &str,
min_prefix: Option<usize>,
max_tokens: Option<usize>,
self_baseline: bool,
) -> Result<()> {
let evals_dir = Path::new("tests/evals");
let ref_dir = evals_dir.join("reference");
let prompt_file = evals_dir.join("prompts").join(format!("{prompt_name}.txt"));
anyhow::ensure!(
prompt_file.exists(),
"Prompt file not found: {}",
prompt_file.display()
);
let prompt_text = std::fs::read_to_string(&prompt_file)?.trim().to_string();
let ref_suffix = if self_baseline { "_hf2q" } else { "_llama" };
let ref_file = ref_dir.join(format!("{prompt_name}{ref_suffix}.txt"));
anyhow::ensure!(
ref_file.exists(),
"Reference file not found: {}",
ref_file.display()
);
let ref_bytes = std::fs::read(&ref_file)?;
let manifest: serde_json::Value =
serde_json::from_str(&std::fs::read_to_string(ref_dir.join("MANIFEST.json"))?)?;
let prompt_meta = &manifest["prompts"][prompt_name];
let tokens =
max_tokens.unwrap_or_else(|| prompt_meta["max_tokens"].as_u64().unwrap_or(1000) as usize);
let threshold = min_prefix.unwrap_or_else(||
prompt_meta["parity_gate"].as_str()
.and_then(|s| s.split(">=").nth(1))
.and_then(|s| s.trim().parse::<usize>().ok())
.unwrap_or(0));
eprintln!("=== Parity Check: {} ===", prompt_name);
eprintln!("Model: {}", model_path.display());
eprintln!("Prompt: {} ({} chars)", prompt_name, prompt_text.len());
eprintln!("Tokens: {}", tokens);
eprintln!("Threshold: {} bytes", threshold);
eprintln!();
let tokenizer_path = find_tokenizer(model_path, None)?;
let mut ctx = gpu::GpuContext::new().map_err(|e| anyhow::anyhow!("GPU init: {e}"))?;
let gguf = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF open: {e}"))?;
let cfg = config::Gemma4Config::from_gguf(&gguf)?;
let mut parity_progress = header::LoadProgress::new(false, 1, 0);
let mut mlx_w = crate::inference::models::gemma4::MlxModelWeights::load_from_gguf(
&gguf,
&cfg,
&mut ctx,
&mut parity_progress,
)?;
let mut tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path)
.map_err(|e| anyhow::anyhow!("Tokenizer: {e}"))?;
tokenizer
.with_truncation(None)
.map_err(|e| anyhow::anyhow!("Tokenizer truncation: {e}"))?;
let rendered = render_chat_template(
&gguf,
&cli::GenerateArgs {
model: model_path.to_path_buf(),
prompt: Some(prompt_text.clone()),
prompt_file: None,
tokenizer: None,
config: None,
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: tokens,
mmproj: None,
image: None,
chat_template: None,
chat_template_file: None,
benchmark: false,
speculative: false,
kv_bits: None,
enable_thinking: false,
no_thinking: false,
ignore_eos: false,
},
Some(&tokenizer),
&prompt_text,
)?;
let encoding = tokenizer
.encode(rendered.as_str(), false)
.map_err(|e| anyhow::anyhow!("Tokenize: {e}"))?;
let prompt_tokens: Vec<u32> = encoding.get_ids().to_vec();
let eos_token_ids: Vec<u32> = vec![1, 106];
let first_token = mlx_w.forward_prefill(&prompt_tokens, tokens, &mut ctx)?;
let mut all_tokens = prompt_tokens.to_vec();
let mut next_token = first_token;
all_tokens.push(next_token);
for _ in 1..tokens {
if eos_token_ids.contains(&next_token) {
break;
}
let pos = all_tokens.len() - 1;
let mut p = None;
next_token = mlx_w.forward_decode(next_token, pos, &mut ctx, &mut p)?;
all_tokens.push(next_token);
}
let gen_tokens = &all_tokens[prompt_tokens.len()..];
let hf2q_text = tokenizer.decode(gen_tokens, false).unwrap_or_default();
let hf2q_bytes = hf2q_text.as_bytes();
let n = ref_bytes.len().min(hf2q_bytes.len());
let mut common = 0;
while common < n && ref_bytes[common] == hf2q_bytes[common] {
common += 1;
}
eprintln!();
let ref_label = if self_baseline {
"frozen hf2q"
} else {
"llama.cpp"
};
println!("Reference: {} bytes ({})", ref_bytes.len(), ref_label);
println!("hf2q: {} bytes", hf2q_bytes.len());
println!("Common: {} bytes", common);
if self_baseline {
let identical = hf2q_bytes.len() == ref_bytes.len() && common == ref_bytes.len();
if identical {
println!(
"PASS: byte-identical to frozen hf2q baseline ({} bytes)",
common
);
} else {
println!("FAIL: not byte-identical to frozen hf2q baseline");
if common < n {
let ctx_start = common;
let ctx_end = (common + 80).min(n);
let ref_snip = String::from_utf8_lossy(&ref_bytes[ctx_start..ctx_end]);
let hf2q_snip =
String::from_utf8_lossy(&hf2q_bytes[ctx_start..ctx_end.min(hf2q_bytes.len())]);
println!();
println!("Divergence at byte {}:", common);
println!(" frozen: {:?}", ref_snip);
println!(" hf2q: {:?}", hf2q_snip);
}
anyhow::bail!("Self-baseline check failed: hf2q differs from frozen baseline");
}
} else {
println!("Threshold: {} bytes", threshold);
if common >= threshold {
println!("PASS: {} >= {}", common, threshold);
if common > threshold {
println!(" ({} bytes above threshold)", common - threshold);
}
} else {
println!("FAIL: {} < {}", common, threshold);
if common < n {
let ctx_start = common;
let ctx_end = (common + 80).min(n);
let ref_snip = String::from_utf8_lossy(&ref_bytes[ctx_start..ctx_end]);
let hf2q_snip =
String::from_utf8_lossy(&hf2q_bytes[ctx_start..ctx_end.min(hf2q_bytes.len())]);
println!();
println!("Divergence at byte {}:", common);
println!(" llama: {:?}", ref_snip);
println!(" hf2q: {:?}", hf2q_snip);
}
anyhow::bail!("Parity check failed: {} < {}", common, threshold);
}
}
Ok(())
}
fn cmd_parity_capture(
model_path: &Path,
output_dir: &Path,
prompt_name: &str,
max_tokens: Option<usize>,
) -> Result<()> {
let evals_dir = Path::new("tests/evals");
let prompts: Vec<String> = if prompt_name == "all" {
vec![
"sourdough".into(),
"short_hello".into(),
"sliding_wrap".into(),
]
} else {
vec![prompt_name.to_string()]
};
let tokenizer_path = find_tokenizer(model_path, None)?;
let mut ctx = gpu::GpuContext::new().map_err(|e| anyhow::anyhow!("GPU init: {e}"))?;
let gguf = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF open: {e}"))?;
let cfg = config::Gemma4Config::from_gguf(&gguf)?;
let _gguf_preload = &gguf; let mut tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path)
.map_err(|e| anyhow::anyhow!("Tokenizer: {e}"))?;
tokenizer
.with_truncation(None)
.map_err(|e| anyhow::anyhow!("Tokenizer truncation: {e}"))?;
std::fs::create_dir_all(output_dir)?;
for pname in &prompts {
let prompt_file = evals_dir.join("prompts").join(format!("{pname}.txt"));
anyhow::ensure!(
prompt_file.exists(),
"Prompt not found: {}",
prompt_file.display()
);
let prompt_text = std::fs::read_to_string(&prompt_file)?.trim().to_string();
let tokens = max_tokens.unwrap_or(match pname.as_str() {
"sourdough" => 1000,
"short_hello" => 50,
"sliding_wrap" => 500,
_ => 200,
});
eprintln!("Capturing: {} ({} tokens)", pname, tokens);
let mut parity_progress = header::LoadProgress::new(false, 1, 0);
let mut mlx_w_fresh = crate::inference::models::gemma4::MlxModelWeights::load_from_gguf(
&gguf,
&cfg,
&mut ctx,
&mut parity_progress,
)?;
let rendered = render_chat_template(
&gguf,
&cli::GenerateArgs {
model: model_path.to_path_buf(),
prompt: Some(prompt_text.clone()),
prompt_file: None,
tokenizer: None,
config: None,
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: tokens,
mmproj: None,
image: None,
chat_template: None,
chat_template_file: None,
benchmark: false,
speculative: false,
kv_bits: None,
enable_thinking: false,
no_thinking: false,
ignore_eos: false,
},
Some(&tokenizer),
&prompt_text,
)?;
let encoding = tokenizer
.encode(rendered.as_str(), false)
.map_err(|e| anyhow::anyhow!("Tokenize: {e}"))?;
let prompt_tokens: Vec<u32> = encoding.get_ids().to_vec();
let eos_token_ids: Vec<u32> = vec![1, 106];
let first_token = mlx_w_fresh.forward_prefill(&prompt_tokens, tokens, &mut ctx)?;
let mut all_tokens = prompt_tokens.to_vec();
let mut next_token = first_token;
all_tokens.push(next_token);
for _ in 1..tokens {
if eos_token_ids.contains(&next_token) {
break;
}
let pos = all_tokens.len() - 1;
let mut p = None;
next_token = mlx_w_fresh.forward_decode(next_token, pos, &mut ctx, &mut p)?;
all_tokens.push(next_token);
}
let gen_tokens = &all_tokens[prompt_tokens.len()..];
let text = tokenizer.decode(gen_tokens, false).unwrap_or_default();
let out_path = output_dir.join(format!("{pname}_hf2q.txt"));
std::fs::write(&out_path, &text)?;
eprintln!(" Wrote {} bytes to {}", text.len(), out_path.display());
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::{
build_chat_template_env, detect_greedy_repetition_loop,
detect_greedy_repetition_loop_with_text, find_special_token_stop,
llama_cpp_special_token_id_for_model, maybe_print_serve_banner, parse_scheduler_config,
render_jinja_template, resolve_enable_thinking, run_decode_loop, should_enable_kv_persist,
DecodeStopReason, RaisePolicy, DEFAULT_MAX_SLOTS_UNDER_INFLIGHT,
FALLBACK_GEMMA4_API_CHAT_TEMPLATE, FALLBACK_GEMMA4_CHAT_TEMPLATE,
};
use crate::cli;
use crate::core::chat_templates::QWEN3_CHATML;
use crate::core::provenance::Provenance;
use crate::serve::load_info::{ArchFamily, ChatTemplateSource, LoadInfo, TokenizerSource};
use std::path::PathBuf;
use std::time::Duration;
#[test]
fn llama_cpp_gpt2_special_token_defaults_match_vocab_cpp() {
assert_eq!(
llama_cpp_special_token_id_for_model("gpt2", "tokenizer.ggml.bos_token_id"),
Some(11)
);
assert_eq!(
llama_cpp_special_token_id_for_model("gpt2", "tokenizer.ggml.eos_token_id"),
Some(11)
);
assert_eq!(
llama_cpp_special_token_id_for_model("gpt2", "tokenizer.ggml.padding_token_id"),
None
);
}
fn synthetic_serve_banner_info() -> LoadInfo {
LoadInfo {
model_id: "serve-test-model".to_string(),
arch_str: "gemma4".to_string(),
arch_family: ArchFamily::Gemma4,
model_path: PathBuf::from("/tmp/serve-test-model.gguf"),
on_disk_bytes: 1024,
backend_chip: "Apple M5 Max".to_string(),
backend: "mlx-native",
n_layers: 2,
hidden_size: 32,
vocab_size: 128,
n_attention_heads: 4,
n_key_value_heads: 2,
head_dim: 8,
sliding_window: Some(16),
full_attention_interval: None,
max_context_length: Some(128),
moe: None,
quant_label: Some("Q4_K".to_string()),
quant_bpw: Some(4.5),
tokenizer_source: TokenizerSource::GgufEmbedded,
eos_token_ids: vec![1],
bos_token_id: Some(2),
chat_template_source: ChatTemplateSource::GgufEmbedded,
provenance: Provenance::External,
vision_projector: None,
load_wall_clock: Duration::from_millis(25),
resident_weight_bytes: None,
kv_cache_budget_bytes: None,
kv_spill_active: false,
tq_kv_active: false,
kv_bytes_per_token_override: None,
}
}
#[test]
fn cmd_serve_banner_emits_on_tty() {
let info = synthetic_serve_banner_info();
let mut buf = Vec::new();
maybe_print_serve_banner(&info, &mut buf, true, false).expect("print serve banner");
let got = String::from_utf8(buf).expect("utf8");
assert_eq!(got.lines().count(), 14);
assert!(got.contains("hf2q load: model = serve-test-model"));
assert!(got.contains("\x1b[2m"));
}
#[test]
fn cmd_serve_banner_silent_when_non_tty() {
let info = synthetic_serve_banner_info();
let mut buf = Vec::new();
maybe_print_serve_banner(&info, &mut buf, false, false).expect("skip serve banner");
assert!(buf.is_empty());
}
#[test]
fn cmd_serve_banner_silent_when_quiet() {
let info = synthetic_serve_banner_info();
let mut buf = Vec::new();
maybe_print_serve_banner(&info, &mut buf, true, true).expect("skip quiet serve banner");
assert!(buf.is_empty());
}
#[test]
fn hf2q_kv_persist_zero_overrides_flag() {
assert!(!should_enable_kv_persist(Some("/path"), Some("0")));
}
#[test]
fn hf2q_kv_persist_unset_respects_flag_present() {
assert!(should_enable_kv_persist(Some("/path"), None));
}
#[test]
fn hf2q_kv_persist_one_respects_flag() {
assert!(should_enable_kv_persist(Some("/path"), Some("1")));
}
#[test]
fn hf2q_kv_persist_no_flag_means_disabled_regardless() {
assert!(!should_enable_kv_persist(None, None));
assert!(!should_enable_kv_persist(None, Some("0")));
assert!(!should_enable_kv_persist(None, Some("1")));
}
#[test]
fn hf2q_kv_persist_zero_with_whitespace_still_disables() {
assert!(!should_enable_kv_persist(Some("/path"), Some(" 0 ")));
assert!(!should_enable_kv_persist(Some("/path"), Some("0\n")));
}
#[test]
fn hf2q_kv_persist_empty_or_other_values_respect_flag() {
assert!(should_enable_kv_persist(Some("/path"), Some("")));
assert!(should_enable_kv_persist(Some("/path"), Some("true")));
assert!(should_enable_kv_persist(Some("/path"), Some("yes")));
assert!(should_enable_kv_persist(Some("/path"), Some("00")));
}
#[test]
fn iter219b_cli_fallback_chat_template_matches_iter217_contract() {
assert!(
!FALLBACK_GEMMA4_CHAT_TEMPLATE.contains("<|think|>"),
"CLI fallback MUST NOT contain `<|think|>` (activates thinking-mode \
that diverges from llama.cpp parity gate). Got: {FALLBACK_GEMMA4_CHAT_TEMPLATE:?}"
);
assert!(
FALLBACK_GEMMA4_CHAT_TEMPLATE.contains("<|channel>thought\n<channel|>"),
"CLI fallback MUST end with empty `<|channel>thought\\n<channel|>` block \
(closes thinking-mode pre-content; mirrors iter-217 API-path fix). \
Got: {FALLBACK_GEMMA4_CHAT_TEMPLATE:?}"
);
assert!(
FALLBACK_GEMMA4_CHAT_TEMPLATE.contains("{{PROMPT}}"),
"CLI fallback MUST keep `{{{{PROMPT}}}}` placeholder for `String::replace` \
rendering. Got: {FALLBACK_GEMMA4_CHAT_TEMPLATE:?}"
);
assert!(
!FALLBACK_GEMMA4_API_CHAT_TEMPLATE.contains("<|think|>"),
"API fallback regression — `<|think|>` reintroduced after iter-217 fix"
);
assert!(
FALLBACK_GEMMA4_API_CHAT_TEMPLATE.contains("<|channel>thought\n<channel|>"),
"API fallback regression — empty channel block missing"
);
}
#[test]
fn jinja_template_renders_single_user_turn() {
let tmpl = "{{ bos_token }}{% for m in messages %}<|turn|>{{ m.role }}\n{{ m.content }}<|end|>\n{% endfor %}{% if add_generation_prompt %}<|turn|>model\n{% endif %}";
let out = render_jinja_template(tmpl, "hello", None).expect("render ok");
assert!(
out.starts_with("<bos>"),
"output should start with bos_token: {out}"
);
assert!(
out.contains("<|turn|>user\nhello<|end|>"),
"user turn missing: {out}"
);
assert!(
out.ends_with("<|turn|>model\n"),
"generation prompt missing: {out}"
);
}
#[test]
fn jinja_template_parse_error_is_reported() {
let tmpl = "{% unclosed"; let err = render_jinja_template(tmpl, "x", None).unwrap_err();
let msg = format!("{err:#}");
assert!(
msg.contains("parse") || msg.contains("Jinja") || msg.contains("template"),
"expected parse error, got: {msg}"
);
}
fn write_minimal_gguf_with_arch(arch: &str) -> tempfile::NamedTempFile {
use std::io::Write;
let mut f = tempfile::Builder::new()
.suffix(".gguf")
.tempfile()
.expect("tempfile");
let mut buf: Vec<u8> = Vec::new();
buf.extend_from_slice(b"GGUF");
buf.extend_from_slice(&3u32.to_le_bytes()); buf.extend_from_slice(&0u64.to_le_bytes()); buf.extend_from_slice(&1u64.to_le_bytes()); let key = b"general.architecture";
buf.extend_from_slice(&(key.len() as u64).to_le_bytes());
buf.extend_from_slice(key);
buf.extend_from_slice(&8u32.to_le_bytes()); let val = arch.as_bytes();
buf.extend_from_slice(&(val.len() as u64).to_le_bytes());
buf.extend_from_slice(val);
f.write_all(&buf).expect("write");
f.flush().expect("flush");
f
}
#[test]
fn load_engine_routes_qwen35_to_qwen35_loaded_model() {
let tmp = write_minimal_gguf_with_arch("qwen35");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(result.is_err(), "0-tensor synthetic GGUF must fail load");
let msg = format!("{:#}", result.err().unwrap());
assert!(
!msg.contains("Phase 4 follow-up: SERVE-side load for general.architecture"),
"wedge-1 bail must not fire post iter-215; got: {msg}"
);
assert!(
!msg.contains("hf2q generate") || !msg.contains("cmd_generate_qwen35"),
"load_engine error must not include the iter-214 wedge-1 workaround \
pointer; got: {msg}"
);
}
#[test]
fn load_engine_routes_qwen35moe_to_qwen35_loaded_model() {
let tmp = write_minimal_gguf_with_arch("qwen35moe");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(result.is_err(), "0-tensor synthetic GGUF must fail load");
let msg = format!("{:#}", result.err().unwrap());
assert!(
!msg.contains("Phase 4 follow-up: SERVE-side load for general.architecture"),
"wedge-1 bail must not fire post iter-215; got: {msg}"
);
}
#[test]
fn iter228a_load_engine_routes_qwen3_vl_to_qwen3vl_text_loader() {
let tmp = write_minimal_gguf_with_arch("qwen3_vl");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(
result.is_err(),
"0-tensor synthetic qwen3_vl GGUF must fail load (no metadata)"
);
let msg = format!("{:#}", result.err().unwrap());
assert!(
msg.contains("Qwen3VlTextConfig::from_gguf")
|| msg.contains("missing core architecture facts"),
"iter-228a must route dense Qwen3-VL through Qwen3VlTextConfig parser; got: {msg}"
);
assert!(
!msg.contains("missing blk.0.ffn_gate_up_exps.weight"),
"iter-228a must intercept BEFORE the Gemma MoE expert load; got: {msg}"
);
assert!(
!msg.contains("iter-227 closes only the dispatch gap"),
"iter-228a lifted the iter-227 dispatch bail for dense Qwen3-VL; got: {msg}"
);
}
#[test]
fn iter228a_load_engine_routes_qwen3vl_upstream_to_qwen3vl_text_loader() {
let tmp = write_minimal_gguf_with_arch("qwen3vl");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(result.is_err());
let msg = format!("{:#}", result.err().unwrap());
assert!(
msg.contains("Qwen3VlTextConfig::from_gguf")
|| msg.contains("missing core architecture facts"),
"iter-228a must route upstream-arch dense Qwen3-VL through Qwen3VlTextConfig parser; got: {msg}"
);
}
#[test]
fn iter227_load_engine_rejects_qwen3vlmoe_with_moe_specific_error() {
let tmp = write_minimal_gguf_with_arch("qwen3vlmoe");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(result.is_err());
let msg = format!("{:#}", result.err().unwrap());
assert!(
msg.contains("Qwen3-VL") && msg.contains("MoE"),
"iter-227 MoE-variant error must include 'MoE' to distinguish from dense; got: {msg}"
);
}
#[test]
fn load_engine_rejects_unknown_arch_without_gemma_fallback() {
let tmp = write_minimal_gguf_with_arch("totally-fake-arch-name");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(result.is_err(), "unknown architecture must fail dispatch");
let msg = format!("{:#}", result.err().unwrap());
assert!(
msg.contains("unsupported GGUF general.architecture")
&& msg.contains("totally-fake-arch-name"),
"unknown architecture error must identify the rejected value; got: {msg}"
);
assert!(
!msg.contains("missing blk.0.ffn_gate_up_exps.weight"),
"unknown architecture must not reach the Gemma loader; got: {msg}"
);
}
#[test]
fn load_engine_routes_deepseek4_to_native_loader_without_gemma_fallback() {
let tmp = write_minimal_gguf_with_arch("deepseek4");
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(
result.is_err(),
"the minimal DeepSeek-V4 fixture must fail validation"
);
let msg = format!("{:#}", result.err().unwrap());
assert!(
msg.contains("DeepSeek-V4 tokenizer") || msg.contains("native DeepSeek-V4 model"),
"DeepSeek-V4 must reach its native loader; got: {msg}"
);
assert!(
!msg.contains("refusing to route it through Gemma")
&& !msg.contains("missing blk.0.ffn_gate_up_exps.weight"),
"DeepSeek-V4 must not reach the Gemma loader; got: {msg}"
);
}
#[test]
fn iter227_does_not_regress_qwen35_dispatch() {
for arch in &["qwen35", "qwen35moe"] {
let tmp = write_minimal_gguf_with_arch(arch);
let cfg = super::multi_model::EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: 4,
warmup_synchronously: false,
kv_metrics_sink: None,
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let result = super::load_engine(tmp.path(), &cfg);
assert!(
result.is_err(),
"0-tensor synthetic GGUF for {arch} must fail load (routed to qwen35 path)"
);
let msg = format!("{:#}", result.err().unwrap());
assert!(
!msg.contains("Qwen3-VL"),
"iter-227 Qwen3-VL dispatch must NOT fire on arch={arch}; got: {msg}"
);
}
}
#[test]
fn detect_repetition_returns_none_below_window_size() {
let toks: Vec<u32> = (0..50).collect();
assert_eq!(detect_greedy_repetition_loop(&toks), None);
}
#[test]
fn detect_repetition_returns_none_for_diverse_tokens() {
let toks: Vec<u32> = (0..200).collect();
assert_eq!(detect_greedy_repetition_loop(&toks), None);
}
#[test]
fn detect_repetition_finds_8_token_cycle() {
let cycle: [u32; 8] = [101, 102, 103, 104, 105, 106, 107, 108];
let mut toks: Vec<u32> = (0..50).collect();
for _ in 0..20 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(result.is_some(), "8-token cycle should be detected");
let (ngram, occ) = result.unwrap();
assert_eq!(ngram, 8, "should detect at the smallest matching size");
assert!(
occ >= 3,
"should report ≥3 occurrences in a saturated window, got {occ}"
);
}
#[test]
fn detect_repetition_finds_2_token_cycle_single_token_loop() {
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..10 {
toks.push(777);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"single-token loop should be detected (was missed by prior detector)"
);
let (ngram, _) = result.unwrap();
assert_eq!(ngram, 2, "should detect at the smallest matching size");
}
#[test]
fn detect_repetition_finds_7_token_cycle() {
let cycle: [u32; 7] = [501, 502, 503, 504, 505, 506, 507];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..10 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"7-token cycle should be detected (was missed by prior detector)"
);
assert_eq!(result.unwrap().0, 7);
}
#[test]
fn detect_repetition_finds_11_token_cycle() {
let cycle: [u32; 11] = [601, 602, 603, 604, 605, 606, 607, 608, 609, 610, 611];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..10 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"11-token cycle should be detected (was missed by prior detector)"
);
assert_eq!(result.unwrap().0, 11);
}
#[test]
fn detect_repetition_finds_18_token_cycle() {
let cycle: [u32; 18] = [
701, 702, 703, 704, 705, 706, 707, 708, 709, 710, 711, 712, 713, 714, 715, 716, 717,
718,
];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..8 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"18-token cycle should be detected (was missed by prior detector)"
);
assert_eq!(result.unwrap().0, 18);
}
#[test]
fn detect_repetition_finds_22_token_cycle() {
let cycle: [u32; 22] = [
801, 802, 803, 804, 805, 806, 807, 808, 809, 810, 811, 812, 813, 814, 815, 816, 817,
818, 819, 820, 821, 822,
];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..6 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"22-token cycle should be detected (was missed by prior detector)"
);
assert_eq!(result.unwrap().0, 22);
}
#[test]
fn detect_repetition_no_false_positive_on_two_repeats() {
let cycle: [u32; 8] = [901, 902, 903, 904, 905, 906, 907, 908];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..2 {
toks.extend_from_slice(&cycle);
}
assert_eq!(detect_greedy_repetition_loop(&toks), None);
}
#[test]
fn detect_repetition_finds_16_token_cycle() {
let cycle: [u32; 16] = [
201, 202, 203, 204, 205, 206, 207, 208, 209, 210, 211, 212, 213, 214, 215, 216,
];
let mut toks: Vec<u32> = (0..50).collect();
for _ in 0..10 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(result.is_some(), "16-token cycle should be detected");
}
#[test]
fn detect_repetition_smart_skips_structural_cycle_at_3_reps() {
let cycle: [u32; 4] = [101, 102, 103, 104];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..3 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop_with_text(&toks, |_cycle| {
"| :--- ".to_string()
});
assert!(
result.is_none(),
"structural cycle (alpha_ratio=0) at only 3 reps should NOT fire"
);
let mut toks2: Vec<u32> = (0..30).collect();
for _ in 0..8 {
toks2.extend_from_slice(&cycle);
}
let result2 =
detect_greedy_repetition_loop_with_text(&toks2, |_cycle| "| :--- ".to_string());
assert!(
result2.is_some(),
"structural cycle at 8 reps should fire (genuine pad-loop)"
);
let (_ng, reps) = result2.unwrap();
assert!(reps >= 8, "structural threshold should be ≥ 8 reps");
}
#[test]
fn detect_repetition_smart_fires_on_content_cycle_at_3_reps() {
let cycle: [u32; 5] = [201, 202, 203, 204, 205];
let mut toks: Vec<u32> = (0..30).collect();
for _ in 0..3 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop_with_text(&toks, |_cycle| {
"Miticidal Agent, ".to_string()
});
assert!(
result.is_some(),
"alphanumeric content cycle at 3 reps MUST fire"
);
let (_ng, reps) = result.unwrap();
assert_eq!(reps, 3, "content threshold should be 3 reps");
}
#[test]
fn detect_repetition_tolerates_minor_drift() {
let cycle: [u32; 12] = [301, 302, 303, 304, 305, 306, 307, 308, 309, 310, 311, 312];
let outlier: u32 = 999;
let mut toks: Vec<u32> = (0..30).collect();
for i in 0..10 {
if i == 4 {
let mut drifted = cycle.to_vec();
drifted[6] = outlier;
toks.extend_from_slice(&drifted);
} else {
toks.extend_from_slice(&cycle);
}
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"minor-drift cycle should still detect (non-consecutive matching), got None"
);
}
#[test]
fn detect_repetition_does_not_overcount_with_overlapping_increment() {
let toks: Vec<u32> = vec![0; 200];
let result = detect_greedy_repetition_loop(&toks);
assert!(result.is_some());
let (ngram, occ) = result.unwrap();
assert!(
occ.saturating_mul(ngram) <= toks.len(),
"non-overlap invariant violated: ngram={ngram} occ={occ}, occ*ngram={} > n={}",
occ.saturating_mul(ngram),
toks.len()
);
}
#[test]
fn detect_repetition_does_not_fire_on_bullet_list() {
let bullet_prefix: [u32; 8] = [42, 32, 32, 32, 1234, 5678, 9012, 3456];
let mut toks: Vec<u32> = (0..50).collect();
for content_seed in 0..6u32 {
toks.extend_from_slice(&bullet_prefix);
for j in 0..(12 + content_seed as usize) {
toks.push(50000 + content_seed * 100 + j as u32);
}
}
assert_eq!(
detect_greedy_repetition_loop(&toks),
None,
"false positive on legitimate bullet-list enumeration"
);
}
#[test]
fn detect_repetition_fires_on_true_consecutive_loop() {
let cycle: [u32; 16] = [
701, 702, 703, 704, 705, 706, 707, 708, 709, 710, 711, 712, 713, 714, 715, 716,
];
let mut toks: Vec<u32> = (0..50).collect();
for _ in 0..30 {
toks.extend_from_slice(&cycle);
}
let result = detect_greedy_repetition_loop(&toks);
assert!(
result.is_some(),
"true consecutive loop must still trigger the detector"
);
let (ngram, _occ) = result.unwrap();
assert_eq!(
ngram, 16,
"16-token cycle's halves are NOT individually periodic at 8, so \
ngram=8 doesn't fire; the smallest matching consecutive period \
is 16"
);
}
#[test]
fn special_token_stop_empty_input() {
assert_eq!(find_special_token_stop(""), None);
}
#[test]
fn special_token_stop_plain_thinking_text_does_not_stop() {
let s = "I should make sure the answer is comprehensive and avoids \
jargon. The user is asking about candy.";
assert_eq!(find_special_token_stop(s), None);
}
#[test]
fn special_token_stop_think_html_tag_does_not_stop() {
let s = "<think>reasoning content</think>\n\nactual answer";
assert_eq!(find_special_token_stop(s), None);
}
#[test]
fn special_token_stop_im_end_alone_is_not_a_stop() {
let s = "answer text<|im_end|>";
assert_eq!(find_special_token_stop(s), None);
}
#[test]
fn special_token_stop_im_start_user_signals_turn_end() {
let s = "[thinking]<|im_end|>\n<|im_start|>assistant\n[answer]<|im_end|>\n<|im_start|>user\nnext";
assert_eq!(find_special_token_stop(s), Some("<|im_start|>user"));
}
#[test]
fn special_token_stop_endoftext_fragment_stops() {
let s = "answer text<|endoftext|>";
assert_eq!(find_special_token_stop(s), Some("<|endoftext|>"));
}
#[test]
fn special_token_stop_end_marker_before_im_start_stops_at_end() {
let s = "thinking content<|end|>\n\n<|im_start|>user\necho";
assert_eq!(find_special_token_stop(s), Some("<|end|>"));
}
#[test]
fn special_token_stop_im_end_then_assistant_role_no_user_does_not_stop() {
let s = "answer<|im_end|>\n<|im_start|>assistant\nleak";
assert_eq!(find_special_token_stop(s), None);
}
#[test]
fn special_token_stop_phi3_end_marker_stops() {
let s = "answer<|end|>";
assert_eq!(find_special_token_stop(s), Some("<|end|>"));
}
#[test]
fn special_token_stop_end_substring_does_not_falsely_match_endoftext() {
let s = "answer<|endoftext|>";
assert_eq!(find_special_token_stop(s), Some("<|endoftext|>"));
}
#[test]
fn special_token_stop_degenerate_im_start_only_does_not_stop() {
let s = "<|im_start|>\n\n<|im_start|>";
assert_eq!(find_special_token_stop(s), None);
}
#[test]
fn special_token_stop_qwen3_thinking_clean_answer_does_not_stop_mid_text() {
let s = "<think>The user wants candy.</think>\n\n\
To make candy you need sugar, water, and flavoring. \
Heat to soft-ball stage at 235°F.";
assert_eq!(find_special_token_stop(s), None);
}
#[test]
fn special_token_stop_earliest_marker_wins() {
let s = "x<|end|>y<|im_end|>z";
assert_eq!(find_special_token_stop(s), Some("<|end|>"));
}
#[test]
fn render_jinja_qwen3_chatml_enable_thinking_true_opens_unfilled_think_block() {
let rendered = render_jinja_template(QWEN3_CHATML, "make me candy", Some(true))
.expect("qwen3-chatml render with enable_thinking=true must succeed");
assert!(
rendered.ends_with("<|im_start|>assistant\n<think>\n"),
"with enable_thinking=true the qwen3-chatml `else` branch (line 152) \
must fire, ending the prompt with `<|im_start|>assistant\\n<think>\\n` \
so the model emits reasoning content next. Actual tail: \
`{}`",
&rendered[rendered.len().saturating_sub(60)..]
);
}
#[test]
fn render_jinja_qwen3_chatml_enable_thinking_false_emits_pre_closed_think_block() {
let rendered = render_jinja_template(QWEN3_CHATML, "make me candy", Some(false))
.expect("qwen3-chatml render with enable_thinking=false must succeed");
assert!(
rendered.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"),
"with enable_thinking=false the qwen3-chatml `if` branch (line 150) \
must fire, emitting a pre-closed `<think>\\n\\n</think>\\n\\n` block \
so the model skips reasoning and emits the answer directly. Actual \
tail: `{}`",
&rendered[rendered.len().saturating_sub(60)..]
);
}
#[test]
fn render_jinja_qwen3_chatml_enable_thinking_none_treats_as_undefined_else_branch() {
let rendered_none = render_jinja_template(QWEN3_CHATML, "make me candy", None)
.expect("qwen3-chatml render with None must succeed");
let rendered_true = render_jinja_template(QWEN3_CHATML, "make me candy", Some(true))
.expect("qwen3-chatml render with Some(true) must succeed");
assert_eq!(
rendered_none, rendered_true,
"with the qwen3-chatml template, None must render byte-identically \
to Some(true) — the `is defined and is false` guard fails for both. \
A divergence here means minijinja serialized None as a non-undefined \
value, which would break the lower-level helper contract."
);
}
fn test_args_default() -> cli::GenerateArgs {
cli::GenerateArgs {
model: std::path::PathBuf::from("/dev/null"),
prompt: None,
prompt_file: None,
tokenizer: None,
config: None,
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
mmproj: None,
image: None,
chat_template: None,
chat_template_file: None,
benchmark: false,
speculative: false,
kv_bits: None,
enable_thinking: false,
no_thinking: false,
ignore_eos: false,
}
}
#[test]
fn iter229_parity_tojson_non_finite_renders_null() {
let mut env = build_chat_template_env(RaisePolicy::Lenient);
env.add_template("t", "{{ x | tojson }}").unwrap();
let out = env
.get_template("t")
.unwrap()
.render(minijinja::context! { x => f64::NAN })
.unwrap();
assert_eq!(out, "null");
}
#[test]
fn iter229_parity_pycompat_strip_resolves() {
let mut env = build_chat_template_env(RaisePolicy::Lenient);
env.add_template("t", "{{ ' x '.strip() }}").unwrap();
let out = env.get_template("t").unwrap().render(()).unwrap();
assert_eq!(out, "x");
}
#[test]
fn iter229_parity_raise_exception_lenient_warns_and_continues() {
let mut env = build_chat_template_env(RaisePolicy::Lenient);
env.add_template("t", "{{ raise_exception('boom') }}x")
.unwrap();
let out = env.get_template("t").unwrap().render(()).unwrap();
assert_eq!(out, "x", "Lenient policy must render through the raise");
}
#[test]
fn iter229_parity_raise_exception_strict_hard_errors_with_message() {
let mut env = build_chat_template_env(RaisePolicy::Strict);
env.add_template("t", "{{ raise_exception('boom') }}x")
.unwrap();
let err = env.get_template("t").unwrap().render(()).unwrap_err();
assert!(err.to_string().contains("boom"), "err={err}");
}
use super::template_supports_enable_thinking;
#[test]
fn template_supports_enable_thinking_true_for_qwen3_chatml() {
assert!(
template_supports_enable_thinking(QWEN3_CHATML),
"qwen3-chatml.jinja MUST be detected as thinking-capable; \
render-and-diff against vendor lines 147-153 should differ"
);
}
#[test]
fn template_supports_enable_thinking_false_for_thinking_agnostic_template() {
let template = "{% for m in messages %}{{ m.role }}: {{ m.content }}\n\
{% endfor %}{% if add_generation_prompt %}assistant: {% endif %}";
assert!(
!template_supports_enable_thinking(template),
"thinking-agnostic template MUST NOT be flagged"
);
}
#[test]
fn template_supports_enable_thinking_false_for_plain_qwen3_suppressor_template() {
let template = "{%- for m in messages -%}\
<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n\
{%- endfor -%}\
{%- if add_generation_prompt -%}\
<|im_start|>assistant\n\
{%- if enable_thinking is defined and enable_thinking is false -%}\
<think>\n\n</think>\n\n\
{%- endif -%}\
{%- endif -%}";
assert!(
!template_supports_enable_thinking(template),
"a disabled-only empty think suppressor must not auto-enable thinking"
);
}
#[test]
fn template_supports_enable_thinking_false_for_gemma4_ara_channel_suppressor() {
let template = "{%- for m in messages -%}\
<|turn>{{ m.role }}\n{{ m.content }}<turn|>\n\
{%- endfor -%}\
{%- if add_generation_prompt -%}\
<|turn>model\n\
{%- if not (enable_thinking | default(false)) -%}\
<|channel>thought\n<channel|>\
{%- endif -%}\
{%- endif -%}";
assert!(
!template_supports_enable_thinking(template),
"gemma4-ara `<|channel>thought<channel|>` suppressor must not \
auto-enable thinking — section markers are not `<think>` opens"
);
}
#[test]
fn template_supports_enable_thinking_false_for_empty_branch_template() {
let template = "{% if enable_thinking %}{% endif %}\
{% for m in messages %}{{ m.content }}{% endfor %}";
assert!(
!template_supports_enable_thinking(template),
"empty-branch template MUST NOT be flagged (output bytes \
are identical for true vs false)"
);
}
#[test]
fn template_supports_enable_thinking_false_on_malformed_template() {
let template = "{% unclosed";
assert!(
!template_supports_enable_thinking(template),
"malformed template must produce false (graceful), not panic"
);
}
#[test]
fn resolve_enable_thinking_default_no_template_returns_some_false() {
let args = test_args_default();
assert_eq!(resolve_enable_thinking(&args, None), Some(false));
}
#[test]
fn resolve_enable_thinking_default_thinking_template_auto_enables() {
let args = test_args_default();
assert_eq!(
resolve_enable_thinking(&args, Some(QWEN3_CHATML)),
Some(true),
"qwen3-chatml.jinja default render MUST auto-enable thinking"
);
}
#[test]
fn resolve_enable_thinking_default_non_thinking_template_returns_some_false() {
let template = "{% for m in messages %}<|im_start|>{{ m.role }}\n\
{{ m.content }}<|im_end|>\n{% endfor %}\
{% if add_generation_prompt %}<|im_start|>assistant\n{% endif %}";
let args = test_args_default();
assert_eq!(resolve_enable_thinking(&args, Some(template)), Some(false));
}
#[test]
fn resolve_enable_thinking_explicit_enable_overrides_auto_detect() {
let mut args = test_args_default();
args.enable_thinking = true;
let plain = "{{ messages[0].content }}";
assert_eq!(resolve_enable_thinking(&args, Some(plain)), Some(true));
}
#[test]
fn resolve_enable_thinking_explicit_no_thinking_overrides_auto_detect() {
let mut args = test_args_default();
args.no_thinking = true;
assert_eq!(
resolve_enable_thinking(&args, Some(QWEN3_CHATML)),
Some(false)
);
}
#[test]
fn user_candy_regression_default_non_thinking_template_renders_pre_closed() {
let mut args = test_args_default();
args.no_thinking = true;
let resolved = resolve_enable_thinking(&args, Some(QWEN3_CHATML));
let rendered = render_jinja_template(QWEN3_CHATML, "make me candy", resolved)
.expect("--no-thinking render must succeed");
assert!(
rendered.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"),
"--no-thinking on qwen3-chatml MUST emit pre-closed think block. \
Actual tail: `{}`",
&rendered[rendered.len().saturating_sub(80)..]
);
}
#[test]
fn user_we_want_both_default_on_qwen3_chatml_renders_open_block() {
let args = test_args_default();
let resolved = resolve_enable_thinking(&args, Some(QWEN3_CHATML));
assert_eq!(resolved, Some(true), "auto-detect must enable thinking");
let rendered = render_jinja_template(QWEN3_CHATML, "make me candy", resolved)
.expect("default render must succeed");
assert!(
rendered.ends_with("<|im_start|>assistant\n<think>\n"),
"default cmdline on qwen3-chatml MUST emit open think block. \
Actual tail: `{}`",
&rendered[rendered.len().saturating_sub(80)..]
);
}
#[test]
fn render_jinja_qwen3_chatml_enable_thinking_does_not_affect_user_message() {
for et in [None, Some(true), Some(false)] {
let rendered = render_jinja_template(QWEN3_CHATML, "MARKER_USER_BODY_42", et)
.expect("render must succeed");
assert!(
rendered.contains("MARKER_USER_BODY_42"),
"user message body missing from rendered prompt with \
enable_thinking={et:?}"
);
assert!(
rendered.contains("<|im_start|>user\n"),
"user-turn opener missing from rendered prompt with \
enable_thinking={et:?}"
);
}
}
fn canned_step_model(tokens: Vec<u32>) -> impl FnMut(u32, i32, &[u32]) -> super::Result<u32> {
let mut iter = tokens.into_iter();
move |_prev, _pos, _generated| {
let t = iter.next().expect("test bug: canned stream exhausted");
Ok(t)
}
}
fn ascii_decode(toks: &[u32]) -> String {
toks.iter().map(|&t| char::from(t as u8)).collect()
}
#[test]
fn run_decode_loop_max_tokens_reached() {
let outcome = run_decode_loop(
b'A' as u32,
10,
1024,
5,
&[999],
canned_step_model(vec![b'B' as u32, b'C' as u32, b'D' as u32, b'E' as u32]),
ascii_decode,
|_event| {},
)
.expect("step_model never errors here");
assert_eq!(outcome.stop_reason, DecodeStopReason::MaxTokensReached);
assert_eq!(outcome.generated, 5);
}
#[test]
fn run_decode_loop_stops_on_integer_eos() {
let outcome = run_decode_loop(
b'A' as u32,
10,
1024,
10,
&[42],
canned_step_model(vec![b'B' as u32, 42, b'D' as u32]),
ascii_decode,
|_event| {},
)
.expect("ok");
assert_eq!(outcome.stop_reason, DecodeStopReason::EosTokenId(42));
assert_eq!(outcome.generated, 3);
}
#[test]
fn run_decode_loop_stops_on_special_token_leak() {
let mut step = 0;
let outcome = run_decode_loop(
b'A' as u32,
10,
1024,
10,
&[999],
canned_step_model(vec![b'B' as u32, b'C' as u32, b'D' as u32]),
move |toks| {
step += 1;
let mut s: String = toks.iter().map(|&t| char::from(t as u8)).collect();
if step >= 3 {
s.push_str("<|im_start|>user");
}
s
},
|_event| {},
)
.expect("ok");
assert_eq!(
outcome.stop_reason,
DecodeStopReason::SpecialTokenLeak("<|im_start|>user"),
);
}
#[test]
fn run_decode_loop_streams_raw_deltas_then_stops_on_full_marker() {
let mut deltas: Vec<String> = Vec::new();
const TARGET: &str = "<|im_start|>user";
let outcome = run_decode_loop(
b'<' as u32,
10,
1024,
32,
&[999],
canned_step_model(TARGET.bytes().skip(1).map(|b| b as u32).collect()),
ascii_decode,
|event| deltas.push(event.delta.to_string()),
)
.expect("ok");
assert_eq!(
outcome.stop_reason,
DecodeStopReason::SpecialTokenLeak("<|im_start|>user")
);
let joined: String = deltas.concat();
assert_eq!(joined, TARGET);
}
#[test]
fn run_decode_loop_stops_on_repetition() {
let cycle: Vec<u32> = (b'A'..=b'H').map(|c| c as u32).collect();
let mut stream: Vec<u32> = Vec::new();
for _ in 0..40 {
stream.extend(cycle.iter().copied());
}
let outcome = run_decode_loop(
b'Z' as u32,
10,
1024,
500,
&[999],
canned_step_model(stream),
ascii_decode,
|_event| {},
)
.expect("ok");
match outcome.stop_reason {
DecodeStopReason::RepetitionLoop { ngram, repeats } => {
assert!(ngram >= 8, "expected ngram >= 8, got {ngram}");
assert!(repeats >= 3, "expected repeats >= 3, got {repeats}");
}
other => panic!("expected RepetitionLoop, got {other:?}"),
}
}
#[test]
fn run_decode_loop_stops_on_max_seq_overrun() {
let outcome = run_decode_loop(
b'A' as u32,
8,
10,
10,
&[999],
canned_step_model(vec![b'B' as u32, b'C' as u32, b'D' as u32]),
ascii_decode,
|_event| {},
)
.expect("ok");
assert_eq!(outcome.stop_reason, DecodeStopReason::MaxSeq);
assert_eq!(outcome.generated, 3);
}
#[test]
fn run_decode_loop_emits_step_zero_seed_event() {
let mut events: Vec<(usize, u32, String, String)> = Vec::new();
let _ = run_decode_loop(
b'A' as u32,
10,
1024,
2,
&[999],
canned_step_model(vec![b'B' as u32]),
ascii_decode,
|event| {
events.push((
event.step,
event.token,
event.cumulative_text.to_string(),
event.delta.to_string(),
));
},
)
.expect("ok");
assert_eq!(events.len(), 2, "expected 2 events (step 0 + step 1)");
assert_eq!(events[0].0, 0);
assert_eq!(events[0].1, b'A' as u32);
assert_eq!(events[0].2, "A");
assert_eq!(events[0].3, "A", "step 0 delta MUST equal cumulative");
}
#[test]
fn run_decode_loop_delta_is_cumulative_minus_previous() {
let mut deltas: Vec<String> = Vec::new();
let _ = run_decode_loop(
b'X' as u32,
10,
1024,
4,
&[999],
canned_step_model(vec![b'Y' as u32, b'Z' as u32, b'!' as u32]),
ascii_decode,
|event| deltas.push(event.delta.to_string()),
)
.expect("ok");
assert_eq!(deltas, vec!["X", "Y", "Z", "!"]);
}
#[test]
fn run_decode_loop_propagates_step_model_error() {
let outcome = run_decode_loop(
b'A' as u32,
10,
1024,
5,
&[999],
|_prev, pos, _generated| -> super::Result<u32> {
Err(anyhow::anyhow!("synthetic forward failure at pos {pos}"))
},
ascii_decode,
|_event| {},
);
assert!(outcome.is_err());
let msg = format!("{:#}", outcome.unwrap_err());
assert!(
msg.contains("synthetic forward failure"),
"expected error to propagate, got: {msg}"
);
}
#[test]
fn run_decode_loop_zero_max_tokens_emits_seed_only() {
let mut events_count = 0usize;
let outcome = run_decode_loop(
b'A' as u32,
10,
1024,
1,
&[999],
canned_step_model(vec![]), ascii_decode,
|_event| events_count += 1,
)
.expect("ok");
assert_eq!(events_count, 1, "only the seed event must fire");
assert_eq!(outcome.generated, 1);
assert_eq!(outcome.stop_reason, DecodeStopReason::MaxTokensReached);
}
#[test]
fn run_decode_loop_initial_token_is_eos_stops_at_top_of_step_one() {
let mut events_count = 0usize;
let outcome = run_decode_loop(
42,
10,
1024,
5,
&[42],
canned_step_model(vec![]),
ascii_decode,
|_event| events_count += 1,
)
.expect("ok");
assert_eq!(events_count, 1);
assert_eq!(outcome.generated, 1);
assert_eq!(outcome.stop_reason, DecodeStopReason::EosTokenId(42));
}
#[test]
fn run_decode_loop_phi3_end_marker_caught_via_special_token_leak() {
let mut step = 0;
let outcome = run_decode_loop(
b'A' as u32,
10,
1024,
10,
&[999],
canned_step_model(vec![b'B' as u32, b'C' as u32, b'D' as u32]),
move |toks| {
step += 1;
let mut s: String = toks.iter().map(|&t| char::from(t as u8)).collect();
if step >= 4 {
s.push_str("<|end|>");
}
s
},
|_event| {},
)
.expect("ok");
assert_eq!(
outcome.stop_reason,
DecodeStopReason::SpecialTokenLeak("<|end|>")
);
}
#[test]
fn c4_scheduler_env_unset_defaults_to_fifo_serial() {
let mode = parse_scheduler_config(None, None, None, None)
.expect("env-absence must yield SerialFifo, not an error");
assert_eq!(
mode,
crate::serve::api::engine::EngineMode::SerialFifo,
"ADR-040 §3.6: env-absence MUST be byte-equivalent to \
pre-ADR-040 (= EngineMode::SerialFifo). Got {mode:?}."
);
}
#[test]
fn c4_scheduler_env_fifo_serial_lowercase_matches() {
let mode = parse_scheduler_config(None, Some("fifo_serial"), None, None)
.expect("HF2Q_SCHEDULER=fifo_serial must parse");
assert_eq!(mode, crate::serve::api::engine::EngineMode::SerialFifo);
}
#[test]
fn c4_scheduler_env_inflight_batched_matches() {
let mode = parse_scheduler_config(None, Some("inflight_batched"), None, None)
.expect("HF2Q_SCHEDULER=inflight_batched must parse");
assert_eq!(
mode,
crate::serve::api::engine::EngineMode::SlotAware {
max_slots: DEFAULT_MAX_SLOTS_UNDER_INFLIGHT,
}
);
}
#[test]
fn c4_scheduler_env_case_insensitive() {
let mode = parse_scheduler_config(None, Some("INFLIGHT_BATCHED"), None, None)
.expect("HF2Q_SCHEDULER=INFLIGHT_BATCHED must parse case-insensitively");
assert!(matches!(
mode,
crate::serve::api::engine::EngineMode::SlotAware { .. }
));
let mode = parse_scheduler_config(None, Some("FiFo_SeRiAl"), None, None)
.expect("HF2Q_SCHEDULER=FiFo_SeRiAl must parse case-insensitively");
assert_eq!(mode, crate::serve::api::engine::EngineMode::SerialFifo);
let mode = parse_scheduler_config(None, Some(" "), None, None)
.expect("HF2Q_SCHEDULER=` ` must be treated as unset");
assert_eq!(mode, crate::serve::api::engine::EngineMode::SerialFifo);
}
#[test]
fn c4_scheduler_env_unknown_value_errors() {
let err = parse_scheduler_config(None, Some("foo"), None, None)
.expect_err("unknown HF2Q_SCHEDULER value must error");
assert!(err.contains("HF2Q_SCHEDULER"), "msg: {err}");
assert!(err.contains("foo"), "msg must echo bad value: {err}");
assert!(
err.contains("fifo_serial") && err.contains("inflight_batched"),
"msg must name supported values: {err}"
);
}
#[test]
fn c4_max_slots_env_unset_defaults_to_4_under_inflight() {
let mode = parse_scheduler_config(None, Some("inflight_batched"), None, None)
.expect("inflight_batched without HF2Q_MAX_SLOTS must default");
assert_eq!(
mode,
crate::serve::api::engine::EngineMode::SlotAware { max_slots: 4 },
"ADR-040 §3.4: default max_slots under inflight_batched MUST be 4"
);
assert_eq!(
DEFAULT_MAX_SLOTS_UNDER_INFLIGHT, 4,
"ADR-040 §3.4: the named constant MUST hold the 4 default"
);
}
#[test]
fn c4_max_slots_env_unset_defaults_to_1_under_fifo_serial() {
let mode = parse_scheduler_config(Some(cli::SchedulerArg::FifoSerial), None, Some(8), None)
.expect("max_slots on fifo_serial must be ignored, not error");
assert_eq!(
mode,
crate::serve::api::engine::EngineMode::SerialFifo,
"ADR-040: under fifo_serial, max_slots is ignored — \
resolved mode MUST be SerialFifo (single-slot by definition). \
Got {mode:?}."
);
let mode2 = parse_scheduler_config(None, None, None, None).expect("default ok");
assert_eq!(mode2, crate::serve::api::engine::EngineMode::SerialFifo);
}
#[test]
fn c4_max_slots_env_zero_normalizes_or_errors() {
let err = parse_scheduler_config(
Some(cli::SchedulerArg::InflightBatched),
None,
Some(0),
None,
)
.expect_err("max_slots=0 must be rejected, not silently coerced");
assert!(err.contains("max-slots"), "msg: {err}");
assert!(err.contains('0'), "msg must echo the bad value: {err}");
assert!(
err.contains("F3a") || err.contains(".max(1)"),
"msg should cite the iter-2.5 F3a discipline: {err}"
);
let err = parse_scheduler_config(None, Some("inflight_batched"), None, Some("0"))
.expect_err("HF2Q_MAX_SLOTS=0 must be rejected, not silently coerced");
assert!(err.contains("max-slots"), "msg: {err}");
let err = parse_scheduler_config(None, None, None, Some("not-a-number"))
.expect_err("non-u32 HF2Q_MAX_SLOTS must error");
assert!(err.contains("HF2Q_MAX_SLOTS"), "msg: {err}");
assert!(err.contains("u32"), "msg must name u32: {err}");
}
#[test]
fn c4_scheduler_cli_wins_over_env() {
let mode = parse_scheduler_config(
Some(cli::SchedulerArg::FifoSerial),
Some("inflight_batched"),
None,
None,
)
.expect("cli wins, must resolve to SerialFifo");
assert_eq!(mode, crate::serve::api::engine::EngineMode::SerialFifo);
let mode = parse_scheduler_config(
Some(cli::SchedulerArg::InflightBatched),
Some("fifo_serial"),
Some(7),
None,
)
.expect("cli wins, must resolve to SlotAware{7}");
assert_eq!(
mode,
crate::serve::api::engine::EngineMode::SlotAware { max_slots: 7 }
);
}
#[test]
fn c4_max_slots_cli_wins_over_env() {
let mode = parse_scheduler_config(
Some(cli::SchedulerArg::InflightBatched),
None,
Some(2),
Some("16"),
)
.expect("cli max_slots wins over env");
assert_eq!(
mode,
crate::serve::api::engine::EngineMode::SlotAware { max_slots: 2 }
);
}
}