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::SpecPhaseNs;
use cudarc::driver::CudaSlice;
type Res<T> = Result<T, Box<dyn std::error::Error>>;
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_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_tp_on() -> bool {
std::env::var("MEMRA_GLM5_SPEC_TP").as_deref() == Ok("1")
}
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;
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.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 {
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 Some(slot) = sink.layer_ids.iter().position(|&l| l == il) else {
return Ok(());
};
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 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]);
}
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.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.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.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 sampling = sampling.filter(|sp| sp.temp > 0.0);
if let Some(sp) = sampling.as_ref()
&& sp.pen_on()
{
return Err(
"glm5 spec has no penalty arm yet: penalized sampled 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 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!(
"[glm5-spec] draft head TRIMMED to {} rows (FR-Spec d2t engaged)",
map.len(),
);
}
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 mut cache = crate::pp::new_cache_planned(e, &self.cfg, &self.plan, ctx_cap)?;
let plen = prompt.len();
let n_embd = self.cfg.n_embd as usize;
let tap_layers = match dflash_src {
Some(dr) => {
let taps = glm5_dflash_tap_layers(&dr.draft, self.layers.len())?;
cache.hc_taps = Some(HcTapSink::new(taps.clone(), n_embd, plen));
Some(taps)
}
None => None,
};
let (logits0, _seed, hiddens) = self.prime_cache(e, prompt, &mut cache, 0)?;
let eh = self.glm5_head_engine(e)?;
let prefix_capture = if glm5_spec_prefix_on() && dflash_src.is_some() {
self.glm5_prefix_boundary_capture(e, eh, &cache, &logits0, &hiddens, plen)
} else {
None
};
let mut sctr = 0u32;
let anchor = match sampling.as_ref() {
Some(sp) => {
crate::spec::sample_boundary_token(eh, &logits0, sp, &[], &mut sctr, "glm5-prime")?
}
None => argmax(&logits0) as u32,
};
let (draft, pending) = match (source_kind, dflash_src, tap_layers) {
(crate::spec::DraftSourceKind::Dflash2, Some(dr), Some(taps)) => {
let sink = cache
.hc_taps
.take()
.ok_or("glm5 dflash prime tap sink vanished")?;
let kv = DflashKv::new(eh, &dr.draft.cfg, ctx_cap)?;
(
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 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,
sctr,
uctr: 0,
rounds: 0,
done: false,
max_ctx: ctx_cap,
mtp_il,
prefix_capture,
})
}
fn glm5_prefix_boundary_capture(
&self,
e: &Engine,
eh: &Engine,
cache: &Cache,
logits0: &[f32],
hiddens: &CudaSlice<f32>,
plen: 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, plen, 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],
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 sampling = sampling.filter(|sp| sp.temp > 0.0);
if let Some(sp) = sampling.as_ref()
&& sp.pen_on()
{
return Err(
"glm5 spec has no penalty arm: penalized requests serve on the plain path".into(),
);
}
if fed.is_empty() || suffix.is_empty() {
return Err(
"restored glm5 spec session needs a non-empty restored prefix AND a \
non-empty suffix (empty-suffix full-cover hits keep the plain \
boundary-logits resume)"
.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!(
"[glm5-spec] draft head TRIMMED to {} rows (FR-Spec d2t engaged)",
map.len(),
);
}
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())?;
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)?;
let eh = self.glm5_head_engine(e)?;
let prefix_capture = if glm5_spec_prefix_on() {
self.glm5_prefix_boundary_capture(e, eh, &cache, &logits_s, &hiddens, cache.pos)
} else {
None
};
let mut sctr = 0u32;
let anchor = match sampling.as_ref() {
Some(sp) => crate::spec::sample_boundary_token(
eh,
&logits_s,
sp,
&[],
&mut sctr,
"glm5-restore",
)?,
None => argmax(&logits_s) as u32,
};
let sink = cache
.hc_taps
.take()
.ok_or("glm5 restored-session suffix tap sink vanished")?;
let mut committed = Vec::with_capacity(fed.len() + suffix.len());
committed.extend_from_slice(fed);
committed.extend_from_slice(suffix);
eprintln!(
"[glm5-spec] RESTORED session: {} prefix tokens + {} suffix from cache — no \
cold prime (drafter tail rows {})",
fed.len(),
suffix.len(),
dkv.len,
);
Ok(Glm5SpecSession {
cache,
committed,
anchor,
anchor_emitted: false,
pending: Vec::new(),
draft: Glm5DraftState::Dflash2 {
kv: dkv,
pending: sink.rows,
taps,
},
sampling,
sctr,
uctr: 0,
rounds: 0,
done: false,
max_ctx: ctx_cap,
mtp_il: None,
prefix_capture,
})
}
fn glm5_d2t(&self) -> Option<&[u32]> {
self.mtp
.as_ref()
.and_then(|head| head.d2t.as_deref())
.filter(|map| !map.is_empty())
}
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_gated(e, sess, target, k, eos, &mut Glm5SpecKnobs::default())
}
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)> {
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);
if !sess.anchor_emitted {
out.push(sess.anchor);
sess.anchor_emitted = true;
if eos.contains(&sess.anchor) {
sess.done = true;
}
}
while out.len() < target && !sess.done {
if sess.cache.pos + k + 2 > sess.max_ctx {
sess.done = true;
break;
}
let (round_tokens, n_drafted) =
self.glm5_spec_round(e, sess, k, d2t, sp_on.as_ref(), knobs, phase.as_mut())?;
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(ph) = phase.as_ref() {
ph.emit("glm5-phase", "glm5-phase-v", k);
}
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));
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;
}
};
}
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)?;
(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);
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);
let (j, bonus) = if let (true, Some(sp)) = (drafts.is_empty(), sp) {
(
0,
self.glm5_sampled_bonus(eh, sess, sp, &vlogits, 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(&vlogits, 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, &vlogits, &drafts, &vam, j)?;
}
(j, vam[j])
}
(
Some(sp),
Glm5DraftQ::Mtp {
draft_idx,
draft_logits,
draft_stats,
},
) => self.glm5_sampled_accept(
eh,
sess,
sp,
&vlogits,
&drafts,
draft_idx,
draft_logits,
draft_stats,
d2t,
drafts.len(),
)?,
(Some(sp), Glm5DraftQ::Selector { prop, dl }) => {
let (m, next) = crate::dflash::dspark_accept_sampled(
eh,
&vlogits,
&rows,
rows.len(),
n_vocab,
dl,
prop,
sp,
&[],
&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")
}
}
};
bump!(accept);
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);
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);
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,
) -> 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,
..
} = 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;
debug_assert_eq!(pending.len() % row_w, 0, "ragged pending feature rows");
let n_new = pending.len() / row_w;
let mut r0 = 0usize;
while r0 < n_new {
let t_c = (n_new - r0).min(256);
let chunk = eh.htod(&pending[r0 * row_w..(r0 + t_c) * row_w])?;
let feats = draft.ctx_features(eh, &chunk, t_c)?;
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)?;
r0 += t_c;
}
pending.clear();
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.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
.mtp
.as_ref()
.filter(|m| m.d2t_from_target_head)
.and_then(|m| m.shared_head_head.as_ref().zip(m.d2t.as_ref()))
.filter(|(_, d2t)| !d2t.is_empty());
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.as_slice());
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>,
sctr: u32,
uctr: u32,
pub rounds: usize,
done: bool,
max_ctx: usize,
mtp_il: Option<usize>,
prefix_capture: Option<crate::spec::SpecBoundaryCapture>,
}
impl Glm5SpecSession {
pub fn cache_max_ctx(&self) -> usize {
self.max_ctx
}
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()
}
}
impl HybridModel {
pub fn glm5_spec_into_demoted(
&self,
e: &Engine,
mut sess: Glm5SpecSession,
) -> Res<(Cache, u32)> {
if !sess.demote_eligible() {
return Err(
"glm5 demote: 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)"
.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)]
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);
}
}