use anyhow::{anyhow, Context, Result};
use mlx_native::ops::elementwise::elementwise_add;
use mlx_native::ops::rms_norm::dispatch_rms_norm;
use mlx_native::{DType, GgmlType, KernelRegistry, MlxBuffer, MlxDevice};
use crate::inference::models::qwen35::gpu_ffn::{
build_dense_ffn_layer_gpu_q_into, DenseFfnWeightsGpuQ,
};
use crate::inference::models::qwen35::gpu_full_attn::{
apply_imrope, apply_linear_projection_f32, apply_q_or_k_per_head_rms_norm,
apply_sdpa_causal_from_seq_major, download_f32, upload_f32,
};
use crate::inference::models::qwen35::io_heads::embed_tokens;
use crate::inference::vision::image_token_residual_add::image_token_residual_add_gpu;
use crate::serve::forward_prefill::{DeepstackInjection, SoftTokenInjection};
use super::Qwen3VlTextModel;
pub const QWEN3VL_TEXT_FORWARD_PENDING_SENTINEL: &str = "qwen3vl_text_forward_pending";
pub const QWEN3VL_TEXT_FORWARD_PENDING_MESSAGE: &str =
"Qwen3-VL text-LM dense forward is implemented in `forward.rs::forward_text_prefill_logits_last` \
(iter-8a-2 LANDED) but the engine seam wire-up is iter-9b scope. The Generate / GenerateStream / \
GenerateWithSoftTokens dispatch arms in `serve/api/engine.rs::worker_run` still return this \
sentinel; iter-9b replaces them with calls into the prefill forward path. For text-only chat \
today, use a Qwen3.5/3.6 GGUF (full chat path) or a Gemma 4 GGUF (full chat + image path).";
pub fn qwen3vl_text_forward_pending_err<T>() -> Result<T> {
Err(anyhow::anyhow!(
"{}: {}",
QWEN3VL_TEXT_FORWARD_PENDING_SENTINEL,
QWEN3VL_TEXT_FORWARD_PENDING_MESSAGE
))
}
pub fn forward_text_prefill_logits_last(
model: &mut Qwen3VlTextModel,
tokens: &[u32],
positions_flat: &[i32],
deepstack: Option<&DeepstackInjection<'_>>,
soft_tokens: &[SoftTokenInjection<'_>],
) -> Result<Vec<f32>> {
if tokens.is_empty() {
return Err(anyhow!(
"forward_text_prefill_logits_last: empty token list (need at least 1 prompt token)"
));
}
let seq_len_usize = tokens.len();
if positions_flat.len() != 4 * seq_len_usize {
return Err(anyhow!(
"forward_text_prefill_logits_last: positions_flat length ({}) must equal \
4 * tokens.len() ({}) — IMROPE expects 4 axes (t/y/x/pad) per token",
positions_flat.len(),
4 * seq_len_usize,
));
}
let seq_len = seq_len_usize as u32;
if let Some(ds) = deepstack {
let n_image_tokens = ds.n_image_tokens();
let n_chunks = ds.n_deepstack();
let n_required = model.cfg.n_deepstack_layers;
if n_chunks < n_required {
return Err(anyhow!(
"forward_text_prefill_logits_last: deepstack.chunks.len() ({}) < \
cfg.n_deepstack_layers ({}). The peer dispatch \
(`qwen3vl.cpp:146-150`) reads chunk `il` for every LM layer \
`il < n_deepstack_layers`; missing chunks would silently skip \
their dispatch and break image conditioning.",
n_chunks,
n_required
));
}
if n_image_tokens == 0 {
return Err(anyhow!(
"forward_text_prefill_logits_last: deepstack supplied with \
image_token_positions.len() == 0; pass `deepstack = None` for \
text-only prompts (the per-image-token kernel rejects empty \
position arrays)."
));
}
for (i, &pos) in ds.image_token_positions.iter().enumerate() {
if (pos as usize) >= seq_len_usize {
return Err(anyhow!(
"forward_text_prefill_logits_last: deepstack.image_token_positions[{}] \
= {} >= seq_len ({}); position out of range",
i,
pos,
seq_len_usize
));
}
}
}
let cfg = model.cfg.clone();
let hidden = cfg.hidden_size;
let head_dim = cfg.head_dim;
let n_heads = cfg.num_attention_heads;
let n_kv_heads = cfg.num_key_value_heads;
let kv_dim = n_kv_heads * head_dim;
let intermediate = cfg.intermediate_size;
let n_layers = cfg.num_hidden_layers as usize;
let vocab = cfg.vocab_size;
let rms_eps = cfg.rms_norm_eps;
let rope_theta = cfg.rope_theta;
let mrope_section = cfg.mrope_section;
let tied = cfg.tied_word_embeddings;
let rotary_dim = head_dim;
let weights = &model.weights;
let token_embd_cpu =
download_f32(&weights.token_embd).context("download token_embd for embedding lookup")?;
let mut hidden_cpu = embed_tokens(tokens, &token_embd_cpu, vocab, hidden);
drop(token_embd_cpu);
if !soft_tokens.is_empty() {
let h = hidden as usize;
let mut occupied: Vec<bool> = vec![false; seq_len_usize];
for inj in soft_tokens {
if inj.range.start >= inj.range.end {
return Err(anyhow!(
"forward_text: SoftTokenInjection has empty/inverted range \
[{}, {}); each injection must cover ≥ 1 token",
inj.range.start,
inj.range.end
));
}
if inj.range.end > seq_len_usize {
return Err(anyhow!(
"forward_text: SoftTokenInjection range [{}, {}) extends past \
prompt_len ({})",
inj.range.start,
inj.range.end,
seq_len_usize
));
}
for p in inj.range.clone() {
if occupied[p] {
return Err(anyhow!(
"forward_text: SoftTokenInjection ranges overlap at position {p}"
));
}
occupied[p] = true;
}
let required_bytes = inj.range.len() * h * 4;
let span = inj
.embeddings
.byte_len()
.saturating_sub(inj.embeddings.byte_offset() as usize);
if span < required_bytes {
return Err(anyhow!(
"forward_text: SoftTokenInjection embeddings span {} < required {} \
(range.len()={} * hidden={} * 4)",
span,
required_bytes,
inj.range.len(),
hidden
));
}
if inj.embeddings.dtype() != DType::F32 {
return Err(anyhow!(
"forward_text: SoftTokenInjection embeddings dtype must be F32; got {:?}",
inj.embeddings.dtype()
));
}
}
for inj in soft_tokens {
let chunk = inj
.embeddings
.as_slice::<f32>()
.map_err(|e| anyhow!("soft-token embeddings as_slice: {e}"))?;
for (i, p) in inj.range.clone().enumerate() {
let src_off = i * h;
let dst_off = p * h;
hidden_cpu[dst_off..dst_off + h].copy_from_slice(&chunk[src_off..src_off + h]);
}
}
}
let (executor, registry) = model.ctx.split();
let device = executor.device();
crate::inference::vision::image_token_residual_add::register_image_token_residual_add_shader(
registry,
);
let mut hidden_gpu =
upload_f32(&hidden_cpu, device).context("upload initial residual stream")?;
drop(hidden_cpu);
let positions_gpu = {
let byte_len = positions_flat.len() * 4;
let mut buf = device
.alloc_buffer(byte_len, DType::I32, vec![positions_flat.len()])
.map_err(|e| anyhow!("alloc IMROPE positions buffer: {e}"))?;
buf.as_mut_slice::<i32>()
.map_err(|e| anyhow!("positions as_mut_slice: {e}"))?
.copy_from_slice(positions_flat);
buf
};
for il in 0..n_layers {
let lw = &weights.layers[il];
let (q_rope, k_rope, v_seq) = {
let mut session_a = executor
.begin()
.with_context(|| format!("layer {il}: begin Phase A session"))?;
let enc = session_a.encoder_mut();
let attn_normed = rms_norm_2d(
enc,
registry,
device,
&hidden_gpu,
&lw.attn_norm,
seq_len,
hidden,
rms_eps,
)
.with_context(|| format!("layer {il}: pre-attn rms_norm"))?;
enc.memory_barrier();
let q_seq = apply_linear_projection_f32(
enc,
registry,
device,
&attn_normed,
&lw.attn_q,
seq_len,
hidden,
hidden,
)
.with_context(|| format!("layer {il}: Q proj"))?;
let k_seq = apply_linear_projection_f32(
enc,
registry,
device,
&attn_normed,
&lw.attn_k,
seq_len,
hidden,
kv_dim,
)
.with_context(|| format!("layer {il}: K proj"))?;
let v_seq = apply_linear_projection_f32(
enc,
registry,
device,
&attn_normed,
&lw.attn_v,
seq_len,
hidden,
kv_dim,
)
.with_context(|| format!("layer {il}: V proj"))?;
enc.memory_barrier();
let q_normed = apply_q_or_k_per_head_rms_norm(
enc,
registry,
device,
&q_seq,
&lw.attn_q_norm,
seq_len,
n_heads,
head_dim,
rms_eps,
)
.with_context(|| format!("layer {il}: Q per-head rms_norm"))?;
let k_normed = apply_q_or_k_per_head_rms_norm(
enc,
registry,
device,
&k_seq,
&lw.attn_k_norm,
seq_len,
n_kv_heads,
head_dim,
rms_eps,
)
.with_context(|| format!("layer {il}: K per-head rms_norm"))?;
enc.memory_barrier();
let q_rope = apply_imrope(
enc,
registry,
device,
&q_normed,
&positions_gpu,
seq_len,
n_heads,
head_dim,
rotary_dim,
rope_theta,
mrope_section,
)
.with_context(|| format!("layer {il}: Q IMROPE"))?;
let k_rope = apply_imrope(
enc,
registry,
device,
&k_normed,
&positions_gpu,
seq_len,
n_kv_heads,
head_dim,
rotary_dim,
rope_theta,
mrope_section,
)
.with_context(|| format!("layer {il}: K IMROPE"))?;
session_a
.finish()
.with_context(|| format!("layer {il}: finish Phase A session"))?;
(q_rope, k_rope, v_seq)
};
let attn_out = {
let mut session_b = executor
.begin()
.with_context(|| format!("layer {il}: begin Phase B session"))?;
let enc = session_b.encoder_mut();
let attn_out = apply_sdpa_causal_from_seq_major(
enc, registry, device, &q_rope, &k_rope, &v_seq, seq_len, n_heads, n_kv_heads,
head_dim,
)
.with_context(|| format!("layer {il}: SDPA"))?;
attn_out
};
let l_out = {
let mut session_c = executor
.begin()
.with_context(|| format!("layer {il}: begin Phase C session"))?;
let enc = session_c.encoder_mut();
let attn_proj = apply_linear_projection_f32(
enc,
registry,
device,
&attn_out,
&lw.attn_output,
seq_len,
hidden,
hidden,
)
.with_context(|| format!("layer {il}: output proj"))?;
enc.memory_barrier();
let ffn_inp = elementwise_add_f32_2d(
enc,
registry,
device,
&hidden_gpu,
&attn_proj,
seq_len,
hidden,
)
.with_context(|| format!("layer {il}: post-attn residual add"))?;
enc.memory_barrier();
let ffn_normed = rms_norm_2d(
enc,
registry,
device,
&ffn_inp,
&lw.ffn_norm,
seq_len,
hidden,
rms_eps,
)
.with_context(|| format!("layer {il}: ffn rms_norm"))?;
enc.memory_barrier();
let ffn_weights = DenseFfnWeightsGpuQ {
gate_q: lw.ffn_gate.clone(),
up_q: lw.ffn_up.clone(),
down_q: lw.ffn_down.clone(),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: intermediate,
hidden_size: hidden,
};
let l_out = build_dense_ffn_layer_gpu_q_into(
enc,
device,
registry,
&ffn_normed,
&ffn_weights,
Some(&ffn_inp),
)
.with_context(|| format!("layer {il}: dense SwiGLU FFN"))?;
session_c
.finish()
.with_context(|| format!("layer {il}: finish Phase C session"))?;
l_out
};
if let Some(ds) = deepstack {
if il < cfg.n_deepstack_layers {
let chunk = ds.chunks[il];
let n_image_tokens = ds.n_image_tokens() as u32;
{
let mut session_d = executor.begin().with_context(|| {
format!("layer {il}: begin Phase D (deepstack) session")
})?;
let enc = session_d.encoder_mut();
image_token_residual_add_gpu(
enc,
registry,
device,
&l_out,
chunk,
&ds.image_token_positions,
seq_len,
n_image_tokens,
hidden,
)
.with_context(|| format!("layer {il}: deepstack residual add (slab {il})"))?;
session_d.finish().with_context(|| {
format!("layer {il}: finish Phase D (deepstack) session")
})?;
}
}
}
hidden_gpu = l_out;
}
let final_residual_cpu = hidden_gpu
.as_slice::<f32>()
.map_err(|e| anyhow!("final residual as_slice: {e}"))?;
let last_row_start = (seq_len_usize - 1) * (hidden as usize);
let last_row_end = last_row_start + (hidden as usize);
let last_row: Vec<f32> = final_residual_cpu[last_row_start..last_row_end].to_vec();
let last_row_gpu =
upload_f32(&last_row, device).context("upload last-row residual for output head")?;
drop(last_row);
let lm_head_weight: &MlxBuffer = if tied {
&weights.token_embd
} else {
weights.output.as_ref().ok_or_else(|| {
anyhow!(
"Qwen3-VL text-LM weights inconsistent: tied_word_embeddings=false but \
output is None — config and weights disagree"
)
})?
};
let logits_buf = {
let mut session_head = executor.begin().context("begin output-head session")?;
let enc = session_head.encoder_mut();
let final_normed = rms_norm_2d(
enc,
registry,
device,
&last_row_gpu,
&weights.output_norm,
1,
hidden,
rms_eps,
)
.context("final output rms_norm")?;
enc.memory_barrier();
let logits_buf = apply_linear_projection_f32(
enc,
registry,
device,
&final_normed,
lm_head_weight,
1,
hidden,
vocab,
)
.context("LM head matmul")?;
session_head
.finish()
.context("finish output-head session")?;
logits_buf
};
let logits_full = logits_buf
.as_slice::<f32>()
.map_err(|e| anyhow!("logits as_slice: {e}"))?;
let v = vocab as usize;
if logits_full.len() < v {
return Err(anyhow!(
"logits buffer too small: got {} elements, expected {}",
logits_full.len(),
v
));
}
Ok(logits_full[..v].to_vec())
}
fn rms_norm_2d(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
rows: u32,
dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
let out = device
.alloc_buffer(
(rows * dim) as usize * 4,
DType::F32,
vec![rows as usize, dim as usize],
)
.map_err(|e| anyhow!("alloc rms_norm output: {e}"))?;
let mut params = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc rms_norm params: {e}"))?;
{
let s = params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("rms_norm params as_mut_slice: {e}"))?;
s[0] = eps;
s[1] = dim as f32;
}
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
weight,
&out,
¶ms,
rows,
dim,
)
.context("dispatch_rms_norm")?;
Ok(out)
}
fn elementwise_add_f32_2d(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
a: &MlxBuffer,
b: &MlxBuffer,
rows: u32,
cols: u32,
) -> Result<MlxBuffer> {
let n_elements = (rows as usize) * (cols as usize);
let out = device
.alloc_buffer(
n_elements * 4,
DType::F32,
vec![rows as usize, cols as usize],
)
.map_err(|e| anyhow!("alloc elementwise_add output: {e}"))?;
elementwise_add(
encoder,
registry,
device.metal_device(),
a,
b,
&out,
n_elements,
DType::F32,
)
.map_err(|e| anyhow!("elementwise_add: {e}"))?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sentinel_is_stable_across_iters() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_eq!(
QWEN3VL_TEXT_FORWARD_PENDING_SENTINEL,
"qwen3vl_text_forward_pending"
);
}
#[test]
fn pending_err_carries_sentinel_substring() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err: Result<()> = qwen3vl_text_forward_pending_err();
let msg = format!("{:#}", err.unwrap_err());
assert!(
msg.contains(QWEN3VL_TEXT_FORWARD_PENDING_SENTINEL),
"error message must carry the sentinel substring; got: {msg}"
);
assert!(
msg.contains("iter-9b"),
"iter-8a-2 pending message must point at iter-9b for engine seam wire-up; got: {msg}"
);
}
#[test]
fn pending_err_message_is_operator_actionable() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err: Result<()> = qwen3vl_text_forward_pending_err();
let msg = format!("{:#}", err.unwrap_err());
assert!(
msg.contains("Qwen3.5") || msg.contains("Gemma"),
"error message must name a working alternative; got: {msg}"
);
}
#[test]
fn forward_text_prefill_shape_finite_when_operator_gated() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if std::env::var("HF2Q_QWEN3VL_LM_LOAD").ok().as_deref() != Some("1") {
eprintln!("skip: HF2Q_QWEN3VL_LM_LOAD!=1");
return;
}
let p =
std::path::PathBuf::from("/opt/hf2q/.cfa-archive/wedge4f-out/qwen3-vl-2b-q4_0.gguf");
if !p.exists() {
eprintln!("skip: real GGUF fixture not present at {}", p.display());
return;
}
let gguf = mlx_native::gguf::GgufFile::open(&p).expect("open real Qwen3-VL-2B GGUF");
let mut progress = crate::serve::header::LoadProgress::new(false, 0, 28);
let mut model = Qwen3VlTextModel::load_from_gguf(&gguf, &mut progress)
.expect("load real Qwen3-VL-2B text-LM model");
let tokens: Vec<u32> = vec![100, 200, 300, 400];
let seq = tokens.len();
let mut positions = vec![0i32; 4 * seq];
for axis in 0..4 {
for t in 0..seq {
positions[axis * seq + t] = t as i32;
}
}
let logits = forward_text_prefill_logits_last(&mut model, &tokens, &positions, None, &[])
.expect("forward must succeed on the canonical GGUF");
assert_eq!(
logits.len(),
model.cfg.vocab_size as usize,
"logits length must equal vocab_size"
);
let n_finite = logits.iter().filter(|x| x.is_finite()).count();
let n_nan = logits.iter().filter(|x| x.is_nan()).count();
let n_inf = logits.iter().filter(|x| x.is_infinite()).count();
assert_eq!(
n_nan, 0,
"logits must contain no NaN; got {n_nan} NaN entries"
);
assert_eq!(
n_inf, 0,
"logits must contain no Inf; got {n_inf} Inf entries"
);
assert_eq!(
n_finite,
logits.len(),
"all logits must be finite (no NaN/Inf); got {n_finite}/{}",
logits.len()
);
let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let min = logits.iter().cloned().fold(f32::INFINITY, f32::min);
assert!(
(max - min) > 0.1,
"logits should have nonzero spread; got max={max}, min={min}"
);
eprintln!(
"forward_text_prefill_shape_finite: vocab={} max_logit={:.4} min_logit={:.4}",
logits.len(),
max,
min
);
}
}