use anyhow::Result;
use mlx_native::graph::GraphSession;
use mlx_native::{DType, KernelRegistry, MlxBuffer, MlxDevice};
use crate::debug::INVESTIGATION_ENV;
use crate::quantize::imatrix::ImatrixHint;
use crate::serve::forward_mlx_shared::dispatch_qmatmul;
use crate::serve::gpu::GpuContext;
use super::model::MlxModelWeights;
static FIRST_HEAD_TRACE: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
static RERANK_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static RERANK_CAND: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static RERANK_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static RERANK_PROFILE_ON: std::sync::LazyLock<bool> =
std::sync::LazyLock::new(|| std::env::var("HF2Q_RERANK_PROFILE").as_deref() == Ok("1"));
pub fn rerank_profile() -> (u64, u64, u64) {
use std::sync::atomic::Ordering::Relaxed;
(
RERANK_NS.load(Relaxed),
RERANK_CAND.load(Relaxed),
RERANK_CALLS.load(Relaxed),
)
}
pub fn rerank_profile_reset() {
use std::sync::atomic::Ordering::Relaxed;
RERANK_NS.store(0, Relaxed);
RERANK_CAND.store(0, Relaxed);
RERANK_CALLS.store(0, Relaxed);
}
pub struct BatchedHeadOut {
pub logits: Vec<f32>,
pub normed: Vec<f32>,
pub gpu_sample: Option<GpuSampleOut>,
}
pub struct GpuSampleOut {
pub top1_idx: Vec<u32>, pub top1_val: Vec<f32>, pub cand_count: Vec<u32>, pub overflow: Vec<u32>, pub cand_ids: Vec<u32>, pub cap: usize,
}
pub struct GpuSampleBuffers {
pub top1_idx: mlx_native::MlxBuffer,
pub top1_val: mlx_native::MlxBuffer,
pub cand_count: mlx_native::MlxBuffer,
pub overflow: mlx_native::MlxBuffer,
pub cand_ids: mlx_native::MlxBuffer,
pub params: mlx_native::MlxBuffer,
pub cap: usize,
}
impl MlxModelWeights {
pub(crate) fn finalize_token_from_logits(
&self,
logits_row: &[f32],
normed_row: &[f32],
gpu_top1: u32,
top1_val: f32,
) -> Result<u32> {
let vocab_size = self.vocab_size;
if !self.rerank_active() {
return Ok(gpu_top1);
}
let delta: f32 = 0.5;
let threshold = top1_val - delta;
let mut candidates: Vec<u32> = Vec::with_capacity(64);
for (i, &v) in logits_row[..vocab_size].iter().enumerate() {
if v >= threshold {
candidates.push(i as u32);
}
}
self.rerank_candidates(candidates, normed_row, gpu_top1)
}
pub(crate) fn finalize_token_from_gpu_candidates(
&self,
gpu_cand_ids: &[u32],
normed_row: &[f32],
gpu_top1: u32,
) -> Result<u32> {
if !self.rerank_active() {
return Ok(gpu_top1);
}
self.rerank_candidates(gpu_cand_ids.to_vec(), normed_row, gpu_top1)
}
#[inline]
fn rerank_active(&self) -> bool {
(self.lm_head_q8.is_some() || self.lm_head_q6k.is_some())
&& !INVESTIGATION_ENV.lmhead_rerank_disabled
}
fn rerank_candidates(
&self,
mut candidates: Vec<u32>,
normed_row: &[f32],
gpu_top1: u32,
) -> Result<u32> {
let vocab_size = self.vocab_size;
let hs = self.hidden_size;
let embed_f32: &[f32] = self
.embed_weight
.as_slice()
.map_err(|e| anyhow::anyhow!("finalize rerank embed read: {e}"))?;
for sp in [0u32, 1, 2, 105, 106] {
if (sp as usize) < vocab_size {
candidates.push(sp);
}
}
candidates.sort_unstable();
candidates.dedup();
let _rr_prof = *RERANK_PROFILE_ON;
let _rr_t = if _rr_prof {
Some(std::time::Instant::now())
} else {
None
};
let mut best_tok: u32 = gpu_top1;
let mut best_logit: f32 = f32::NEG_INFINITY;
for &tok in &candidates {
let row_off = (tok as usize) * hs;
if row_off + hs > embed_f32.len() {
continue;
}
let row = &embed_f32[row_off..row_off + hs];
let mut acc: f64 = 0.0;
for i in 0..hs {
acc += (normed_row[i] as f64) * (row[i] as f64);
}
let l = acc as f32;
if l > best_logit {
best_logit = l;
best_tok = tok;
}
}
if let Some(t) = _rr_t {
RERANK_NS.fetch_add(
t.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
RERANK_CAND.fetch_add(
candidates.len() as u64,
std::sync::atomic::Ordering::Relaxed,
);
RERANK_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
Ok(best_tok)
}
pub fn lm_head_batched(
&self,
hidden_rows: &[f32],
n: usize,
gpu: &mut GpuContext,
) -> Result<BatchedHeadOut> {
let hs = self.hidden_size;
let vocab = self.vocab_size;
if n == 0 {
return Ok(BatchedHeadOut {
logits: Vec::new(),
normed: Vec::new(),
gpu_sample: None,
});
}
if hidden_rows.len() != n * hs {
anyhow::bail!(
"lm_head_batched: hidden_rows len {} != n*hidden_size {}*{}",
hidden_rows.len(),
n,
hs
);
}
let q6k = self.lm_head_q6k.as_ref().ok_or_else(|| {
anyhow::anyhow!("lm_head_batched requires a Q6_K lm_head (production decode path)")
})?;
let (exec, reg) = gpu.split();
let dev = exec.device();
let metal_dev = dev.metal_device();
let mut hidden_b = dev
.alloc_buffer(n * hs * 4, DType::F32, vec![n * hs])
.map_err(|e| anyhow::anyhow!("lm_head_batched alloc hidden_b: {e}"))?;
hidden_b
.as_mut_slice::<f32>()
.map_err(|e| anyhow::anyhow!("lm_head_batched write hidden_b: {e}"))?
.copy_from_slice(hidden_rows);
let normed_b = dev
.alloc_buffer(n * hs * 4, DType::F32, vec![n * hs])
.map_err(|e| anyhow::anyhow!("lm_head_batched alloc normed_b: {e}"))?;
let logits_b = dev
.alloc_buffer(n * vocab * 4, DType::F32, vec![n * vocab])
.map_err(|e| anyhow::anyhow!("lm_head_batched alloc logits_b: {e}"))?;
let mut softcap_params_b = dev
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow::anyhow!("lm_head_batched alloc softcap params: {e}"))?;
if let Some(cap) = self.final_logit_softcapping {
let p: &mut [f32] = softcap_params_b
.as_mut_slice()
.map_err(|e| anyhow::anyhow!("lm_head_batched softcap params slice: {e}"))?;
let total = n
.checked_mul(vocab)
.expect("lm_head softcap: n*vocab overflow");
p[0] = cap;
p[1] = f32::from_bits(total as u32);
}
let mut s = exec
.begin()
.map_err(|e| anyhow::anyhow!("lm_head_batched session begin: {e}"))?;
let _bc_dbg = std::env::var("HF2Q_MVN_BARRIER_TRACE").as_deref() == Ok("1");
let _bc_before = if _bc_dbg {
mlx_native::barrier_count()
} else {
0
};
let _enc_trace = std::env::var("HF2Q_MVN_ENCODE_TRACE").as_deref() == Ok("1")
&& FIRST_HEAD_TRACE.swap(false, std::sync::atomic::Ordering::Relaxed);
if _enc_trace {
eprintln!("[ENCODE-TRACE] === lm_head_batched BEGIN n={} ===", n);
mlx_native::set_encode_trace(true);
}
s.barrier_between(&[&hidden_b, &self.final_norm], &[&normed_b]);
s.rms_norm(
reg,
metal_dev,
&hidden_b,
&self.final_norm,
&normed_b,
&self.activations.norm_params,
n as u32,
hs as u32,
)
.map_err(|e| anyhow::anyhow!("lm_head_batched final norm: {e}"))?;
s.barrier_between(&[&normed_b, &q6k.buffer], &[&logits_b]);
dispatch_qmatmul(
&mut s,
reg,
dev,
&normed_b,
q6k,
&logits_b,
n as u32,
ImatrixHint::Global("output.weight"),
)?;
if let Some(cap) = self.final_logit_softcapping {
s.barrier_between(&[&logits_b], &[&logits_b]);
mlx_native::ops::softcap::dispatch_softcap(
s.encoder_mut(),
reg,
metal_dev,
&logits_b,
&logits_b,
&softcap_params_b,
cap,
)
.map_err(|e| anyhow::anyhow!("lm_head_batched softcap: {e}"))?;
}
if _enc_trace {
mlx_native::set_encode_trace(false);
eprintln!("[ENCODE-TRACE] === lm_head_batched END n={} ===", n);
}
if _bc_dbg {
let after = mlx_native::barrier_count();
eprintln!(
"[BARRIER-TRACE] lm_head_batched n={} barriers_emitted={}",
n,
after - _bc_before
);
}
use crate::inference::models::gemma4::batched_body::host_phases;
let _hp = std::time::Instant::now();
s.finish()
.map_err(|e| anyhow::anyhow!("lm_head_batched session finish: {e}"))?;
host_phases::add(
host_phases::Phase::LmheadWait,
_hp.elapsed().as_nanos() as u64,
);
let _hp = std::time::Instant::now();
let logits: Vec<f32> = logits_b
.as_slice::<f32>()
.map_err(|e| anyhow::anyhow!("lm_head_batched read logits: {e}"))?
.to_vec();
let normed: Vec<f32> = normed_b
.as_slice::<f32>()
.map_err(|e| anyhow::anyhow!("lm_head_batched read normed_b: {e}"))?
.to_vec();
host_phases::add(
host_phases::Phase::LmheadReadback,
_hp.elapsed().as_nanos() as u64,
);
Ok(BatchedHeadOut {
logits,
normed,
gpu_sample: None,
})
}
pub(crate) fn encode_lm_head_into(
&self,
s: &mut GraphSession,
hidden_buf: &MlxBuffer,
n: usize,
dev: &MlxDevice,
reg: &mut KernelRegistry,
) -> Result<(MlxBuffer, MlxBuffer, MlxBuffer, Option<GpuSampleBuffers>)> {
let hs = self.hidden_size;
let vocab = self.vocab_size;
let metal_dev = dev.metal_device();
let q6k = self.lm_head_q6k.as_ref().ok_or_else(|| {
anyhow::anyhow!("encode_lm_head_into requires a Q6_K lm_head (production decode path)")
})?;
let normed_b = dev
.alloc_buffer(n * hs * 4, DType::F32, vec![n * hs])
.map_err(|e| anyhow::anyhow!("encode_lm_head_into alloc normed_b: {e}"))?;
let logits_b = dev
.alloc_buffer(n * vocab * 4, DType::F32, vec![n * vocab])
.map_err(|e| anyhow::anyhow!("encode_lm_head_into alloc logits_b: {e}"))?;
let mut softcap_params_b = dev
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow::anyhow!("encode_lm_head_into alloc softcap params: {e}"))?;
if let Some(cap) = self.final_logit_softcapping {
let p: &mut [f32] = softcap_params_b
.as_mut_slice()
.map_err(|e| anyhow::anyhow!("encode_lm_head_into softcap params slice: {e}"))?;
let total = n
.checked_mul(vocab)
.expect("encode_lm_head_into softcap: n*vocab overflow");
p[0] = cap;
p[1] = f32::from_bits(total as u32);
}
s.barrier_between(&[hidden_buf, &self.final_norm], &[&normed_b]);
s.rms_norm(
reg,
metal_dev,
hidden_buf,
&self.final_norm,
&normed_b,
&self.activations.norm_params,
n as u32,
hs as u32,
)
.map_err(|e| anyhow::anyhow!("encode_lm_head_into final norm: {e}"))?;
s.barrier_between(&[&normed_b, &q6k.buffer], &[&logits_b]);
dispatch_qmatmul(
s,
reg,
dev,
&normed_b,
q6k,
&logits_b,
n as u32,
ImatrixHint::Global("output.weight"),
)?;
if let Some(cap) = self.final_logit_softcapping {
s.barrier_between(&[&logits_b], &[&logits_b]);
mlx_native::ops::softcap::dispatch_softcap(
s.encoder_mut(),
reg,
metal_dev,
&logits_b,
&logits_b,
&softcap_params_b,
cap,
)
.map_err(|e| anyhow::anyhow!("encode_lm_head_into softcap: {e}"))?;
}
let gpu_sample =
if std::env::var("HF2Q_GPU_SAMPLE").as_deref() != Ok("0") && self.rerank_active() {
const CAP: usize = 1024;
let top1_idx = dev
.alloc_buffer(n * 4, DType::U32, vec![n])
.map_err(|e| anyhow::anyhow!("gpu_sample alloc top1_idx: {e}"))?;
let top1_val = dev
.alloc_buffer(n * 4, DType::F32, vec![n])
.map_err(|e| anyhow::anyhow!("gpu_sample alloc top1_val: {e}"))?;
let cand_count = dev
.alloc_buffer(n * 4, DType::U32, vec![n])
.map_err(|e| anyhow::anyhow!("gpu_sample alloc cand_count: {e}"))?;
let overflow = dev
.alloc_buffer(n * 4, DType::U32, vec![n])
.map_err(|e| anyhow::anyhow!("gpu_sample alloc overflow: {e}"))?;
let cand_ids = dev
.alloc_buffer(n * CAP * 4, DType::U32, vec![n, CAP])
.map_err(|e| anyhow::anyhow!("gpu_sample alloc cand_ids: {e}"))?;
let mut params = dev
.alloc_buffer(8, DType::U32, vec![2])
.map_err(|e| anyhow::anyhow!("gpu_sample alloc params: {e}"))?;
params
.as_mut_slice::<u32>()
.map_err(|e| anyhow::anyhow!("gpu_sample params slice: {e}"))?
.copy_from_slice(&[vocab as u32, CAP as u32]);
s.barrier_between(
&[&logits_b],
&[&top1_idx, &top1_val, &cand_count, &overflow, &cand_ids],
);
mlx_native::ops::gpu_sample::dispatch_gpu_sample_argmax_candidates(
s.encoder_mut(),
reg,
metal_dev,
&logits_b,
&top1_idx,
&top1_val,
&cand_count,
&overflow,
&cand_ids,
¶ms,
n as u32,
vocab as u32,
CAP as u32,
)
.map_err(|e| anyhow::anyhow!("gpu_sample dispatch: {e}"))?;
Some(GpuSampleBuffers {
top1_idx,
top1_val,
cand_count,
overflow,
cand_ids,
params,
cap: CAP,
})
} else {
None
};
Ok((logits_b, normed_b, softcap_params_b, gpu_sample))
}
}