use std::path::PathBuf;
use std::time::{Duration, Instant};
use anyhow::{anyhow, Context, Result};
use tokenizers::Tokenizer;
use crate::core::provenance::{self, Provenance};
use crate::inference::models::qwen3vl_text::forward::forward_text_prefill_logits_last;
use crate::inference::models::qwen3vl_text::Qwen3VlTextModel;
use crate::serve::forward_prefill::{DeepstackInjection, SoftTokenInjection};
use crate::serve::load_info::{
self, ArchFamily, ChatTemplateSource, LoadInfo, LoadInfoBuilder, TokenizerSource,
};
use crate::serve::sampler_pure::{self, SamplingParams as SamplerPureParams};
use super::engine::{effective_repetition_penalty, GenerationResult, LoadOptions, SamplingParams};
use super::registry::ModelRegistration;
pub struct Qwen3VlTextLoadedModel {
pub model: Qwen3VlTextModel,
pub tokenizer: Tokenizer,
pub chat_template: String,
pub model_id: String,
pub model_path: PathBuf,
pub eos_token_ids: Vec<u32>,
pub hidden_size: usize,
pub vocab_size: usize,
pub context_length: Option<usize>,
pub quant_type: Option<String>,
pub load_duration: Duration,
pub provenance: Provenance,
pub slot_aware_max_slots: Option<u32>,
}
impl Qwen3VlTextLoadedModel {
pub fn provision_multi_seq_kv_for_slot_aware(&mut self, max_slots: u32) -> Result<()> {
if max_slots == 0 {
anyhow::bail!(
"ADR-040 C2e: provision_multi_seq_kv_for_slot_aware called with \
max_slots == 0; spawn_with_mode invariant is max_slots >= 1 \
(EngineMode::SlotAware variant enforces this at the API \
boundary — caller violated)"
);
}
self.slot_aware_max_slots = Some(max_slots);
Ok(())
}
pub fn handle_qwen3vl_slot_aware_n_gt_0_sentinel<T>(
&mut self,
slot_id: crate::serve::multi_seq_kv::SlotId,
arm_sublabel: &str,
) -> Result<T> {
let witness = self.slot_aware_max_slots.take();
let sentinel_result: Result<T> =
crate::inference::models::qwen3vl_text::forward::qwen3vl_text_forward_pending_err();
self.slot_aware_max_slots = witness;
sentinel_result.map_err(|sentinel_err| {
anyhow!(
"capability_unsupported: ADR-040 \
{arm_sublabel} (iter-C2e-cont per ADR-040 §6.1.55 — \
structural worker hot path lift gated on iter-228a \
Qwen3-VL forward path landing past the 501 sentinel; \
worker hot path lift onto the persistent multi-seq \
cache cannot land until the persistent cache itself \
exists). SlotId({}) — sentinel propagated verbatim: \
{sentinel_err}",
slot_id.0,
)
})
}
}
impl Qwen3VlTextLoadedModel {
pub fn load(opts: &LoadOptions) -> Result<Self> {
let load_start = Instant::now();
let model_path = &opts.model_path;
anyhow::ensure!(
model_path.exists(),
"Model not found: {}",
model_path.display()
);
let gguf = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF open: {e}"))?;
let provenance = provenance::detect(&gguf);
let cfg_preview = Qwen3VlTextModel::load_config_only(&gguf).context("config preview")?;
let tokenizer_path =
crate::serve::find_tokenizer(model_path, opts.tokenizer_path.as_deref())?;
let stderr_is_tty = std::io::IsTerminal::is_terminal(&std::io::stderr());
let verbosity = if tracing::enabled!(tracing::Level::INFO) {
1
} else {
0
};
let mut progress = crate::serve::header::LoadProgress::new(
stderr_is_tty,
verbosity,
cfg_preview.num_hidden_layers as usize,
);
let model = Qwen3VlTextModel::load_from_gguf(&gguf, &mut progress)
.context("Qwen3VlTextModel::load_from_gguf")?;
let eos_token: u32 = gguf
.metadata_u32("tokenizer.ggml.eos_token_id")
.unwrap_or(151645);
let eos_token_ids: Vec<u32> = vec![eos_token];
let mut tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path).map_err(|e| {
anyhow::anyhow!(
"Failed to load tokenizer.json from {}: {e}",
tokenizer_path.display()
)
})?;
tokenizer
.with_truncation(None)
.map_err(|e| anyhow::anyhow!("Failed to disable tokenizer truncation: {e}"))?;
let chat_template = gguf
.metadata_string("tokenizer.chat_template")
.map(|s| s.to_string())
.unwrap_or_default();
let model_id = gguf
.metadata_string("general.name")
.map(|s| s.to_string())
.unwrap_or_else(|| {
model_path
.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| "qwen3vl-text-model".to_string())
});
let hidden_size = model.cfg.hidden_size as usize;
let vocab_size = model.cfg.vocab_size as usize;
let context_length = if model.cfg.max_position_embeddings > 0 {
Some(model.cfg.max_position_embeddings as usize)
} else {
None
};
let quant_type = crate::serve::load_info::infer_quant_label(&gguf);
let load_duration = load_start.elapsed();
Ok(Self {
model,
tokenizer,
chat_template,
model_id,
model_path: model_path.clone(),
eos_token_ids,
hidden_size,
vocab_size,
context_length,
quant_type,
load_duration,
provenance,
slot_aware_max_slots: None,
})
}
}
impl LoadInfoBuilder for Qwen3VlTextLoadedModel {
fn build_load_info(
&self,
gguf: &mlx_native::gguf::GgufFile,
load_wall_clock: Duration,
kv_cache_budget_bytes: Option<u64>,
kv_spill_active: bool,
) -> LoadInfo {
let cfg = &self.model.cfg;
LoadInfo {
model_id: self.model_id.clone(),
arch_str: load_info::arch_str_from_gguf(gguf),
arch_family: ArchFamily::Qwen3VlText,
model_path: self.model_path.clone(),
on_disk_bytes: load_info::on_disk_bytes(&self.model_path),
backend_chip: mlx_native::MlxDevice::new()
.map(|d| d.name())
.unwrap_or_else(|_| "Apple GPU".to_string()),
backend: "mlx-native",
n_layers: cfg.num_hidden_layers,
hidden_size: self.hidden_size as u32,
vocab_size: self.vocab_size as u32,
n_attention_heads: cfg.num_attention_heads,
n_key_value_heads: cfg.num_key_value_heads,
head_dim: cfg.head_dim,
sliding_window: None,
full_attention_interval: None,
max_context_length: self.context_length.map(|v| v as u32),
moe: None,
quant_label: self.quant_type.clone(),
quant_bpw: load_info::compute_bpw(gguf),
tokenizer_source: TokenizerSource::GgufEmbedded,
eos_token_ids: self.eos_token_ids.clone(),
bos_token_id: gguf.metadata_u32("tokenizer.ggml.bos_token_id"),
chat_template_source: if gguf.metadata_string("tokenizer.chat_template").is_some() {
ChatTemplateSource::GgufEmbedded
} else {
ChatTemplateSource::None
},
provenance: self.provenance.clone(),
vision_projector: None,
load_wall_clock,
resident_weight_bytes: None,
kv_cache_budget_bytes,
kv_spill_active,
tq_kv_active: false,
kv_bytes_per_token_override: None,
}
}
}
fn text_only_positions_for(prompt_len: usize) -> 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
}
fn is_greedy_eligible_qwen3vl(params: &SamplingParams) -> bool {
!(params.temperature > 0.0
|| params.top_k > 0
|| params.top_p < 1.0
|| params.repetition_penalty != 1.0
|| params.seed.is_some())
}
fn argmax_u32(logits: &[f32]) -> u32 {
let mut best_idx: usize = 0;
let mut best_val: f32 = f32::NEG_INFINITY;
for (i, &v) in logits.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i;
}
}
best_idx as u32
}
fn sample_logits_qwen3vl(logits: &mut [f32], params: &SamplingParams, generated: &[u32]) -> u32 {
let sp = SamplerPureParams {
temperature: params.temperature as f64,
top_p: params.top_p as f64,
top_k: params.top_k,
min_p: params.min_p as f64,
repetition_penalty: effective_repetition_penalty(params),
max_tokens: params.max_tokens,
};
sampler_pure::sample_token(logits, &sp, generated)
}
fn decode_to_text(tokenizer: &Tokenizer, decoded_tokens: &[u32]) -> Result<String> {
tokenizer
.decode(decoded_tokens, false)
.map_err(|e| anyhow!("Qwen3-VL tokenizer decode: {e}"))
}
pub fn generate_qwen3vl_text_once(
qwen: &mut Qwen3VlTextLoadedModel,
prompt_tokens: &[u32],
params: &SamplingParams,
_registration: Option<&ModelRegistration>,
) -> Result<GenerationResult> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen3vl_text_once: empty prompt_tokens"
);
let prompt_len = prompt_tokens.len();
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible_qwen3vl(params);
let mut tokens_so_far: Vec<u32> = prompt_tokens.to_vec();
let mut decoded_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
let mut finish_reason: &'static str = "length";
let prefill_start = Instant::now();
let mut prefill_duration: Duration = Duration::ZERO;
let decode_start_outer = Instant::now();
for step in 0..max_tokens {
let positions = text_only_positions_for(tokens_so_far.len());
let mut logits = forward_text_prefill_logits_last(
&mut qwen.model,
&tokens_so_far,
&positions,
None, &[], )
.with_context(|| format!("forward step {step}"))?;
if step == 0 {
prefill_duration = prefill_start.elapsed();
}
let next_token: u32 = if is_greedy {
argmax_u32(&logits)
} else {
sample_logits_qwen3vl(&mut logits, params, &decoded_tokens)
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
decoded_tokens.push(next_token);
tokens_so_far.push(next_token);
}
let decode_duration = decode_start_outer
.elapsed()
.saturating_sub(prefill_duration);
let text = decode_to_text(&qwen.tokenizer, &decoded_tokens)?;
Ok(GenerationResult {
text,
reasoning_text: None,
prompt_tokens: prompt_len,
completion_tokens: decoded_tokens.len(),
reasoning_tokens: None,
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: 0,
logprobs: None,
})
}
pub fn generate_qwen3vl_text_with_soft_tokens_once(
qwen: &mut Qwen3VlTextLoadedModel,
prompt_tokens: &[u32],
soft_tokens: &[SoftTokenInjection<'_>],
deepstack: Option<&DeepstackInjection<'_>>,
positions_flat: &[i32],
params: &SamplingParams,
_registration: Option<&ModelRegistration>,
) -> Result<GenerationResult> {
anyhow::ensure!(
!prompt_tokens.is_empty(),
"generate_qwen3vl_text_with_soft_tokens_once: empty prompt_tokens"
);
let prompt_len = prompt_tokens.len();
if positions_flat.len() != 4 * prompt_len {
return Err(anyhow!(
"generate_qwen3vl_text_with_soft_tokens_once: positions_flat.len()={} != \
4 * prompt_len ({})",
positions_flat.len(),
4 * prompt_len
));
}
let max_tokens = params.max_tokens.max(1);
let is_greedy = is_greedy_eligible_qwen3vl(params);
let mut tokens_so_far: Vec<u32> = prompt_tokens.to_vec();
let mut decoded_tokens: Vec<u32> = Vec::with_capacity(max_tokens);
let mut finish_reason: &'static str = "length";
let mut current_positions: Vec<i32> = Vec::with_capacity(4 * (prompt_len + max_tokens));
let prefill_start = Instant::now();
let mut prefill_duration: Duration = Duration::ZERO;
let decode_start_outer = Instant::now();
for step in 0..max_tokens {
let total_len = prompt_len + step;
current_positions.resize(4 * total_len, 0i32);
for axis in 0..4 {
for t in 0..prompt_len {
current_positions[axis * total_len + t] = positions_flat[axis * prompt_len + t];
}
}
if step > 0 {
let last_prompt_t = positions_flat[prompt_len - 1];
for s in 0..step {
let p = prompt_len + s;
let t_val = last_prompt_t + 1 + s as i32;
for axis in 0..4 {
current_positions[axis * total_len + p] = t_val;
}
}
}
let mut logits = forward_text_prefill_logits_last(
&mut qwen.model,
&tokens_so_far,
¤t_positions,
deepstack,
soft_tokens,
)
.with_context(|| format!("forward step {step} (multimodal)"))?;
if step == 0 {
prefill_duration = prefill_start.elapsed();
}
let next_token: u32 = if is_greedy {
argmax_u32(&logits)
} else {
sample_logits_qwen3vl(&mut logits, params, &decoded_tokens)
};
if qwen.eos_token_ids.contains(&next_token) {
finish_reason = "stop";
break;
}
decoded_tokens.push(next_token);
tokens_so_far.push(next_token);
}
let decode_duration = decode_start_outer
.elapsed()
.saturating_sub(prefill_duration);
let text = decode_to_text(&qwen.tokenizer, &decoded_tokens)?;
Ok(GenerationResult {
text,
reasoning_text: None,
prompt_tokens: prompt_len,
completion_tokens: decoded_tokens.len(),
reasoning_tokens: None,
finish_reason,
prefill_duration,
decode_duration,
cached_tokens: 0,
logprobs: None,
})
}