use crate::Engine;
use crate::cache::{Cache, HcTapSink};
use crate::dflash::{DflashDraft, DflashKv, DsparkDraftSample};
use crate::forward::argmax;
use crate::hybrid::{HybridModel, Mixer};
use crate::spec::SpecSampling;
use crate::spec_phase::{
ProfClock, SPEC_PROF_ROUNDS, SpecFirstTokenProf, SpecPhaseNs, SpecRoundProf, SpecRoundsLog,
V_SEQ_ROWS, spec_prof_on,
};
use cudarc::driver::CudaSlice;
type Res<T> = Result<T, Box<dyn std::error::Error>>;
pub type CommitHook<'a> = &'a mut dyn FnMut(&[u32]);
pub fn glm5_spec_on() -> bool {
use std::sync::OnceLock;
static ON: OnceLock<bool> = OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_GLM5_SPEC").as_deref() == Ok("1"))
}
pub fn glm5_verify_batch_on() -> bool {
std::env::var("MEMRA_GLM5_VERIFY_BATCH").as_deref() != Ok("0")
}
pub fn glm5_draft_prime_v2_on() -> bool {
std::env::var("MEMRA_GLM5_DRAFT_PRIME_V2").as_deref() == Ok("1")
}
pub fn glm5_draft_taps_device_on() -> bool {
std::env::var("MEMRA_GLM5_DRAFT_TAPS_DEVICE").as_deref() == Ok("1")
}
struct Glm5DraftPrimeInflight {
kv: DflashKv,
taps: Vec<usize>,
n_embd: usize,
ring: usize,
stage: Vec<Option<CudaSlice<f32>>>,
rows_dev: Option<CudaSlice<f32>>,
prof_on: bool,
copy_ms: f64,
feat_ms: f64,
kv_ms: f64,
chunks: usize,
}
pub fn glm5_draft_prime_lazy_on() -> bool {
std::env::var("MEMRA_GLM5_DRAFT_PRIME_LAZY").as_deref() == Ok("1")
}
fn dflash_kv_bytes(cfg: &crate::dflash::DflashCfg, cap: usize) -> usize {
2 * cfg.n_layer * (cap + cfg.block_size) * cfg.n_kv * cfg.head_dim * std::mem::size_of::<f32>()
}
fn pinned_f32(buf: &crate::PinnedHostBuf, n: usize) -> &[f32] {
debug_assert!(n * std::mem::size_of::<f32>() <= buf.len());
unsafe { std::slice::from_raw_parts(buf.as_slice().as_ptr() as *const f32, n) }
}
fn pinned_f32_mut(buf: &mut crate::PinnedHostBuf, n: usize) -> &mut [f32] {
debug_assert!(n * std::mem::size_of::<f32>() <= buf.len());
unsafe { std::slice::from_raw_parts_mut(buf.as_mut_slice().as_mut_ptr() as *mut f32, n) }
}
pub fn glm5_spec_prefix_on() -> bool {
use std::sync::OnceLock;
static ON: OnceLock<bool> = OnceLock::new();
*ON.get_or_init(|| {
let sp = std::env::var("MEMRA_GLM5_SPEC_PREFIX").as_deref() == Ok("1");
let pl = std::env::var("MEMRA_PREFIX_LATENT").as_deref() == Ok("1");
if sp && !pl {
eprintln!(
"[glm5-spec] MEMRA_GLM5_SPEC_PREFIX=1 is INERT: it requires \
MEMRA_PREFIX_LATENT=1 (latent entries could never publish without it)"
);
}
sp && pl
})
}
pub fn glm5_spec_fullcover_on() -> bool {
std::env::var("MEMRA_GLM5_SPEC_FULLCOVER").as_deref() == Ok("1")
}
pub fn glm5_spec_tp_on() -> bool {
std::env::var("MEMRA_GLM5_SPEC_TP").as_deref() == Ok("1")
}
pub fn glm5_spec_penalty_on() -> bool {
std::env::var("MEMRA_SPEC_PENALTY").as_deref() == Ok("1")
}
pub fn glm5_spec_warm_on() -> bool {
std::env::var("MEMRA_SPEC_WARM").as_deref() == Ok("1")
}
#[derive(Clone, Copy, Debug)]
pub struct Glm5Penalty {
last_n: usize,
rep: f32,
freq: f32,
present: f32,
}
impl Glm5Penalty {
pub fn of(sp: &SpecSampling) -> Option<Self> {
sp.pen_on().then_some(Self {
last_n: sp.penalty_last_n,
rep: sp.penalty_repeat,
freq: sp.penalty_freq,
present: sp.penalty_present,
})
}
fn win(&self) -> usize {
self.last_n.min(crate::spec::PEN_WINDOW_MAX)
}
}
static PENALTY_ARM_ANNOUNCED: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
fn glm5_penalty_admit(sampling: Option<&SpecSampling>) -> Res<Option<Glm5Penalty>> {
let Some(pen) = sampling.and_then(Glm5Penalty::of) else {
return Ok(None);
};
if !glm5_spec_penalty_on() {
return Err(
"glm5 spec penalty arm is DARK (MEMRA_SPEC_PENALTY unset): penalized requests serve \
on the plain path (worker admission owns the exclusion; silently dropping the \
request's penalties is the failure class this refusal prevents)"
.into(),
);
}
if !PENALTY_ARM_ANNOUNCED.swap(true, std::sync::atomic::Ordering::AcqRel) {
eprintln!(
"[glm5-spec] penalty arm ENGAGED (MEMRA_SPEC_PENALTY=1): verify rows penalized on \
device over the session window (rep={} freq={} present={} last_n={}); printed \
once per process, every penalized request carries penalized=1 on its route line",
pen.rep, pen.freq, pen.present, pen.last_n,
);
}
Ok(Some(pen))
}
fn glm5_pen_window_seed(pen: Option<&Glm5Penalty>, committed: &[u32]) -> Vec<u32> {
match pen {
Some(p) => crate::spec::pen_window_seed(&[], committed, p.last_n),
None => Vec::new(),
}
}
#[allow(clippy::too_many_arguments)]
fn glm5_anchor(
eh: &Engine,
logits: &[f32],
sampling: Option<&SpecSampling>,
pen: Option<&Glm5Penalty>,
pen_hist: &[u32],
sctr: &mut u32,
site: &str,
) -> Res<u32> {
match (sampling, pen) {
(Some(sp), _) => crate::spec::sample_boundary_token(eh, logits, sp, pen_hist, sctr, site),
(None, Some(p)) => {
let n = logits.len();
let mut col = eh.htod(logits)?;
let w0 = pen_hist.len().saturating_sub(p.win());
let hist = &pen_hist[w0..];
if !hist.is_empty() {
let hd = eh.htod_u32_v(hist)?;
eh.penalize_logits(&mut col, &hd, hist.len(), p.rep, p.freq, p.present, n)?;
}
let td = eh.argmax_token_device(&col, n)?;
crate::spec::guard_vocab_token(
eh.dtoh_u32_one(&td)?,
n,
&format!("glm5 penalized greedy anchor ({site})"),
)
}
(None, None) => Ok(argmax(logits) as u32),
}
}
pub(crate) fn glm5_pmin() -> f32 {
use std::sync::OnceLock;
static P: OnceLock<f32> = OnceLock::new();
*P.get_or_init(|| {
std::env::var("MEMRA_SPEC_PMIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0)
})
}
pub(crate) fn glm5_pmin0() -> bool {
use std::sync::OnceLock;
static P: OnceLock<bool> = OnceLock::new();
*P.get_or_init(|| std::env::var("MEMRA_SPEC_PMIN0").as_deref() == Ok("1"))
}
pub use crate::spec::spec_conf_keep as glm5_conf_keep;
pub use crate::dflash::DflashDrafter as Glm5DflashDrafter;
static GLM5_RANK_TRIMMED_DRAFT_ROUNDS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub fn glm5_rank_trimmed_draft_rounds() -> u64 {
GLM5_RANK_TRIMMED_DRAFT_ROUNDS.load(std::sync::atomic::Ordering::Relaxed)
}
pub(crate) enum Glm5DraftState {
NativeMtp,
Dflash2 {
kv: DflashKv,
pending: Vec<f32>,
taps: Vec<usize>,
},
}
enum Glm5DraftQ {
None,
Mtp {
draft_idx: Vec<u32>,
draft_logits: Vec<CudaSlice<f32>>,
draft_stats: Vec<(f32, f32, f32)>,
},
Selector {
prop: DsparkDraftSample,
dl: CudaSlice<f32>,
},
}
fn glm5_dflash_tap_layers(draft: &DflashDraft, n_trunk: usize) -> Res<Vec<usize>> {
let shift = match std::env::var("MEMRA_GLM5_DFLASH_GATE_RED").ok().as_deref() {
Some("tap-shift") => {
eprintln!(
"[glm5-spec] RED-ARM tap-shift: drafter tap layers shifted +1 (gate \
instrument, never a serving flag)"
);
1
}
Some("") | None => 0,
Some(other) => {
return Err(format!(
"MEMRA_GLM5_DFLASH_GATE_RED={other:?}: unknown red arm (want tap-shift)"
)
.into());
}
};
Ok(crate::dflash::resolve_tap_layers(
&draft.cfg.target_layer_ids,
n_trunk,
shift,
"glm5 DFlash2",
)?)
}
pub struct Glm5VerifyCkpt {
pos: usize,
latent_len: Vec<Option<usize>>,
kda_conv_cols: Vec<Option<Vec<CudaSlice<f32>>>>,
kda_ssm_snap: Vec<Option<CudaSlice<f32>>>,
kda_scan_stash: Vec<Option<Vec<crate::kda::KdaScanInputs>>>,
kda_rows: Vec<Option<crate::kda::KdaRowsStash>>,
kda_tp: Vec<Option<crate::glm5_tp::Glm5TpKdaVerifyStash>>,
rows: usize,
}
impl Glm5VerifyCkpt {
pub fn kda_stash_kinds(&self) -> (usize, usize) {
(
self.kda_rows.iter().filter(|s| s.is_some()).count(),
self.kda_conv_cols.iter().filter(|s| s.is_some()).count(),
)
}
}
struct Glm5VerifyPos {
pos0: usize,
t: usize,
all: CudaSlice<i32>,
rows: Vec<CudaSlice<i32>>,
}
impl Glm5VerifyPos {
fn new(e: &Engine, pos0: usize, t: usize) -> Res<Self> {
let v: Vec<i32> = (0..t as i32).map(|r| pos0 as i32 + r).collect();
let all = e.htod_i32(&v)?;
let rows = if glm5_verify_batch_on() && t > 1 {
Vec::new()
} else {
(0..t)
.map(|r| e.htod_i32(&[(pos0 + r) as i32]))
.collect::<Result<_, _>>()?
};
Ok(Self { pos0, t, all, rows })
}
}
impl HybridModel {
pub fn glm5_verify_rows(
&self,
e: &Engine,
tokens: &[u32],
cache: &mut Cache,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>, Glm5VerifyCkpt)> {
let topology = *self
.hyper
.as_ref()
.ok_or("glm5_verify_rows on a model with no HyperConnections topology")?;
let t = tokens.len();
let cap = Self::hyper_batch_cap();
if t == 0 {
return Err("glm5_verify_rows: empty verify row set".into());
}
if t > cap {
return Err(format!(
"glm5_verify_rows: t={t} > cap {cap} — at t >= PRIME_MIN_T (16) the MoE \
shared-expert trio crosses off the decode-exact class (the batched-decode \
gate's measured B=16 knee), so per-row bit-identity vs plain decode breaks. \
K <= cap-1 drafts per round"
)
.into());
}
let mut any_sharded = false;
for (il, layer) in self.layers.iter().enumerate() {
match &layer.mixer {
Mixer::Kda(la) => any_sharded |= la.tp.is_some(),
Mixer::Mla(mla) => any_sharded |= mla.tp.is_some(),
_ => {
return Err(format!(
"glm5_verify_rows: trunk layer {il} is not a KDA or MLA mixer — the \
rollback contract below is built and gated for glm5_next's two state \
classes only; a Full/Linear arm needs its own ckpt plane and gate"
)
.into());
}
}
}
if any_sharded && t > 1 && !glm5_verify_batch_on() {
return Err(
"glm5_verify_rows: the trunk is glm5-TP-SHARDED and MEMRA_GLM5_VERIFY_BATCH=0 — \
the per-row rollback seam carries no TP arm; the spec x TP composition \
requires the batched verify walk (unset MEMRA_GLM5_VERIFY_BATCH or run \
without the TP door)"
.into(),
);
}
let n_embd = self.cfg.n_embd as usize;
let pos0 = cache.pos;
let mut ckpt = Glm5VerifyCkpt {
pos: pos0,
latent_len: cache
.latent
.iter()
.take(self.layers.len())
.map(|plane| plane.as_ref().map(|plane| plane.len))
.collect(),
kda_conv_cols: (0..self.layers.len()).map(|_| None).collect(),
kda_ssm_snap: (0..self.layers.len()).map(|_| None).collect(),
kda_scan_stash: (0..self.layers.len()).map(|_| None).collect(),
kda_rows: (0..self.layers.len()).map(|_| None).collect(),
kda_tp: (0..self.layers.len()).map(|_| None).collect(),
rows: t,
};
if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
return Err("pipeline rewrite is not qualified for this ModelPlan".into());
}
return self.glm5_verify_rows_ppn(e, tokens, cache, ckpt, &topology, &fence);
}
let pos = Glm5VerifyPos::new(e, pos0, t)?;
let embedded = e.htod(&self.embd.try_gather(n_embd, tokens)?)?;
let x = crate::hyper::expand(e, &topology, &embedded, t, n_embd)?;
let x = self.glm5_verify_range(
e,
&topology,
x,
0,
self.layers.len(),
&pos,
cache,
&mut ckpt,
)?;
let (logits, collapsed) = self.glm5_verify_head(e, &topology, &x, t)?;
Ok((logits, collapsed, ckpt))
}
#[allow(clippy::too_many_arguments)]
fn glm5_verify_range(
&self,
e: &Engine,
topology: &crate::hyper::HyperTopology,
mut x: CudaSlice<f32>,
lo: usize,
hi: usize,
pos: &Glm5VerifyPos,
cache: &mut Cache,
ckpt: &mut Glm5VerifyCkpt,
) -> Res<CudaSlice<f32>> {
let t = pos.t;
let n_embd = self.cfg.n_embd as usize;
let eps = self.cfg.rms_eps;
let batch = glm5_verify_batch_on() && t > 1;
{
static SAID: std::sync::Once = std::sync::Once::new();
SAID.call_once(|| {
if batch {
eprintln!(
"[glm5-spec] verify walk BATCHED per layer: kda=one t-call (scan \
sequential in-kernel), mla=rows-exact t-call, head=rows-exact, \
moe=pairs rows-call where qualified \
(MEMRA_GLM5_VERIFY_BATCH default ON)"
);
} else {
eprintln!("[glm5-spec] verify walk PER-ROW (MEMRA_GLM5_VERIFY_BATCH=0 or t=1)");
}
});
}
let trace_v = crate::spec_phase::spec_trace_level() >= 2;
let vclock = |on: bool| -> Option<std::time::Instant> {
on.then(|| {
let _ = e.stream().synchronize();
std::time::Instant::now()
})
};
for il in lo..hi {
let layer = &self.layers[il];
let hyper = layer.hyper.as_ref().ok_or_else(|| {
format!("layer {il} carries no hyper-connection weights under an hc plan")
})?;
let (y, mix) = crate::hyper::pre_exact(e, topology, &hyper.attn, &x, t, n_embd)?;
let mut h = e.uninit(t * n_embd)?;
e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
let layer_batched = batch
&& match &layer.mixer {
Mixer::Kda(_) => true,
Mixer::Mla(mla) => mla.index.is_some(),
Mixer::Full(_) | Mixer::Linear(_) => unreachable!("refused at walk entry"),
};
let mixed = if layer_batched {
match &layer.mixer {
Mixer::Kda(la) if la.tp.is_some() => {
let t0 = vclock(trace_v);
let mut scan_ns = 0u64;
let (out, stash) = crate::glm5_tp::kda_tp_verify_rows(
e,
la,
&h,
t,
eps,
cache,
il,
trace_v.then_some(&mut scan_ns),
)?;
ckpt.kda_tp[il] = Some(stash);
if let Some(t0) = t0 {
let _ = e.stream().synchronize();
use std::sync::atomic::Ordering;
crate::spec_phase::V_KDA_NS
.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
crate::spec_phase::V_KDA_SCAN_NS.fetch_add(scan_ns, Ordering::Relaxed);
}
out
}
Mixer::Mla(mla) if mla.tp.is_some() => {
let t0 = vclock(trace_v);
let out =
self.mla_tp_attn_cached(e, mla, &h, &pos.all, t, il, cache, true)?;
if let Some(t0) = t0 {
let _ = e.stream().synchronize();
use std::sync::atomic::Ordering;
crate::spec_phase::V_MLA_NS
.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
}
out
}
Mixer::Kda(la) => {
{
let rl = cache.recur[il]
.as_ref()
.ok_or("glm5 verify KDA layer has no recurrent state")?;
ckpt.kda_ssm_snap[il] = Some(e.clone_dtod(&rl.ssm_state)?);
}
let t0 = vclock(trace_v);
let mut scan_ns = 0u64;
let (out, stash) = crate::kda::kda_verify_rows_cached(
e,
la,
&h,
t,
eps,
cache,
il,
trace_v.then_some(&mut scan_ns),
)?;
ckpt.kda_rows[il] = Some(stash);
if let Some(t0) = t0 {
let _ = e.stream().synchronize();
use std::sync::atomic::Ordering;
crate::spec_phase::V_KDA_NS
.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
crate::spec_phase::V_KDA_SCAN_NS.fetch_add(scan_ns, Ordering::Relaxed);
}
out
}
Mixer::Mla(mla) => {
let t0 = vclock(trace_v);
let out =
self.mla_attn_cached_rows_exact(e, mla, &h, &pos.all, t, il, cache)?;
if let Some(t0) = t0 {
let _ = e.stream().synchronize();
use std::sync::atomic::Ordering;
crate::spec_phase::V_MLA_NS
.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
}
out
}
Mixer::Full(_) | Mixer::Linear(_) => unreachable!("refused at walk entry"),
}
} else {
V_SEQ_ROWS.fetch_add(t as u64, std::sync::atomic::Ordering::Relaxed);
let mut mixed = e.uninit(t * n_embd)?;
let mut h_row = e.uninit(n_embd)?;
#[allow(clippy::needless_range_loop)]
for r in 0..t {
e.dtod_copy_view(&h.slice(r * n_embd..(r + 1) * n_embd), &mut h_row)?;
let pos_row: CudaSlice<i32>;
let pos_r = if let Some(p) = pos.rows.get(r) {
p
} else {
pos_row = e.htod_i32(&[(pos.pos0 + r) as i32])?;
&pos_row
};
let out_row = match &layer.mixer {
Mixer::Kda(la) if la.tp.is_some() => {
if t > 1 {
return Err(format!(
"glm5 verify per-row arm reached a sharded KDA \
layer {il} at t={t}: no per-rank rollback stash \
exists on this arm (walk-entry guard bypassed?)"
)
.into());
}
crate::glm5_tp::kda_tp_cached(
e,
la,
&h_row,
1,
eps,
cache,
il,
crate::kda::ConvArm::Decode,
)?
}
Mixer::Mla(mla) if mla.tp.is_some() => {
self.mla_tp_attn_cached(e, mla, &h_row, pos_r, 1, il, cache, false)?
}
Mixer::Kda(la) => {
if r == 0 && t > 1 {
let rl = cache.recur[il]
.as_ref()
.ok_or("glm5 verify KDA layer has no recurrent state")?;
ckpt.kda_ssm_snap[il] = Some(e.clone_dtod(&rl.ssm_state)?);
}
if r + 1 < t {
let (out, inputs) = crate::kda::kda_decode_cached_stash(
e, la, &h_row, eps, cache, il,
)?;
let rl = cache.recur[il]
.as_ref()
.ok_or("glm5 verify KDA layer has no recurrent state")?;
ckpt.kda_conv_cols[il]
.get_or_insert_with(Vec::new)
.push(e.clone_dtod(&rl.conv_state)?);
ckpt.kda_scan_stash[il]
.get_or_insert_with(Vec::new)
.push(inputs);
out
} else {
crate::kda::kda_decode_cached(e, la, &h_row, eps, cache, il)?
}
}
Mixer::Mla(mla) => {
self.mla_attn_cached(e, mla, &h_row, pos_r, 1, il, cache)?
}
Mixer::Full(_) | Mixer::Linear(_) => unreachable!("refused at walk entry"),
};
e.copy_into(&mut mixed, r * n_embd, &out_row, n_embd)?;
}
mixed
};
x = crate::hyper::post(e, topology, &mixed, &x, &mix, t, n_embd)?;
let (y, mix) = crate::hyper::pre_exact(e, topology, &hyper.mlp, &x, t, n_embd)?;
let mut z = e.uninit(t * n_embd)?;
e.rms_norm(
&y,
layer.post_attn_norm.float_data(),
&mut z,
n_embd,
t,
eps,
)?;
let t0 = vclock(trace_v && batch);
let ffn_out = self.hyper_ffn_branch_batch(e, layer, &z, t, il, batch)?;
if let Some(t0) = t0 {
let _ = e.stream().synchronize();
use std::sync::atomic::Ordering;
crate::spec_phase::V_FFN_NS
.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
}
x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, t, n_embd)?;
self.glm5_hc_tap(e, cache, topology, il, &x, t)?;
}
Ok(x)
}
pub(crate) fn glm5_hc_tap(
&self,
e: &Engine,
cache: &mut Cache,
topology: &crate::hyper::HyperTopology,
il: usize,
x: &CudaSlice<f32>,
t: usize,
) -> Res<()> {
let Some(sink) = cache.hc_taps.as_mut() else {
return Ok(());
};
let base = sink.base;
self.glm5_hc_tap_into(e, sink, base, topology, il, x, t)
}
#[allow(clippy::too_many_arguments)] pub(crate) fn glm5_hc_tap_into(
&self,
e: &Engine,
sink: &mut HcTapSink,
base_abs: usize,
topology: &crate::hyper::HyperTopology,
il: usize,
x: &CudaSlice<f32>,
t: usize,
) -> Res<()> {
let Some(slot) = sink.layer_ids.iter().position(|&l| l == il) else {
return Ok(());
};
let saved = sink.base;
sink.base = base_abs;
let out = self.glm5_hc_tap_slot(e, sink, slot, topology, x, t);
sink.base = saved;
out
}
fn glm5_hc_tap_slot(
&self,
e: &Engine,
sink: &mut HcTapSink,
slot: usize,
topology: &crate::hyper::HyperTopology,
x: &CudaSlice<f32>,
t: usize,
) -> Res<()> {
let h = sink.hidden;
let n_taps = sink.layer_ids.len();
let base = sink.base.checked_sub(sink.origin).ok_or_else(|| {
format!(
"hc tap base {} below sink origin {} (walk outside the sink's window)",
sink.base, sink.origin,
)
})?;
debug_assert!(
base + t <= sink.t,
"hc tap window {base}+{t} exceeds sink {}",
sink.t
);
let contracted = crate::hyper::contract_mean(e, topology, x, t, h)?;
if sink.device_stage {
if sink.dev[slot].is_none() {
sink.dev[slot] = Some(e.uninit(sink.t * h)?);
}
let buf = sink.dev[slot].as_mut().expect("just filled");
e.copy_into(buf, base * h, &contracted, t * h)?;
return Ok(());
}
let t_dtoh = std::time::Instant::now();
let host = e.dtoh(&contracted)?;
for r in 0..t {
let dst = (base + r) * n_taps * h + slot * h;
sink.rows[dst..dst + h].copy_from_slice(&host[r * h..(r + 1) * h]);
}
sink.dtoh_ns += t_dtoh.elapsed().as_nanos() as u64;
Ok(())
}
fn glm5_tap_drain(&self, e: &Engine, sink: &mut HcTapSink) -> Res<()> {
if !sink.device_stage {
return Ok(());
}
let h = sink.hidden;
let n_taps = sink.layer_ids.len();
let split = match crate::pp::pp_cuts(self.layers.len()) {
Some(fence) if !crate::pp::pp2_streams_off() => {
Some((crate::pp::PpNRt::get(e)?, fence))
}
_ => None,
};
for slot in 0..n_taps {
let Some(buf) = sink.dev[slot].take() else {
continue;
};
let il = sink.layer_ids[slot];
let es = match split.as_ref() {
Some((rt, fence)) => {
let stage = fence
.windows(2)
.position(|w| il >= w[0] && il < w[1])
.ok_or_else(|| format!("tap layer {il} outside every stage range"))?;
rt.engine(stage, e)
}
None => e,
};
let host = es.dtoh(&buf)?;
for r in 0..sink.t {
let dst = r * n_taps * h + slot * h;
sink.rows[dst..dst + h].copy_from_slice(&host[r * h..(r + 1) * h]);
}
}
Ok(())
}
fn glm5_verify_head(
&self,
e: &Engine,
topology: &crate::hyper::HyperTopology,
x: &CudaSlice<f32>,
t: usize,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
let n_embd = self.cfg.n_embd as usize;
let eps = self.cfg.rms_eps;
let collapsed =
crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, t, n_embd)?;
let mut hn = e.uninit(t * n_embd)?;
e.rms_norm(
&collapsed,
self.output_norm.float_data(),
&mut hn,
n_embd,
t,
eps,
)?;
let logits = if glm5_verify_batch_on() && t > 1 {
e.matmul_rows_exact(&self.output, &hn, t)?
} else {
e.matmul_decode_exact(&self.output, &hn, t)?
};
Ok((logits, collapsed))
}
#[allow(clippy::too_many_arguments)]
fn glm5_verify_rows_ppn(
&self,
e: &Engine,
tokens: &[u32],
cache: &mut Cache,
mut ckpt: Glm5VerifyCkpt,
topology: &crate::hyper::HyperTopology,
fence: &[usize],
) -> Res<(CudaSlice<f32>, CudaSlice<f32>, Glm5VerifyCkpt)> {
let t = tokens.len();
let n_embd = self.cfg.n_embd as usize;
let pos0 = ckpt.pos;
let payload = t * topology.streams * n_embd;
let pos_on = |eng: &Engine| -> Res<Glm5VerifyPos> { Glm5VerifyPos::new(eng, pos0, t) };
if crate::pp::pp2_streams_off() {
let pos = pos_on(e)?;
let embedded = e.htod(&self.embd.try_gather(n_embd, tokens)?)?;
let mut x = crate::hyper::expand(e, topology, &embedded, t, n_embd)?;
x =
self.glm5_verify_range(e, topology, x, fence[0], fence[1], &pos, cache, &mut ckpt)?;
for s in 1..fence.len() - 1 {
let boundary_tx = e.clone_dtod(&x)?;
let boundary_rx = e.clone_dtod(&boundary_tx)?;
x = self.glm5_verify_range(
e,
topology,
boundary_rx,
fence[s],
fence[s + 1],
&pos,
cache,
&mut ckpt,
)?;
}
let (logits, collapsed) = self.glm5_verify_head(e, topology, &x, t)?;
return Ok((logits, collapsed, ckpt));
}
let rt = crate::pp::PpNRt::get(e)?;
let n_st = fence.len() - 1;
assert_eq!(
rt.n_stages(),
n_st,
"PpNRt stage count {} != fence stages {n_st}",
rt.n_stages()
);
rt.fence_stages_behind(&e.stream())?;
let mut slot = {
let _st0 = rt.enter(0);
let e0 = rt.engine(0, e);
let pos = pos_on(e0)?;
let embedded = e0.htod(&self.embd.try_gather(n_embd, tokens)?)?;
let x = crate::hyper::expand(e0, topology, &embedded, t, n_embd)?;
let x = self
.glm5_verify_range(e0, topology, x, fence[0], fence[1], &pos, cache, &mut ckpt)?;
rt.tx(0, &x, payload)?
};
for s in 1..n_st - 1 {
let _st = rt.enter(s);
let es = rt.engine(s, e);
let pos = pos_on(es)?;
let x = rt.rx(s - 1, slot, payload)?;
let x = self.glm5_verify_range(
es,
topology,
x,
fence[s],
fence[s + 1],
&pos,
cache,
&mut ckpt,
)?;
slot = rt.tx(s, &x, payload)?;
}
let _stl = rt.enter(n_st - 1);
let el = rt.engine(n_st - 1, e);
let pos = pos_on(el)?;
let x = rt.rx(n_st - 2, slot, payload)?;
let x = self.glm5_verify_range(
el,
topology,
x,
fence[n_st - 1],
fence[n_st],
&pos,
cache,
&mut ckpt,
)?;
let (logits, collapsed) = self.glm5_verify_head(el, topology, &x, t)?;
el.stream().synchronize()?;
drop(_stl);
self.glm5_publish_stages(e)?;
Ok((logits, collapsed, ckpt))
}
fn glm5_head_engine<'e>(&self, e: &'e Engine) -> Res<&'e Engine> {
match crate::pp::pp_cuts(self.layers.len()) {
Some(fence) if !crate::pp::pp2_streams_off() => {
let rt = crate::pp::PpNRt::get(e)?;
Ok(rt.engine(fence.len() - 2, e))
}
_ => Ok(e),
}
}
pub fn glm5_verify_rollback(
&self,
e: &Engine,
cache: &mut Cache,
ckpt: &Glm5VerifyCkpt,
keep: usize,
) -> Res<()> {
if keep == 0 || keep > ckpt.rows {
return Err(format!(
"glm5_verify_rollback: keep={keep} outside 1..={} (the anchor row is always \
committed; keep = accepted drafts + 1)",
ckpt.rows
)
.into());
}
match crate::pp::pp_cuts(self.layers.len()) {
Some(fence) if !crate::pp::pp2_streams_off() => {
let rt = crate::pp::PpNRt::get(e)?;
for s in 0..fence.len() - 1 {
let _st = rt.enter(s);
let es = rt.engine(s, e);
for il in fence[s]..fence[s + 1] {
self.glm5_rollback_layer(es, cache, ckpt, keep, il)?;
}
}
self.glm5_publish_stages(e)?;
}
_ => {
for il in 0..self.layers.len() {
self.glm5_rollback_layer(e, cache, ckpt, keep, il)?;
}
}
}
cache.pos = ckpt.pos + keep;
Ok(())
}
fn glm5_rollback_layer(
&self,
e: &Engine,
cache: &mut Cache,
ckpt: &Glm5VerifyCkpt,
keep: usize,
il: usize,
) -> Res<()> {
match &self.layers[il].mixer {
Mixer::Mla(mla) => {
if keep == ckpt.rows {
return Ok(());
}
let saved = ckpt.latent_len[il].ok_or_else(|| {
format!("glm5_verify_rollback: MLA layer {il} missing from the ckpt")
})?;
let plane = cache.latent[il].as_mut().ok_or_else(|| {
format!("glm5_verify_rollback: MLA layer {il} has no latent plane")
})?;
plane.len = saved + keep;
let len_i32 =
i32::try_from(plane.len).map_err(|_| "latent length exceeds i32 mirror")?;
e.i32_mirror_store(&mut plane.len_d, len_i32)?;
if let Some(indexer) = mla.index.as_ref() {
plane.truncate_index_pool_keys(indexer.geom.pool);
}
if let Some(tp) = mla.tp.as_ref() {
let replicas = cache.glm5_tp_latent_peer[il].as_mut().ok_or_else(|| {
format!(
"glm5_verify_rollback: sharded MLA layer {il} has no peer \
latent replicas (the TP verify walk hydrates them; a \
rollback without them would silently restore the canonical \
plane only)"
)
})?;
for (i, replica) in replicas.iter_mut().enumerate() {
replica.len = saved + keep;
tp.rt.peers[i].i32_mirror_store(&mut replica.len_d, len_i32)?;
if let Some(indexer) = mla.index.as_ref() {
replica.truncate_index_pool_keys(indexer.geom.pool);
}
}
}
}
Mixer::Kda(la) if la.tp.is_some() => {
if keep == ckpt.rows {
return Ok(()); }
let stash = ckpt.kda_tp[il].as_ref().ok_or_else(|| {
format!(
"glm5_verify_rollback: sharded KDA layer {il} has no per-rank stash \
(the batched TP verify walk fills it; the per-row arm is refused \
at walk entry)"
)
})?;
crate::glm5_tp::kda_tp_verify_rollback(e, la, stash, keep, cache, il)?;
}
Mixer::Kda(la) => {
if keep == ckpt.rows {
return Ok(()); }
if let Some(stash) = ckpt.kda_rows[il].as_ref() {
let snap = ckpt.kda_ssm_snap[il].as_ref().ok_or_else(|| {
format!("glm5_verify_rollback: KDA layer {il} has no ssm snapshot")
})?;
return crate::kda::kda_verify_rollback_rows(
e, la, snap, stash, keep, cache, il,
);
}
let conv_cols = ckpt.kda_conv_cols[il].as_ref().ok_or_else(|| {
format!("glm5_verify_rollback: KDA layer {il} has no conv columns")
})?;
let conv = &conv_cols[keep - 1];
{
let rl = cache.recur[il].as_mut().ok_or_else(|| {
format!("glm5_verify_rollback: KDA layer {il} has no recurrent state")
})?;
e.copy_into(&mut rl.conv_state, 0, conv, conv.len())?;
}
let snap = ckpt.kda_ssm_snap[il].as_ref().ok_or_else(|| {
format!("glm5_verify_rollback: KDA layer {il} has no ssm snapshot")
})?;
let stash = ckpt.kda_scan_stash[il].as_ref().ok_or_else(|| {
format!("glm5_verify_rollback: KDA layer {il} has no scan stash")
})?;
crate::kda::kda_scan_replay(e, la, snap, &stash[..keep], cache, il)?;
}
Mixer::Full(_) | Mixer::Linear(_) => {
return Err(format!(
"glm5_verify_rollback: layer {il} mixer class was refused at walk \
entry and cannot appear in a ckpt"
)
.into());
}
}
Ok(())
}
fn glm5_mtp_plane_reset(&self, e: &Engine, cache: &mut Cache, len: usize) -> Res<()> {
let e = self.glm5_head_engine(e)?;
let mtp = self
.mtp
.as_ref()
.ok_or("glm5_mtp_plane_reset with no MTP head loaded")?;
let il = self
.plan
.mtp_blocks
.first()
.ok_or("ModelPlan declares no MTP block")?
.layer
.index as usize;
let plane = cache
.latent
.get_mut(il)
.and_then(|plane| plane.as_mut())
.ok_or_else(|| format!("MTP block layer {il} has no latent cache plane"))?;
if len > plane.len {
return Err(format!(
"glm5_mtp_plane_reset: target {len} is past the plane's {} rows — a reset \
only ever shortens",
plane.len
)
.into());
}
plane.len = len;
let len_i32 = i32::try_from(len).map_err(|_| "latent length exceeds i32 mirror")?;
e.stream().memcpy_htod(&[len_i32], &mut plane.len_d)?;
if let Mixer::Mla(mla) = &mtp.mixer
&& let Some(indexer) = mla.index.as_ref()
{
plane.truncate_index_pool_keys(indexer.geom.pool);
}
Ok(())
}
pub fn generate_spec_glm5(
&self,
e: &Engine,
prompt: &[u32],
max_new: usize,
k: usize,
) -> Res<(Vec<u32>, usize, usize)> {
self.generate_spec_glm5_gated(e, prompt, max_new, k, Glm5SpecKnobs::default())
}
pub fn generate_spec_glm5_gated(
&self,
e: &Engine,
prompt: &[u32],
max_new: usize,
k: usize,
mut knobs: Glm5SpecKnobs<'_>,
) -> Res<(Vec<u32>, usize, usize)> {
let cap = Self::hyper_batch_cap();
if k == 0 || k + 1 > cap {
return Err(format!(
"generate_spec_glm5: k={k} outside 1..={} (verify rows = k+1 must stay \
inside the decode-exact knee, cap {cap})",
cap - 1
)
.into());
}
if max_new == 0 {
return Ok((Vec::new(), 0, 0));
}
let max_ctx = prompt.len() + max_new + k + 8;
let mut sess = self.glm5_spec_session_new(e, prompt, max_ctx, None)?;
let mut out: Vec<u32> = Vec::with_capacity(max_new + k);
let mut drafted = 0usize;
let mut accepted = 0usize;
while out.len() < max_new && !sess.finished() {
let (burst, d, a) = self.glm5_spec_session_burst_gated(
e,
&mut sess,
max_new - out.len(),
k,
&[],
&mut knobs,
)?;
if burst.is_empty() {
break; }
out.extend(burst);
drafted += d;
accepted += a;
}
out.truncate(max_new);
Ok((out, drafted, accepted))
}
fn glm5_mtp_plane_fill(
&self,
e: &Engine,
tokens_next: &[u32],
hiddens: &CudaSlice<f32>,
t: usize,
cache: &mut Cache,
) -> Res<()> {
let mtp = self
.mtp
.as_ref()
.ok_or("glm5_mtp_plane_fill with no MTP head loaded")?;
let il = self
.plan
.mtp_blocks
.first()
.ok_or("ModelPlan declares no MTP block")?
.layer
.index as usize;
let Mixer::Mla(mla) = &mtp.mixer else {
return Err("glm5_mtp_plane_fill serves MLA-mixer MTP blocks only".into());
};
if tokens_next.len() < t {
return Err(format!(
"glm5_mtp_plane_fill: {t} rows requested over {} successor tokens",
tokens_next.len()
)
.into());
}
let n_embd = self.cfg.n_embd as usize;
let eps = self.cfg.rms_eps;
const CHUNK: usize = 512;
let mut done = 0usize;
while done < t {
let tc = (t - done).min(CHUNK);
let e_emb = e.htod(
&self
.embd
.try_gather(n_embd, &tokens_next[done..done + tc])?,
)?;
let mut e_norm = e.uninit(tc * n_embd)?;
e.rms_norm(&e_emb, mtp.enorm.float_data(), &mut e_norm, n_embd, tc, eps)?;
let hv = e.view(hiddens, (done + tc) * n_embd);
let mut h_rows = e.uninit(tc * n_embd)?;
e.copy_view_into(
&mut h_rows,
0,
&hv.slice(done * n_embd..(done + tc) * n_embd),
tc * n_embd,
)?;
let mut h_norm = e.uninit(tc * n_embd)?;
e.rms_norm(
&h_rows,
mtp.hnorm.float_data(),
&mut h_norm,
n_embd,
tc,
eps,
)?;
let mut concat = e.uninit(tc * 2 * n_embd)?;
e.place_rows_strided(&e_norm, &mut concat, n_embd, tc, 2 * n_embd, 0)?;
e.place_rows_strided(&h_norm, &mut concat, n_embd, tc, 2 * n_embd, n_embd)?;
let inp_sa = e.matmul(&mtp.eh_proj, &concat, tc)?;
let mut a_norm = e.uninit(tc * n_embd)?;
e.rms_norm(
&inp_sa,
mtp.attn_norm.float_data(),
&mut a_norm,
n_embd,
tc,
eps,
)?;
let pos: Vec<i32> = (done as i32..(done + tc) as i32).collect();
let pos_d = e.htod_i32(&pos)?;
let _ = self.mla_attn_cached(e, mla, &a_norm, &pos_d, tc, il, cache)?;
done += tc;
}
Ok(())
}
fn glm5_seed_row(
&self,
e: &Engine,
src: &CudaSlice<f32>,
rows: usize,
row: usize,
) -> Res<CudaSlice<f32>> {
let n_embd = self.cfg.n_embd as usize;
let stack = e.view(src, rows * n_embd);
let view = stack.slice(row * n_embd..(row + 1) * n_embd);
let mut seed = e.uninit(n_embd)?;
e.copy_view_into(&mut seed, 0, &view, n_embd)?;
Ok(seed)
}
pub fn glm5_spec_session_new(
&self,
e: &Engine,
prompt: &[u32],
ctx_cap: usize,
sampling: Option<SpecSampling>,
) -> Res<Glm5SpecSession> {
if self.hyper.is_none() {
return Err("generate_spec_glm5 requires a HyperConnections trunk".into());
}
let tp_sharded = self.layers.iter().any(|l| match &l.mixer {
Mixer::Kda(la) => la.tp.is_some(),
Mixer::Mla(mla) => mla.tp.is_some(),
_ => false,
});
if tp_sharded {
if !glm5_spec_tp_on() {
return Err(
"glm5 spec is co-refused on a MEMRA_GLM5_TP-sharded model: set \
MEMRA_GLM5_SPEC_TP=1 to run the gated spec x TP composition \
(default OFF — zero real-artifact receipts; every other admission \
law still holds)"
.into(),
);
}
if !glm5_verify_batch_on() {
return Err("MEMRA_GLM5_SPEC_TP=1 requires the BATCHED verify walk \
(MEMRA_GLM5_VERIFY_BATCH must not be 0): the per-row rollback seam \
carries no TP arm"
.into());
}
}
let dflash_src = self.glm5_dflash.as_ref();
let source_kind = crate::spec::resolve_draft_source_kind(
self.plan.draft_source,
self.mtp.is_some(),
dflash_src.is_some(),
)
.map_err(|why| {
format!(
"generate_spec_glm5 cannot select a draft source ({why}). Arm one: the \
embedded MTP head (MEMRA_GLM5_MTP=1; a full MoE layer, unloaded by default) \
or the DFlash2 drafter (MEMRA_GLM5_DFLASH=<dir-or-hf-spec>)"
)
})?;
if prompt.len() < 2 {
return Err(
"generate_spec_glm5 needs a prompt of >= 2 tokens (the MTP plane warms on \
(token[i+1], hidden[i]) pairs)"
.into(),
);
}
if dflash_src.is_none() && crate::spec::spec_hpost() {
return Err(
"generate_spec_glm5 has no MEMRA_SPEC_HPOST arm: the flag flips the MTP \
carrier to the post-norm hidden, but this loop seeds every committed pair \
from the trunk's PRE-output_norm collapsed rows (LANE.md §A). Mixing the \
two silently degrades drafts; the HPOST twin needs its own gate before it \
may run"
.into(),
);
}
if crate::pp::pp_cuts(self.layers.len()).is_some()
&& !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline)
{
return Err("pipeline rewrite is not qualified for this ModelPlan".into());
}
let pen = glm5_penalty_admit(sampling.as_ref())?;
let sampling = sampling.filter(|sp| sp.temp > 0.0);
if prompt.len() + 4 > ctx_cap {
return Err(format!(
"glm5 spec session needs ctx for prompt {} + anchor + one verify round, \
cap {ctx_cap}",
prompt.len()
)
.into());
}
let n_vocab = self.output.out_features();
if let Some(map) = self.glm5_d2t() {
if map.iter().any(|&t| t as usize >= n_vocab) {
return Err(format!(
"glm5 FR-Spec d2t carries a token id >= n_vocab {n_vocab} — the ranks \
artifact was minted for a different vocabulary"
)
.into());
}
eprintln!("{}", self.glm5_trim_engagement_line(map));
}
if glm5_pmin() > 0.0 {
eprintln!(
"[glm5-spec] draft confidence gate armed: PMIN={:.3} PMIN0={} (native \
chain p-of-pick; DFlash2 selector-q tau-slot truncation)",
glm5_pmin(),
glm5_pmin0() as u8,
);
}
let mtp_il = match dflash_src {
Some(_) => None,
None => Some(
self.plan
.mtp_blocks
.first()
.ok_or("ModelPlan declares no MTP block")?
.layer
.index as usize,
),
};
let eh = self.glm5_head_engine(e)?;
let mut prof = spec_prof_on().then(SpecFirstTokenProf::default);
let mut pclk = prof.as_ref().map(|_| ProfClock::start(e, eh));
if let Some(pf) = prof.as_mut() {
pf.free_mb_before = self.glm5_free_mb(e);
}
let mut cache = crate::pp::new_cache_planned(e, &self.cfg, &self.plan, ctx_cap)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.cache_alloc_ms = ck.lap(e, eh);
}
let plen = prompt.len();
let n_embd = self.cfg.n_embd as usize;
let tap_layers = match dflash_src {
Some(dr) => Some(glm5_dflash_tap_layers(&dr.draft, self.layers.len())?),
None => None,
};
let mut v2_kv: Option<DflashKv> = None;
let (logits0, hiddens, hidden_rows) = match (dflash_src, tap_layers.as_ref()) {
(Some(dr), Some(taps)) if glm5_draft_taps_device_on() => {
let kv = DflashKv::new(eh, &dr.draft.cfg, ctx_cap)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.draft_alloc_ms = ck.lap(e, eh);
pf.draft_kv_mb = dflash_kv_bytes(&dr.draft.cfg, ctx_cap) as f64 / 1e6;
}
let ring = crate::hybrid_forward::hyper_prime_call_rows(
plen,
self.layers.len(),
self.gdn_prime_grid_on(),
);
let mut sink = HcTapSink::new_device_staged_at(taps.clone(), n_embd, ring, 0);
sink.ingest_state = Some(Box::new(Glm5DraftPrimeInflight {
kv,
taps: taps.clone(),
n_embd,
ring,
stage: (0..taps.len()).map(|_| None).collect(),
rows_dev: None,
prof_on: prof.is_some(),
copy_ms: 0.0,
feat_ms: 0.0,
kv_ms: 0.0,
chunks: 0,
}));
cache.hc_taps = Some(sink);
let (l, _seed, h) = self.prime_cache(e, prompt, &mut cache, 0)?;
let walk_ms = pclk.as_mut().map(|ck| ck.lap(e, eh));
let mut sink = cache
.hc_taps
.take()
.ok_or("device-resident drafter prime: tap sink vanished")?;
let st = sink
.ingest_state
.take()
.ok_or("device-resident drafter prime: ingest state vanished")?
.downcast::<Glm5DraftPrimeInflight>()
.map_err(|_| "device-resident drafter prime: ingest state of the wrong type")?;
let st = *st;
if st.kv.len != plen {
return Err(format!(
"device-resident drafter prime covered {} of {plen} prompt rows \
(the prime's range loop must hand every range to the ingest)",
st.kv.len
)
.into());
}
if let (Some(pf), Some(walk)) = (prof.as_mut(), walk_ms) {
let ingest = st.copy_ms + st.feat_ms + st.kv_ms;
pf.prime_ms = (walk - ingest).max(0.0);
pf.draft_prime_ms = ingest;
pf.draft_prime_h2d_ms = st.copy_ms;
pf.draft_prime_feat_ms = st.feat_ms;
pf.draft_prime_kv_ms = st.kv_ms;
pf.draft_prime_rows = plen;
pf.draft_prime_chunks = st.chunks;
pf.draft_prime_arm = "device";
}
v2_kv = Some(st.kv);
(l, h, plen)
}
(Some(dr), Some(taps)) if glm5_draft_prime_v2_on() => {
let mut kv = DflashKv::new(eh, &dr.draft.cfg, ctx_cap)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.draft_alloc_ms = ck.lap(e, eh);
pf.draft_kv_mb = dflash_kv_bytes(&dr.draft.cfg, ctx_cap) as f64 / 1e6;
}
let out = self.glm5_draft_prime_chunked(
e,
eh,
&dr.draft,
prompt,
&mut cache,
&mut kv,
taps,
prof.as_mut(),
pclk.as_mut(),
)?;
v2_kv = Some(kv);
out
}
_ => {
if let Some(taps) = tap_layers.as_ref() {
let t_sink = std::time::Instant::now();
cache.hc_taps = Some(HcTapSink::new(taps.clone(), n_embd, plen));
if let Some(pf) = prof.as_mut() {
pf.sink_alloc_ms = t_sink.elapsed().as_secs_f64() * 1e3;
}
}
let (l, _seed, h) = self.prime_cache(e, prompt, &mut cache, 0)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.prime_ms = ck.lap(e, eh);
pf.prime_tap_dtoh_ms = cache
.hc_taps
.as_ref()
.map(|sk| sk.dtoh_ns as f64 / 1e6)
.unwrap_or(0.0);
}
(l, h, plen)
}
};
let prefix_capture = if glm5_spec_prefix_on() && dflash_src.is_some() {
self.glm5_prefix_boundary_capture(e, eh, &cache, &logits0, &hiddens, plen, hidden_rows)
} else {
None
};
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.capture_ms = ck.lap(e, eh);
}
let pen_hist = glm5_pen_window_seed(pen.as_ref(), prompt);
let mut sctr = 0u32;
let anchor = glm5_anchor(
eh,
&logits0,
sampling.as_ref(),
pen.as_ref(),
&pen_hist,
&mut sctr,
"glm5-prime",
)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.anchor_ms = ck.lap(e, eh);
}
let (mut draft, pending) = match (source_kind, dflash_src, tap_layers) {
(crate::spec::DraftSourceKind::Dflash2, Some(dr), Some(taps)) => {
if let Some(kv) = v2_kv.take() {
debug_assert_eq!(kv.len, plen, "chunked drafter prime must cover the prompt");
(
Glm5DraftState::Dflash2 {
kv,
pending: Vec::new(),
taps,
},
Vec::new(),
)
} else {
let sink = cache
.hc_taps
.take()
.ok_or("glm5 dflash prime tap sink vanished")?;
let kv = DflashKv::new(eh, &dr.draft.cfg, ctx_cap)?;
if let Some(pf) = prof.as_mut() {
pf.draft_kv_mb = dflash_kv_bytes(&dr.draft.cfg, ctx_cap) as f64 / 1e6;
}
(
Glm5DraftState::Dflash2 {
kv,
pending: sink.rows,
taps,
},
Vec::new(),
)
}
}
(crate::spec::DraftSourceKind::Dflash2, _, _) => {
return Err(
"the draft-source law selected DFlash2 but this session resolved no tap \
layers — a load-path bug, refused instead of silently drafting from the \
MTP plane (a VANISHED tap sink is a different failure, caught by name in \
the Dflash2 arm itself)"
.into(),
);
}
(crate::spec::DraftSourceKind::NativeMtp, _, _) => {
self.glm5_mtp_plane_fill(eh, &prompt[1..], &hiddens, plen - 1, &mut cache)?;
let pending = vec![(anchor, self.glm5_seed_row(eh, &hiddens, plen, plen - 1)?)];
(Glm5DraftState::NativeMtp, pending)
}
};
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.draft_alloc_ms += ck.lap(e, eh);
}
if !glm5_draft_prime_lazy_on()
&& let (Glm5DraftState::Dflash2 { kv, pending, taps }, Some(dr)) =
(&mut draft, dflash_src)
&& !pending.is_empty()
{
let rows = std::mem::take(pending);
let stats = self.glm5_dflash_ingest_rows(
eh,
&dr.draft,
kv,
&rows,
taps.len() * n_embd,
pclk.as_mut(),
)?;
if let Some(pf) = prof.as_mut() {
stats.write(pf, "eager");
}
}
if let Some(pf) = prof.as_mut() {
pf.free_mb_after = self.glm5_free_mb(e);
}
if tp_sharded {
eprintln!(
"[glm5-spec] spec x TP composition ARMED (MEMRA_GLM5_SPEC_TP=1): verify \
rows ride the TP shards; rollback restores per-rank planes \
performance_claim=false"
);
}
Ok(Glm5SpecSession {
cache,
committed: prompt.to_vec(),
anchor,
anchor_emitted: false,
pending,
draft,
sampling,
pen,
pen_hist,
sctr,
uctr: 0,
rounds: 0,
rank_trimmed_rounds: 0,
done: false,
max_ctx: ctx_cap,
mtp_il,
prefix_capture,
prof_rounds: prof.as_ref().map(|_| SpecRoundsLog::default()),
prof,
})
}
pub(crate) fn glm5_taps_range_begin(&self, cache: &mut Cache, start: usize) {
if let Some(sink) = cache.hc_taps.as_mut()
&& sink.ingest_state.is_some()
{
sink.origin = start;
}
}
pub(crate) fn glm5_taps_range_done(
&self,
e: &Engine,
cache: &mut Cache,
start: usize,
end: usize,
) -> Res<()> {
let Some(sink) = cache.hc_taps.as_mut() else {
return Ok(());
};
let Some(state) = sink.ingest_state.take() else {
return Ok(());
};
let mut st = match state.downcast::<Glm5DraftPrimeInflight>() {
Ok(st) => st,
Err(other) => {
sink.ingest_state = Some(other);
return Ok(());
}
};
let res = self.glm5_taps_ingest_range(e, sink, &mut st, start, end);
sink.ingest_state = Some(st);
res
}
fn glm5_taps_ingest_range(
&self,
e: &Engine,
sink: &mut HcTapSink,
st: &mut Glm5DraftPrimeInflight,
start: usize,
end: usize,
) -> Res<()> {
let t = end - start;
if t == 0 {
return Ok(());
}
if t > st.ring {
return Err(format!(
"device-resident tap ingest: range {start}..{end} ({t} rows) exceeds the \
sink ring of {} rows",
st.ring
)
.into());
}
if st.kv.len != start {
return Err(format!(
"device-resident tap ingest: drafter KV holds {} rows but the range starts \
at {start} (a range was skipped or handed twice)",
st.kv.len
)
.into());
}
let eh = self.glm5_head_engine(e)?;
let dr = self
.glm5_dflash
.as_ref()
.ok_or("device-resident tap ingest without a loaded drafter")?;
let h = st.n_embd;
let n_taps = st.taps.len();
let mut pclk = st.prof_on.then(|| ProfClock::start(e, eh));
if st.rows_dev.is_none() {
st.rows_dev = Some(eh.uninit(st.ring * n_taps * h)?);
}
let rows_dev = st.rows_dev.as_mut().expect("just filled");
for (slot, &il) in st.taps.iter().enumerate() {
let buf = sink.dev[slot].as_ref().ok_or_else(|| {
format!(
"device-resident tap ingest: slot {slot} (layer {il}) was never written \
over rows {start}..{end}"
)
})?;
let es = self.glm5_tap_slot_engine(e, il)?;
if es.ctx().ordinal() == eh.ctx().ordinal() {
eh.copy_2d_dtod_async(rows_dev, slot * h, n_taps * h, buf, h, h, t)?;
} else {
if st.stage[slot].is_none() {
st.stage[slot] = Some(eh.uninit(st.ring * h)?);
}
let staging = st.stage[slot].as_mut().expect("just filled");
eh.copy_peer_from_async(staging, es, buf, t * h)?;
es.stream().synchronize()?;
eh.copy_2d_dtod_async(rows_dev, slot * h, n_taps * h, staging, h, h, t)?;
}
}
if let Some(ck) = pclk.as_mut() {
st.copy_ms += ck.lap(e, eh);
}
let feats = dr.draft.ctx_features(eh, rows_dev, t)?;
if let Some(ck) = pclk.as_mut() {
st.feat_ms += ck.lap(e, eh);
}
let pos: Vec<i32> = ((st.kv.len as i32)..(st.kv.len + t) as i32).collect();
dr.draft.ingest_ctx(eh, &mut st.kv, &feats, &pos, t)?;
if let Some(ck) = pclk.as_mut() {
st.kv_ms += ck.lap(e, eh);
}
st.chunks += 1;
Ok(())
}
fn glm5_dflash_ingest_rows(
&self,
eh: &Engine,
draft: &DflashDraft,
kv: &mut DflashKv,
rows: &[f32],
row_w: usize,
mut pclk: Option<&mut ProfClock>,
) -> Res<DraftIngestStats> {
debug_assert_eq!(rows.len() % row_w, 0, "ragged feature rows");
let n_new = rows.len() / row_w;
let mut st = DraftIngestStats {
rows: n_new,
..Default::default()
};
let mut r0 = 0usize;
while r0 < n_new {
let t_c = (n_new - r0).min(256);
let chunk = eh.htod(&rows[r0 * row_w..(r0 + t_c) * row_w])?;
if let Some(ck) = pclk.as_deref_mut() {
st.h2d_ms += ck.lap(eh, eh);
}
let feats = draft.ctx_features(eh, &chunk, t_c)?;
if let Some(ck) = pclk.as_deref_mut() {
st.feat_ms += ck.lap(eh, eh);
}
let pos_c: Vec<i32> = ((kv.len as i32)..(kv.len + t_c) as i32).collect();
draft.ingest_ctx(eh, kv, &feats, &pos_c, t_c)?;
if let Some(ck) = pclk.as_deref_mut() {
st.kv_ms += ck.lap(eh, eh);
}
r0 += t_c;
st.chunks += 1;
}
Ok(st)
}
fn glm5_free_mb(&self, e: &Engine) -> Vec<(usize, u64)> {
let mut out = vec![(e.ctx().ordinal(), e.free_mem_mb())];
if let Ok(rt) = crate::pp::PpNRt::get(e) {
for stage in 0..rt.n_stages() {
let se = rt.engine(stage, e);
let ord = se.ctx().ordinal();
if !out.iter().any(|(o, _)| *o == ord) {
out.push((ord, se.free_mem_mb()));
}
}
}
out
}
fn glm5_tap_slot_engine<'e>(&self, e: &'e Engine, il: usize) -> Res<&'e Engine> {
match crate::pp::pp_cuts(self.layers.len()) {
Some(fence) if !crate::pp::pp2_streams_off() => {
let rt = crate::pp::PpNRt::get(e)?;
let stage = fence
.windows(2)
.position(|w| il >= w[0] && il < w[1])
.ok_or_else(|| format!("tap layer {il} outside every stage range"))?;
Ok(rt.engine(stage, e))
}
_ => Ok(e),
}
}
#[allow(clippy::too_many_arguments)]
fn glm5_draft_prime_chunked(
&self,
e: &Engine,
eh: &Engine,
draft: &DflashDraft,
prompt: &[u32],
cache: &mut Cache,
kv: &mut DflashKv,
taps: &[usize],
mut prof: Option<&mut SpecFirstTokenProf>,
mut pclk: Option<&mut ProfClock>,
) -> Res<(Vec<f32>, CudaSlice<f32>, usize)> {
let plen = prompt.len();
let n_embd = self.cfg.n_embd as usize;
let n_taps = taps.len();
let ranges = crate::hybrid_forward::hyper_prime_ranges(
plen,
self.layers.len(),
self.gdn_prime_grid_on(),
);
let t_max = ranges.iter().map(|&(a, b)| b - a).max().unwrap_or(0);
let f32b = std::mem::size_of::<f32>();
let mut slot_buf = crate::PinnedHostBuf::new(t_max * n_embd * f32b)?;
let mut rows_buf = crate::PinnedHostBuf::new(t_max * n_taps * n_embd * f32b)?;
let (mut prime_ms, mut h2d_ms, mut feat_ms, mut kv_ms) = (0f64, 0f64, 0f64, 0f64);
let mut last: Option<(Vec<f32>, CudaSlice<f32>, usize)> = None;
for &(start, end) in &ranges {
let t = end - start;
cache.hc_taps = Some(HcTapSink::new_device_staged_at(
taps.to_vec(),
n_embd,
t,
start,
));
let (l, _seed, h) = self.prime_cache(e, &prompt[start..end], cache, plen - end)?;
if let Some(ck) = pclk.as_deref_mut() {
prime_ms += ck.lap(e, eh);
}
let mut sink = cache
.hc_taps
.take()
.ok_or("chunked drafter prime: chunk tap sink vanished")?;
for (slot, &il) in taps.iter().enumerate() {
let buf = sink.dev[slot].take().ok_or_else(|| {
format!(
"chunked drafter prime: tap slot {slot} (layer {il}) was never \
written by the prime walk over rows {start}..{end}"
)
})?;
let es = self.glm5_tap_slot_engine(e, il)?;
es.dtoh_f32_into_pinned(&buf, &mut slot_buf, t * n_embd)?;
let src = pinned_f32(&slot_buf, t * n_embd);
let dst = pinned_f32_mut(&mut rows_buf, t * n_taps * n_embd);
for r in 0..t {
let d0 = (r * n_taps + slot) * n_embd;
dst[d0..d0 + n_embd].copy_from_slice(&src[r * n_embd..(r + 1) * n_embd]);
}
}
let feats_in = eh.htod_f32_from_pinned_async(&rows_buf, t * n_taps * n_embd)?;
if let Some(ck) = pclk.as_deref_mut() {
h2d_ms += ck.lap(e, eh);
}
let feats = draft.ctx_features(eh, &feats_in, t)?;
if let Some(ck) = pclk.as_deref_mut() {
feat_ms += ck.lap(e, eh);
}
let pos: Vec<i32> = ((kv.len as i32)..(kv.len + t) as i32).collect();
draft.ingest_ctx(eh, kv, &feats, &pos, t)?;
eh.stream().synchronize()?;
if let Some(ck) = pclk.as_deref_mut() {
kv_ms += ck.lap(e, eh);
}
last = Some((l, h, t));
}
debug_assert_eq!(kv.len, plen, "chunked drafter prime must cover the prompt");
if let Some(pf) = prof.as_mut() {
pf.prime_ms = prime_ms;
pf.draft_prime_ms = h2d_ms + feat_ms + kv_ms;
pf.draft_prime_h2d_ms = h2d_ms;
pf.draft_prime_feat_ms = feat_ms;
pf.draft_prime_kv_ms = kv_ms;
pf.draft_prime_rows = plen;
pf.draft_prime_chunks = ranges.len();
pf.draft_prime_arm = "chunked";
}
last.ok_or_else(|| "chunked drafter prime: empty prime schedule".into())
}
#[allow(clippy::too_many_arguments)]
fn glm5_prefix_boundary_capture(
&self,
e: &Engine,
eh: &Engine,
cache: &Cache,
logits0: &[f32],
hiddens: &CudaSlice<f32>,
plen: usize,
hidden_rows: usize,
) -> Option<crate::spec::SpecBoundaryCapture> {
debug_assert_eq!(
cache.pos, plen,
"boundary capture must sit at the prime boundary"
);
let snap = match cache.snapshot(e) {
Ok(s) => s,
Err(err) => {
eprintln!("[glm5-spec] prefix boundary capture SKIPPED (cache snapshot: {err})");
return None;
}
};
let mtp_plane_il = self.plan.mtp_blocks.first().map(|b| b.layer.index as usize);
let mut latent_tails = Vec::with_capacity(cache.latent.len());
for (il, l) in cache.latent.iter().enumerate() {
match l {
Some(l) if l.len == 0 && cache.pos > 0 => {
if Some(il) != mtp_plane_il {
eprintln!(
"[glm5-spec] prefix boundary capture SKIPPED (trunk latent \
layer {il} is EMPTY at the boundary — not the MTP plane; a \
capture would publish an absent history for a live layer)"
);
return None;
}
latent_tails.push(None)
}
Some(l) => {
if l.len != plen {
eprintln!(
"[glm5-spec] prefix boundary capture SKIPPED (latent layer {il} \
len {} != boundary {plen})",
l.len,
);
return None;
}
match l.snapshot_tail(e) {
Ok(t) => latent_tails.push(Some(t)),
Err(err) => {
eprintln!(
"[glm5-spec] prefix boundary capture SKIPPED (latent layer \
{il}: {err})"
);
return None;
}
}
}
None => latent_tails.push(None),
}
}
let last_h = crate::spec::capture_boundary_hidden(
eh,
hiddens,
hidden_rows,
self.cfg.n_embd as usize,
);
Some(crate::spec::SpecBoundaryCapture {
snap,
pos: plen,
logits: logits0.to_vec(),
last_h,
latent_tails,
})
}
#[allow(clippy::too_many_arguments)]
pub fn glm5_spec_session_from_restored(
&self,
e: &Engine,
mut cache: Cache,
fed: &[u32],
suffix: &[u32],
boundary_logits: &[f32],
dkv: DflashKv,
ctx_cap: usize,
sampling: Option<SpecSampling>,
) -> Res<Glm5SpecSession> {
if self.hyper.is_none() {
return Err("glm5_spec_session_from_restored requires a HyperConnections trunk".into());
}
let tp_sharded = self.layers.iter().any(|l| match &l.mixer {
Mixer::Kda(la) => la.tp.is_some(),
Mixer::Mla(mla) => mla.tp.is_some(),
_ => false,
});
if tp_sharded {
return Err(
"restored glm5 spec sessions carry no TP arm (the spec x TP composition is \
cold-session gated only); the plain hit serves"
.into(),
);
}
if crate::pp::pp_cuts(self.layers.len()).is_some()
&& !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline)
{
return Err("pipeline rewrite is not qualified for this ModelPlan".into());
}
let dr = self.glm5_dflash.as_ref().ok_or(
"restored glm5 spec sessions require the DFlash2 drafter (MEMRA_GLM5_DFLASH): \
the native MTP plane cannot be re-warmed from restored KV",
)?;
let source = crate::spec::resolve_draft_source_kind(
self.plan.draft_source,
self.mtp.is_some(),
true,
)
.map_err(|why| format!("restored glm5 spec session has no draft source ({why})"))?;
if !matches!(source, crate::spec::DraftSourceKind::Dflash2) {
return Err(
"restored glm5 spec sessions require the DFlash2 draft source; the plan \
selected another"
.into(),
);
}
let pen = glm5_penalty_admit(sampling.as_ref())?;
let sampling = sampling.filter(|sp| sp.temp > 0.0);
if fed.is_empty() {
return Err("restored glm5 spec session needs a non-empty restored prefix".into());
}
if suffix.is_empty() {
if !glm5_spec_fullcover_on() {
return Err(
"restored glm5 spec session needs a non-empty suffix unless \
MEMRA_GLM5_SPEC_FULLCOVER=1 (empty-suffix full-cover hits otherwise \
keep the plain boundary-logits resume)"
.into(),
);
}
if boundary_logits.len() != self.output.out_features() {
return Err(format!(
"full-cover glm5 spec restore needs the entry's boundary logits \
({} rows, got {})",
self.output.out_features(),
boundary_logits.len(),
)
.into());
}
}
if cache.pos != fed.len() {
return Err(format!(
"restored glm5 spec session needs a whole-entry trunk cache: cache.pos {} \
!= restored prefix {}",
cache.pos,
fed.len(),
)
.into());
}
if dkv.len != fed.len() {
return Err(format!(
"restored draft KV len {} != restored prefix {}",
dkv.len,
fed.len(),
)
.into());
}
if dkv.cap != ctx_cap {
return Err(
format!("restored draft KV cap {} != session ctx {ctx_cap}", dkv.cap,).into(),
);
}
if fed.len() + suffix.len() + 4 > ctx_cap {
return Err(format!(
"restored glm5 spec session needs ctx for prefix {} + suffix {} + anchor + \
one verify round, cap {ctx_cap}",
fed.len(),
suffix.len(),
)
.into());
}
let n_vocab = self.output.out_features();
if let Some(map) = self.glm5_d2t() {
if map.iter().any(|&t| t as usize >= n_vocab) {
return Err(format!(
"glm5 FR-Spec d2t carries a token id >= n_vocab {n_vocab} — the ranks \
artifact was minted for a different vocabulary"
)
.into());
}
eprintln!("{}", self.glm5_trim_engagement_line(map));
}
if glm5_pmin() > 0.0 {
eprintln!(
"[glm5-spec] draft confidence gate armed: PMIN={:.3} PMIN0={} (native \
chain p-of-pick; DFlash2 selector-q tau-slot truncation)",
glm5_pmin(),
glm5_pmin0() as u8,
);
}
let n_embd = self.cfg.n_embd as usize;
let taps = glm5_dflash_tap_layers(&dr.draft, self.layers.len())?;
let eh = self.glm5_head_engine(e)?;
crate::pp::PpNRt::order_engine_behind(e, eh)?;
let mut prof = spec_prof_on().then(SpecFirstTokenProf::default);
let mut pclk = prof.as_ref().map(|_| ProfClock::start(e, eh));
let (logits_s, tap_rows, prefix_capture) = if suffix.is_empty() {
(boundary_logits.to_vec(), Vec::new(), None)
} else {
cache.hc_taps = Some(HcTapSink::new_at(
taps.clone(),
n_embd,
suffix.len(),
fed.len(),
));
let (logits_s, _seed, hiddens) = self.prime_cache(e, suffix, &mut cache, 0)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.prime_ms = ck.lap(e, eh);
}
let capture = if glm5_spec_prefix_on() {
self.glm5_prefix_boundary_capture(
e,
eh,
&cache,
&logits_s,
&hiddens,
cache.pos,
suffix.len(),
)
} else {
None
};
let sink = cache
.hc_taps
.take()
.ok_or("glm5 restored-session suffix tap sink vanished")?;
(logits_s, sink.rows, capture)
};
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.capture_ms = ck.lap(e, eh);
}
let mut committed = Vec::with_capacity(fed.len() + suffix.len());
committed.extend_from_slice(fed);
committed.extend_from_slice(suffix);
let pen_hist = glm5_pen_window_seed(pen.as_ref(), &committed);
let mut sctr = 0u32;
let anchor = glm5_anchor(
eh,
&logits_s,
sampling.as_ref(),
pen.as_ref(),
&pen_hist,
&mut sctr,
"glm5-restore",
)?;
if let (Some(pf), Some(ck)) = (prof.as_mut(), pclk.as_mut()) {
pf.anchor_ms = ck.lap(e, eh);
}
eprintln!(
"[glm5-spec] RESTORED session: {} prefix tokens + {} suffix from cache — no \
cold prime (drafter tail rows {}, arm {})",
fed.len(),
suffix.len(),
dkv.len,
if suffix.is_empty() {
"full-cover"
} else {
"suffix-prime"
},
);
Ok(Glm5SpecSession {
cache,
committed,
anchor,
anchor_emitted: false,
pending: Vec::new(),
draft: Glm5DraftState::Dflash2 {
kv: dkv,
pending: tap_rows,
taps,
},
sampling,
pen,
pen_hist,
sctr,
uctr: 0,
rounds: 0,
rank_trimmed_rounds: 0,
done: false,
max_ctx: ctx_cap,
mtp_il: None,
prefix_capture,
prof_rounds: prof.as_ref().map(|_| SpecRoundsLog::default()),
prof,
})
}
fn glm5_d2t(&self) -> Option<&[u32]> {
if self.glm5_dflash.is_some() {
return self.glm5_dflash_trim().map(|(_, d2t)| d2t);
}
self.mtp
.as_ref()
.and_then(|head| head.d2t.as_deref())
.filter(|map| !map.is_empty())
}
pub fn glm5_dflash_trim(&self) -> Option<(&crate::model::GpuTensor, &[u32])> {
self.mtp
.as_ref()
.filter(|m| m.d2t_from_target_head)
.and_then(|m| m.shared_head_head.as_ref().zip(m.d2t.as_deref()))
.or_else(|| {
self.dflash_trim
.as_ref()
.map(|t| (&t.head, t.d2t.as_slice()))
})
.filter(|(_, d2t)| !d2t.is_empty())
}
fn glm5_trim_engagement_line(&self, map: &[u32]) -> String {
if self.glm5_dflash.is_some() {
format!(
"[glm5-spec] draft head RANK-TRIMMED n_ranks={} src={}",
map.len(),
self.frspec_src_sha16.as_deref().unwrap_or("unknown")
)
} else {
format!(
"[glm5-spec] draft head TRIMMED to {} rows (FR-Spec d2t engaged)",
map.len()
)
}
}
pub fn glm5_spec_session_burst(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
target: usize,
k: usize,
eos: &[u32],
) -> Res<(Vec<u32>, usize, usize)> {
self.glm5_spec_session_burst_inner(
e,
sess,
target,
k,
eos,
&mut Glm5SpecKnobs::default(),
None,
)
}
pub fn glm5_spec_session_burst_streamed(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
target: usize,
k: usize,
eos: &[u32],
on_commit: CommitHook<'_>,
) -> Res<(Vec<u32>, usize, usize)> {
self.glm5_spec_session_burst_inner(
e,
sess,
target,
k,
eos,
&mut Glm5SpecKnobs::default(),
Some(on_commit),
)
}
pub fn glm5_spec_session_burst_gated(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
target: usize,
k: usize,
eos: &[u32],
knobs: &mut Glm5SpecKnobs<'_>,
) -> Res<(Vec<u32>, usize, usize)> {
self.glm5_spec_session_burst_inner(e, sess, target, k, eos, knobs, None)
}
#[allow(clippy::too_many_arguments)]
fn glm5_spec_session_burst_inner(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
target: usize,
k: usize,
eos: &[u32],
knobs: &mut Glm5SpecKnobs<'_>,
mut on_commit: Option<CommitHook<'_>>,
) -> Res<(Vec<u32>, usize, usize)> {
let cap = Self::hyper_batch_cap();
if k == 0 || k + 1 > cap {
return Err(format!(
"glm5_spec_session_burst: k={k} outside 1..={} (verify rows = k+1 must stay \
inside the decode-exact knee, cap {cap})",
cap - 1
)
.into());
}
if let Glm5DraftState::Dflash2 { .. } = sess.draft {
let b = self
.glm5_dflash
.as_ref()
.ok_or("dflash session on a model with no loaded drafter")?
.draft
.cfg
.block_size;
if k + 1 > b {
return Err(format!(
"glm5_spec_session_burst: k={k} exceeds the DFlash2 drafter's block \
(block_size {b} = anchor + {} drafts, the trained mask pattern) — \
the worker clamps operator K pins to {} for this source; refusing \
loudly rather than drafting an untrained shape",
b - 1,
b - 1
)
.into());
}
}
let d2t = self.glm5_d2t();
if d2t.is_some() && knobs.skip_d2t_remap {
eprintln!("[glm5-spec] d2t REMAP SKIPPED — red-arm instrument, drafts are rank ids");
}
let sp_on: Option<SpecSampling> = sess.sampling.filter(|sp| sp.temp > 0.0);
let mut out: Vec<u32> = Vec::with_capacity(target + k);
let mut drafted = 0usize;
let mut accepted = 0usize;
let mut phase: Option<SpecPhaseNs> =
crate::spec_phase::spec_trace_on().then(SpecPhaseNs::default);
let first_burst = (!sess.anchor_emitted && sess.prof.is_some())
.then(|| (std::time::Instant::now(), sess.rounds));
let mut hook_ns: u64 = 0;
let mut flushed = 0usize;
if !sess.anchor_emitted {
out.push(sess.anchor);
sess.anchor_emitted = true;
if eos.contains(&sess.anchor) {
sess.done = true;
}
if let Some(cb) = on_commit.as_mut() {
let t_cb = first_burst.map(|_| std::time::Instant::now());
cb(&out[flushed..]);
if let Some(t_cb) = t_cb {
hook_ns += t_cb.elapsed().as_nanos() as u64;
}
flushed = out.len();
}
}
while out.len() < target && !sess.done {
if sess.cache.pos + k + 2 > sess.max_ctx {
sess.done = true;
break;
}
let log_round = sess.prof_rounds.as_ref().is_some_and(|l| l.wants_more());
let mut rp = log_round.then(SpecPhaseNs::default);
let t_round = log_round.then(std::time::Instant::now);
let seq0 = V_SEQ_ROWS.load(std::sync::atomic::Ordering::Relaxed);
let ctx0 = sess.cache.pos;
let (round_tokens, n_drafted) = {
let ph: Option<&mut SpecPhaseNs> = match (rp.as_mut(), phase.as_mut()) {
(Some(r), _) => Some(r),
(None, p) => p,
};
self.glm5_spec_round(e, sess, k, d2t, sp_on.as_ref(), knobs, ph)?
};
if let (Some(r), Some(t0)) = (rp.as_ref(), t_round) {
if let Some(p) = phase.as_mut() {
p.add(r);
}
if let Some(log) = sess.prof_rounds.as_mut() {
let ms = |ns: u64| ns as f32 / 1e6;
log.push(SpecRoundProf {
wall_ms: t0.elapsed().as_secs_f32() * 1e3,
draft_ms: ms(r.draft),
verify_ms: ms(r.verify),
accept_ms: ms(r.accept),
rest_ms: ms(r.roll + r.maint),
k: n_drafted as u16,
j: (round_tokens.len() - 1) as u16,
ctx: ctx0 as u32,
seq_rows: (V_SEQ_ROWS.load(std::sync::atomic::Ordering::Relaxed) - seq0)
as u32,
});
}
}
drafted += n_drafted;
accepted += round_tokens.len() - 1; for &tok in &round_tokens {
if eos.contains(&tok) {
sess.done = true;
}
}
out.extend_from_slice(&round_tokens);
sess.rounds += 1;
if let Some(cb) = on_commit.as_mut() {
let t_cb = first_burst.map(|_| std::time::Instant::now());
cb(&out[flushed..]);
if let Some(t_cb) = t_cb {
hook_ns += t_cb.elapsed().as_nanos() as u64;
}
flushed = out.len();
}
}
debug_assert!(
on_commit.is_none() || flushed == out.len(),
"every committed token must have been handed to on_commit"
);
if let Some(ph) = phase.as_ref() {
ph.emit("glm5-phase", "glm5-phase-v", k);
}
if let (Some((t0, rounds0)), Some(pf)) = (first_burst, sess.prof.as_mut()) {
pf.first_burst_ms = t0.elapsed().as_secs_f64() * 1e3;
pf.first_burst_hook_ms = hook_ns as f64 / 1e6;
pf.first_burst_rounds = sess.rounds - rounds0;
pf.first_burst_tokens = out.len();
}
Ok((out, drafted, accepted))
}
#[allow(clippy::too_many_arguments)]
fn glm5_spec_round(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
k: usize,
d2t: Option<&[u32]>,
sp: Option<&SpecSampling>,
knobs: &mut Glm5SpecKnobs<'_>,
mut phase: Option<&mut SpecPhaseNs>,
) -> Res<(Vec<u32>, usize)> {
let n_vocab = self.output.out_features();
let n_embd = self.cfg.n_embd as usize;
let eh = self.glm5_head_engine(e)?;
let mut t_mark = phase.as_ref().map(|_| SpecPhaseNs::clock(e, eh));
let mut pclk = (sess.rounds == 0 && sess.prof.is_some()).then(|| ProfClock::start(e, eh));
macro_rules! bump {
($field:ident) => {
if let (Some(ph), Some(t0)) = (phase.as_deref_mut(), t_mark.as_mut()) {
let now = SpecPhaseNs::clock(e, eh);
ph.$field += now.duration_since(*t0).as_nanos() as u64;
*t0 = now;
}
};
}
macro_rules! plap {
($field:ident) => {
if let (Some(ck), Some(pf)) = (pclk.as_mut(), sess.prof.as_mut()) {
pf.$field = ck.lap(e, eh);
}
};
}
let (p_min, pmin0) = knobs
.pmin_override
.unwrap_or_else(|| (glm5_pmin(), glm5_pmin0()));
let (drafts, qside, mtp_committed_len) = match sess.draft {
Glm5DraftState::Dflash2 { .. } => {
let (d, q) = self.glm5_dflash_round_drafts(
eh,
sess,
k,
sp,
knobs,
p_min,
pmin0,
pclk.as_mut(),
)?;
(d, q, 0)
}
Glm5DraftState::NativeMtp => {
let mtp_il = sess.mtp_il.ok_or("native-mtp arm without a plane index")?;
let mut last: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
for (tok, h) in sess.pending.drain(..) {
let plane_len = sess.cache.latent[mtp_il]
.as_ref()
.ok_or("MTP plane missing")?
.len;
last = Some(self.mtp_head_forward_mla_cached(
eh,
0,
tok,
&h,
&mut sess.cache,
plane_len,
)?);
}
let (mut d_logits, mut carrier) =
last.ok_or("glm5 spec round started with no pending committed pair")?;
let mtp_committed_len = sess.cache.latent[mtp_il]
.as_ref()
.ok_or("MTP plane missing")?
.len;
let d_vocab = d2t.map(|m| m.len()).unwrap_or(n_vocab);
let mut drafts: Vec<u32> = Vec::with_capacity(k); let mut draft_idx: Vec<u32> = Vec::with_capacity(k); let mut draft_logits: Vec<CudaSlice<f32>> = Vec::new(); let mut draft_stats: Vec<(f32, f32, f32)> = Vec::new(); for ki in 0..k {
let (idx, sampled_stats) = match sp {
Some(sp) => {
let (idx, stats) =
glm5_sampled_draft(eh, &d_logits, d_vocab, sp, &mut sess.sctr)?;
(idx, Some(stats))
}
None => {
let td = eh.argmax_token_device(&d_logits, d_vocab)?;
let idx = crate::spec::guard_vocab_token(
eh.dtoh_u32_one(&td)?,
d_vocab,
&format!(
"glm5 native draft argmax at round {} ki={ki}",
sess.rounds
),
)?;
(idx, None)
}
};
if p_min > 0.0 {
let tok_d = eh.htod_u32_v(&[idx])?;
let p_d = eh.prob_of_token_device(&d_logits, &tok_d, d_vocab)?;
let p = eh.dtoh(&p_d)?[0];
if p < p_min && (ki > 0 || pmin0) {
break;
}
}
if let Some(stats) = sampled_stats {
draft_stats.push(stats);
draft_logits.push(eh.clone_dtod(&d_logits)?);
}
let mut d = match d2t {
Some(map) if !knobs.skip_d2t_remap => map[idx as usize],
_ => idx,
};
if let Some(over) = knobs.draft_override.as_mut() {
d = over(sess.rounds, ki, d);
}
drafts.push(d);
draft_idx.push(idx);
if ki + 1 < k {
let plane_len = sess.cache.latent[mtp_il]
.as_ref()
.ok_or("MTP plane missing")?
.len;
let (lg, ca) = self.mtp_head_forward_mla_cached(
eh,
0,
d,
&carrier,
&mut sess.cache,
plane_len,
)?;
d_logits = lg;
carrier = ca;
}
}
let q = match sp {
Some(_) => Glm5DraftQ::Mtp {
draft_idx,
draft_logits,
draft_stats,
},
None => Glm5DraftQ::None,
};
(drafts, q, mtp_committed_len)
}
};
bump!(draft);
plap!(first_draft_ms);
if d2t.is_some() {
sess.rank_trimmed_rounds += 1;
GLM5_RANK_TRIMMED_DRAFT_ROUNDS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
if let Glm5DraftState::Dflash2 { taps, .. } = &sess.draft {
sess.cache.hc_taps = Some(HcTapSink::new_device_staged(
taps.clone(),
n_embd,
drafts.len() + 1,
));
}
let mut rows: Vec<u32> = Vec::with_capacity(drafts.len() + 1);
rows.push(sess.anchor);
rows.extend_from_slice(&drafts);
let (vlogits, collapsed, ckpt) = self.glm5_verify_rows(e, &rows, &mut sess.cache)?;
bump!(verify);
plap!(first_verify_ms);
let pvl: Option<CudaSlice<f32>> = match sess.pen {
Some(p) => {
sess.pen_hist.push(sess.anchor);
let w0 = sess.pen_hist.len().saturating_sub(p.win());
let pen_win = &sess.pen_hist[w0..];
let mut hist: Vec<u32> = Vec::with_capacity(pen_win.len() + drafts.len());
hist.extend_from_slice(pen_win);
hist.extend_from_slice(&drafts);
let n_win = pen_win.len();
let hd = eh.htod_u32_v(&hist)?;
let mut buf = eh.clone_dtod(&vlogits)?;
eh.penalize_logits_rows_inc(
&mut buf,
&hd,
n_win,
p.rep,
p.freq,
p.present,
n_vocab,
rows.len(),
p.win(),
)?;
Some(buf)
}
None => None,
};
let plogits: &CudaSlice<f32> = pvl.as_ref().unwrap_or(&vlogits);
let sp_accept: Option<SpecSampling> = sp.map(|s| match sess.pen {
Some(_) => SpecSampling {
penalty_last_n: 0,
..*s
},
None => *s,
});
let (j, bonus) = if let (true, Some(sp)) = (drafts.is_empty(), sp) {
(
0,
self.glm5_sampled_bonus(eh, sess, sp, plogits, 0, n_vocab)?,
)
} else {
match (sp, &qside) {
(None, _) => {
let t = rows.len();
let mut vam_d = eh.alloc_u32_zeroed(t)?;
for r in 0..t {
eh.argmax_token_device_col(plogits, r, n_vocab, &mut vam_d, r)?;
}
let vam = eh.dtoh_u32(&vam_d)?;
let mut j = 0usize;
while j < drafts.len() && drafts[j] == vam[j] {
j += 1;
}
if knobs.accept_probe {
self.glm5_accept_probe(eh, sess.rounds, plogits, &drafts, &vam, j)?;
}
(j, vam[j])
}
(
Some(sp),
Glm5DraftQ::Mtp {
draft_idx,
draft_logits,
draft_stats,
},
) => self.glm5_sampled_accept(
eh,
sess,
sp,
plogits,
&drafts,
draft_idx,
draft_logits,
draft_stats,
d2t,
drafts.len(),
)?,
(Some(_), Glm5DraftQ::Selector { prop, dl }) => {
let (m, next) = crate::dflash::dspark_accept_sampled(
eh,
plogits,
&rows,
rows.len(),
n_vocab,
dl,
prop,
sp_accept.as_ref().expect("sampled arm carries its config"),
&[],
&mut sess.sctr,
&mut sess.uctr,
)?;
let next = crate::spec::guard_vocab_token(
next,
n_vocab,
&format!(
"glm5 dflash2 sampled verify bonus at round {} j={m}",
sess.rounds
),
)?;
(m, next)
}
(Some(_), Glm5DraftQ::None) => {
unreachable!("sampled round without a retained q side")
}
}
};
if sess.pen.is_some() {
sess.pen_hist.extend_from_slice(&drafts[..j]);
}
bump!(accept);
plap!(first_accept_ms);
let mut round_tokens: Vec<u32> = Vec::with_capacity(j + 1);
round_tokens.extend_from_slice(&drafts[..j]);
round_tokens.push(bonus);
let keep = j + 1;
if knobs.disable_rollback {
sess.cache.pos = ckpt.pos + keep;
} else {
self.glm5_verify_rollback(e, &mut sess.cache, &ckpt, keep)?;
}
bump!(roll);
plap!(first_roll_ms);
match &mut sess.draft {
Glm5DraftState::NativeMtp => {
self.glm5_mtp_plane_reset(e, &mut sess.cache, mtp_committed_len)?;
for i in 1..=keep {
let tok = round_tokens[i - 1];
let h = self.glm5_seed_row(eh, &collapsed, rows.len(), i - 1)?;
sess.pending.push((tok, h));
}
}
Glm5DraftState::Dflash2 { pending, taps, .. } => {
let mut sink = sess
.cache
.hc_taps
.take()
.ok_or("glm5 dflash verify tap sink vanished")?;
self.glm5_tap_drain(e, &mut sink)?;
let row_w = taps.len() * n_embd;
pending.extend_from_slice(&sink.rows[..keep * row_w]);
}
}
sess.committed.push(sess.anchor);
sess.committed.extend_from_slice(&drafts[..j]);
sess.anchor = bonus;
bump!(maint);
plap!(first_maint_ms);
if let Some(pf) = sess.prof.as_mut()
&& pclk.is_some()
{
pf.first_round_tokens = round_tokens.len();
}
if let Some(ph) = phase {
ph.rounds += 1;
}
Ok((round_tokens, drafts.len()))
}
fn glm5_publish_stages(&self, e: &Engine) -> Res<()> {
if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
let rt = crate::pp::PpNRt::get(e)?;
let dst = e.stream();
rt.publish_all_to(&dst)?;
}
Ok(())
}
fn glm5_accept_probe(
&self,
eh: &Engine,
round: usize,
vlogits: &CudaSlice<f32>,
drafts: &[u32],
vam: &[u32],
j: usize,
) -> Res<()> {
let n_vocab = self.output.out_features();
let t = vam.len();
let host = eh.dtoh(vlogits)?;
let mut hvam: Vec<u32> = Vec::with_capacity(t);
let mut rows_census: Vec<String> = Vec::with_capacity(t);
for r in 0..t {
let row = &host[r * n_vocab..(r + 1) * n_vocab];
let am = argmax(row) as u32;
hvam.push(am);
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for v in row {
for b in v.to_bits().to_le_bytes() {
h ^= u64::from(b);
h = h.wrapping_mul(0x100_0000_01b3);
}
}
rows_census.push(format!("{r}:{am}:{h:016x}"));
}
eprintln!(
"[glm5-accrace] round={round} t={t} j={j} keep={} drafts={drafts:?} \
dev_vam={vam:?} host_vam={hvam:?} agree={} rows=[{}]",
j + 1,
hvam == vam,
rows_census.join(" ")
);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn glm5_dflash_round_drafts(
&self,
eh: &Engine,
sess: &mut Glm5SpecSession,
k: usize,
sp: Option<&SpecSampling>,
knobs: &mut Glm5SpecKnobs<'_>,
p_min: f32,
pmin0: bool,
pclk: Option<&mut ProfClock>,
) -> Res<(Vec<u32>, Glm5DraftQ)> {
let dr = self
.glm5_dflash
.as_ref()
.ok_or("glm5 dflash draft state without a loaded drafter")?;
let draft = &dr.draft;
let c = &draft.cfg;
let b = c.block_size;
let n_embd = self.cfg.n_embd as usize;
let n_vocab = self.output.out_features();
let Glm5SpecSession {
draft: state,
cache,
anchor,
sctr: _,
uctr,
rounds,
prof,
..
} = sess;
let Glm5DraftState::Dflash2 { kv, pending, taps } = state else {
return Err("glm5_dflash_round_drafts on a native-mtp session".into());
};
let anchor = *anchor;
let row_w = taps.len() * n_embd;
let mut pclk = pclk;
let n_new = pending.len() / row_w;
let st = if n_new > 0 {
let rows = std::mem::take(pending);
self.glm5_dflash_ingest_rows(eh, draft, kv, &rows, row_w, pclk.as_deref_mut())?
} else {
DraftIngestStats::default()
};
if let (Some(ck), Some(pf)) = (pclk, prof.as_mut()) {
let tail = ck.lap(eh, eh);
if *rounds == 0 && n_new > 0 && pf.draft_prime_arm.is_empty() {
st.write(pf, "eager-lazy");
pf.draft_prime_ms += tail;
}
}
let start = cache.pos;
debug_assert_eq!(
kv.len, start,
"drafter ctx rows must equal committed trunk rows at a round boundary"
);
let exact_scope = eh.exact_scope(true);
let mut block: Vec<u32> = vec![c.mask_token_id; b];
block[0] = anchor;
let noise = eh.htod(&self.embd.try_gather(n_embd, &block)?)?;
let pos_block: Vec<i32> = ((start as i32)..(start + b) as i32).collect();
let dh = draft.forward_round(eh, kv, &noise, &pos_block)?;
let nd = b - 1;
let mut rows_buf = eh.uninit(nd * n_embd)?;
{
let dv = eh.view(&dh, b * n_embd);
let tail = dv.slice(n_embd..b * n_embd);
eh.copy_view_into(&mut rows_buf, 0, &tail, nd * n_embd)?;
}
let trim = self.glm5_dflash_trim();
let (dl_head, dl_vocab) = match trim {
Some((head, d2t)) => (head, d2t.len()),
None => (&self.output, n_vocab),
};
let trim_d2t = trim.filter(|_| !knobs.skip_d2t_remap).map(|(_, d2t)| d2t);
let dl = eh.matmul(dl_head, &rows_buf, nd)?;
let (mut drafts, slot_q, qside) = match sp {
None => {
let (path, q) = draft
.dflash2_propose_greedy_q(eh, &dl, &rows_buf, nd, dl_vocab, anchor, trim_d2t)?;
(path, q, Glm5DraftQ::None)
}
Some(sp) => {
let (path, q_chosen, cand, q_rows) = draft.dflash2_propose_sampled(
eh, &dl, &rows_buf, nd, dl_vocab, anchor, sp.temp, sp.seed, uctr, trim_d2t,
)?;
let top_k = draft
.dflash2
.as_ref()
.ok_or("glm5 dflash drafter lost its DFlash2 head")?
.top_k;
(
path,
q_chosen.clone(),
Glm5DraftQ::Selector {
prop: DsparkDraftSample::Selector {
cand,
q_rows,
q_chosen,
top_k,
},
dl,
},
)
}
};
drop(exact_scope);
drafts.truncate(k);
if p_min > 0.0 {
let kc = glm5_conf_keep(&slot_q[..drafts.len()], p_min, pmin0);
drafts.truncate(kc);
}
if let Some(over) = knobs.draft_override.as_mut() {
for (ki, d) in drafts.iter_mut().enumerate() {
*d = over(*rounds, ki, *d);
}
}
Ok((drafts, qside))
}
#[allow(clippy::too_many_arguments)]
fn glm5_sampled_accept(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
sp: &SpecSampling,
vlogits: &CudaSlice<f32>,
drafts: &[u32],
draft_idx: &[u32],
draft_logits: &[CudaSlice<f32>],
draft_stats: &[(f32, f32, f32)],
d2t: Option<&[u32]>,
k: usize,
) -> Res<(usize, u32)> {
let n_vocab = self.output.out_features();
let d_vocab = d2t.map(|m| m.len()).unwrap_or(n_vocab);
let rows_i: Vec<i32> = (0..k as i32).collect();
let rows_d = e.htod_i32(&rows_i)?;
let (mut th_d, mut z_d, mut mx_d) = (e.zeros(k)?, e.zeros(k)?, e.zeros(k)?);
e.filter_stats(
vlogits, n_vocab, &rows_d, &mut th_d, &mut z_d, &mut mx_d, n_vocab, k, sp.temp,
sp.top_k, sp.top_p, sp.min_p,
)?;
let ids_d = e.htod_u32_v(drafts)?;
let mut pj_d = e.zeros(k)?;
e.softmax_gather_filtered(
vlogits, n_vocab, &ids_d, &rows_d, &th_d, &z_d, &mut pj_d, n_vocab, k, sp.temp,
)?;
let pj = e.dtoh(&pj_d)?;
let (thv, zv, mxv) = (e.dtoh(&th_d)?, e.dtoh(&z_d)?, e.dtoh(&mx_d)?);
let mut j = 0usize;
while j < k {
let (_qmx, qth, qz) = draft_stats[j];
let idsd = e.htod_u32_v(&[draft_idx[j]])?;
let rows0 = e.htod_i32(&[0])?;
let thd = e.htod(&[qth])?;
let zd = e.htod(&[qz])?;
let mut outd = e.zeros(1)?;
e.softmax_gather_filtered(
&draft_logits[j],
d_vocab,
&idsd,
&rows0,
&thd,
&zd,
&mut outd,
d_vocab,
1,
sp.temp,
)?;
let qj = e.dtoh(&outd)?[0];
let u = crate::spec::host_u01(sp.seed, sess.uctr);
sess.uctr = sess.uctr.wrapping_add(1);
if (u as f64) * (qj as f64) < pj[j] as f64 {
j += 1;
} else {
break;
}
}
if j == k {
return Ok((
j,
self.glm5_sampled_bonus(e, sess, sp, vlogits, k, n_vocab)?,
));
}
let mut col = e.zeros(n_vocab)?;
let bonus = {
let vv = e.view(vlogits, (k + 1) * n_vocab);
let row = vv.slice(j * n_vocab..(j + 1) * n_vocab);
e.copy_view_into(&mut col, 0, &row, n_vocab)?;
let p_stats = (mxv[j], thv[j], zv[j]);
let q_stats = draft_stats[j];
let sc = sess.sctr;
sess.sctr = sess.sctr.wrapping_add(1);
let mut sample_tok = e.alloc_u32_zeroed(1)?;
match d2t {
Some(map) => {
let map_d = e.htod_u32_v(map)?;
let mut q_full = e.zeros(n_vocab)?;
e.scatter_trim_logits(&draft_logits[j], &map_d, &mut q_full, d_vocab, n_vocab)?;
e.residual_sample_filtered(
&col,
Some(&q_full),
n_vocab,
sp.temp,
sp.seed,
sc,
p_stats,
q_stats,
&mut sample_tok,
)?;
}
None => {
e.residual_sample_filtered(
&col,
Some(&draft_logits[j]),
n_vocab,
sp.temp,
sp.seed,
sc,
p_stats,
q_stats,
&mut sample_tok,
)?;
}
}
e.dtoh_u32(&sample_tok)?[0]
};
let bonus = crate::spec::guard_vocab_token(
bonus,
n_vocab,
&format!("glm5 sampled verify bonus at round {} j={j}", sess.rounds),
)?;
Ok((j, bonus))
}
fn glm5_sampled_bonus(
&self,
e: &Engine,
sess: &mut Glm5SpecSession,
sp: &SpecSampling,
vlogits: &CudaSlice<f32>,
row: usize,
n_vocab: usize,
) -> Res<u32> {
let mut col = e.zeros(n_vocab)?;
let vv = e.view(vlogits, (row + 1) * n_vocab);
let src = vv.slice(row * n_vocab..(row + 1) * n_vocab);
e.copy_view_into(&mut col, 0, &src, n_vocab)?;
let rows0 = e.htod_i32(&[0])?;
let (mut bth, mut bz, mut bmx) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
&col, n_vocab, &rows0, &mut bth, &mut bz, &mut bmx, n_vocab, 1, sp.temp, sp.top_k,
sp.top_p, sp.min_p,
)?;
let (th, mx) = (e.dtoh(&bth)?[0], e.dtoh(&bmx)?[0]);
let mut pb = e.zeros(n_vocab)?;
e.gumbel_perturb_filtered(&col, &mut pb, n_vocab, sp.seed, sess.sctr, sp.temp, mx, th)?;
sess.sctr = sess.sctr.wrapping_add(1);
let td = e.argmax_token_device(&pb, n_vocab)?;
crate::spec::guard_vocab_token(
e.dtoh_u32_one(&td)?,
n_vocab,
&format!(
"glm5 sampled verify bonus at round {} (row {row})",
sess.rounds
),
)
}
}
fn glm5_sampled_draft(
e: &Engine,
dl: &CudaSlice<f32>,
d_vocab: usize,
sp: &SpecSampling,
sctr: &mut u32,
) -> Res<(u32, (f32, f32, f32))> {
let rows0 = e.htod_i32(&[0])?;
let (mut th_d, mut z_d, mut mx_d) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
dl, d_vocab, &rows0, &mut th_d, &mut z_d, &mut mx_d, d_vocab, 1, sp.temp, sp.top_k,
sp.top_p, sp.min_p,
)?;
let (th, z, mx) = (e.dtoh(&th_d)?[0], e.dtoh(&z_d)?[0], e.dtoh(&mx_d)?[0]);
let mut pb = e.zeros(d_vocab)?;
e.gumbel_perturb_filtered(dl, &mut pb, d_vocab, sp.seed, *sctr, sp.temp, mx, th)?;
*sctr = sctr.wrapping_add(1);
let td = e.argmax_token_device(&pb, d_vocab)?;
let idx =
crate::spec::guard_vocab_token(e.dtoh_u32_one(&td)?, d_vocab, "glm5 sampled draft draw")?;
Ok((idx, (mx, th, z)))
}
pub struct Glm5SpecSession {
cache: Cache,
pub committed: Vec<u32>,
anchor: u32,
anchor_emitted: bool,
pending: Vec<(u32, CudaSlice<f32>)>,
draft: Glm5DraftState,
sampling: Option<SpecSampling>,
pen: Option<Glm5Penalty>,
pen_hist: Vec<u32>,
sctr: u32,
uctr: u32,
pub rounds: usize,
pub rank_trimmed_rounds: usize,
done: bool,
max_ctx: usize,
mtp_il: Option<usize>,
prefix_capture: Option<crate::spec::SpecBoundaryCapture>,
prof: Option<SpecFirstTokenProf>,
prof_rounds: Option<SpecRoundsLog>,
}
impl Glm5SpecSession {
pub fn round_log_mut(&mut self) -> Option<&mut SpecRoundsLog> {
self.prof_rounds.as_mut()
}
pub fn round_log_cap() -> usize {
SPEC_PROF_ROUNDS
}
pub fn draft_kv_len(&self) -> Option<usize> {
match &self.draft {
Glm5DraftState::Dflash2 { kv, .. } => Some(kv.len),
_ => None,
}
}
#[allow(clippy::type_complexity)]
pub fn draft_kv_rows_host(
&self,
e: &Engine,
rows: usize,
row_floats: usize,
) -> Option<(Vec<Vec<f32>>, Vec<Vec<f32>>)> {
let Glm5DraftState::Dflash2 { kv, .. } = &self.draft else {
return None;
};
let n = rows * row_floats;
let take = |planes: &[CudaSlice<f32>]| -> Option<Vec<Vec<f32>>> {
planes
.iter()
.map(|p| e.dtoh_view(&p.slice(0..n)).ok())
.collect()
};
Some((take(&kv.k)?, take(&kv.v)?))
}
pub fn cache_max_ctx(&self) -> usize {
self.max_ctx
}
pub fn take_first_token_prof(&mut self) -> Option<SpecFirstTokenProf> {
self.prof.take()
}
pub fn take_prefix_capture(&mut self) -> Option<crate::spec::SpecBoundaryCapture> {
self.prefix_capture.take()
}
pub fn prefix_capture_ready(&self) -> bool {
self.prefix_capture
.as_ref()
.is_some_and(|c| match &self.draft {
Glm5DraftState::Dflash2 { kv, .. } => kv.len >= c.pos,
_ => false,
})
}
pub fn export_draft_tail(
&self,
e: &Engine,
upto: usize,
) -> Option<crate::dflash::DflashKvTail> {
match &self.draft {
Glm5DraftState::Dflash2 { kv, .. } => kv.export_tail(e, upto),
_ => None,
}
}
pub fn cache_ref(&self) -> &Cache {
&self.cache
}
pub fn pos(&self) -> usize {
self.cache.pos
}
pub fn finished(&self) -> bool {
self.done
}
pub fn demote_eligible(&self) -> bool {
self.sampling.is_none() && self.pen.is_none()
}
}
impl HybridModel {
pub fn glm5_spec_into_demoted(
&self,
e: &Engine,
mut sess: Glm5SpecSession,
) -> Res<(Cache, u32)> {
if !sess.demote_eligible() {
let why = if sess.sampling.is_some() {
"sampled sessions stay on spec until they end (session-owned Philox vs the \
worker sampler is an unmeasured distributional seam — the MTP sweep's \
exclusion, verbatim)"
} else {
"penalized sessions stay on spec until they end (this flush's plain argmax \
carries no penalty pass; demoting would silently drop the request's \
penalties for that token, lane/spec-exclusions-20260902)"
};
return Err(format!("glm5 demote: {why}").into());
}
if sess.cache.pos + 1 > sess.max_ctx {
return Err(format!(
"glm5 demote: no room to flush the live anchor ({} + 1 > ctx {})",
sess.cache.pos, sess.max_ctx
)
.into());
}
let logits = self.decode_step(e, sess.anchor, &mut sess.cache)?;
sess.committed.push(sess.anchor);
let next = argmax(&logits) as u32;
Ok((sess.cache, next))
}
}
#[derive(Default, Debug, Clone, Copy)]
struct DraftIngestStats {
h2d_ms: f64,
feat_ms: f64,
kv_ms: f64,
rows: usize,
chunks: usize,
}
impl DraftIngestStats {
fn write(&self, pf: &mut SpecFirstTokenProf, arm: &'static str) {
pf.draft_prime_ms = self.h2d_ms + self.feat_ms + self.kv_ms;
pf.draft_prime_h2d_ms = self.h2d_ms;
pf.draft_prime_feat_ms = self.feat_ms;
pf.draft_prime_kv_ms = self.kv_ms;
pf.draft_prime_rows = self.rows;
pf.draft_prime_chunks = self.chunks;
pf.draft_prime_arm = arm;
}
}
#[derive(Default)]
pub struct Glm5SpecKnobs<'a> {
pub draft_override: Option<&'a mut dyn FnMut(usize, usize, u32) -> u32>,
pub disable_rollback: bool,
pub skip_d2t_remap: bool,
pub pmin_override: Option<(f32, bool)>,
pub accept_probe: bool,
}
#[cfg(test)]
mod conf_keep_tests {
use super::glm5_conf_keep;
#[test]
fn conf_keep_matches_the_spec_rs_break_semantics() {
assert_eq!(glm5_conf_keep(&[0.1, 0.1], 0.0, true), 2);
assert_eq!(glm5_conf_keep(&[0.9, 0.8, 0.7], 0.5, false), 3);
assert_eq!(glm5_conf_keep(&[0.9, 0.2, 0.9], 0.5, false), 1);
assert_eq!(glm5_conf_keep(&[0.9, 0.2, 0.9], 0.5, true), 1);
assert_eq!(glm5_conf_keep(&[0.2, 0.9], 0.5, false), 2);
assert_eq!(glm5_conf_keep(&[0.2, 0.2], 0.5, false), 1);
assert_eq!(glm5_conf_keep(&[0.2, 0.9], 0.5, true), 0);
assert_eq!(glm5_conf_keep(&[0.5, 0.5], 0.5, true), 2);
assert_eq!(glm5_conf_keep(&[], 0.5, true), 0);
}
}