use crate::Engine;
use crate::model::GpuTensor;
use cudarc::driver::CudaSlice;
pub struct DflashCfg {
pub hidden: usize, pub n_head: usize, pub n_kv: usize, pub head_dim: usize, pub n_ff: usize, pub n_layer: usize, pub eps: f32, pub rope_theta: f32, pub block_size: usize, pub mask_token_id: u32, pub target_layer_ids: Vec<usize>, pub sliding_window: usize, pub layer_sliding: Vec<bool>,
pub strategy_dspark: bool,
}
pub struct DflashLayer {
pub wq: GpuTensor, pub wk: GpuTensor, pub wv: GpuTensor, pub wo: GpuTensor, pub w_gate: GpuTensor, pub w_up: GpuTensor, pub w_down: GpuTensor, pub ln_in: CudaSlice<f32>, pub ln_post: CudaSlice<f32>, pub q_norm: CudaSlice<f32>, pub k_norm: CudaSlice<f32>, }
pub struct DflashDraft {
pub cfg: DflashCfg,
pub layers: Vec<DflashLayer>,
pub fc: GpuTensor, pub hidden_norm: CudaSlice<f32>, pub norm: CudaSlice<f32>, pub markov: Option<MarkovHead>,
pub confidence: Option<ConfidenceHead>,
pub rope_yarn: Option<(CudaSlice<f32>, f32)>,
}
pub struct ConfidenceHead {
pub w: Vec<f32>, pub b: f32,
pub in_dim: usize,
pub with_markov: bool,
}
impl ConfidenceHead {
pub fn raw_score(&self, hidden: &[f32], emb: Option<&[f32]>) -> f32 {
let mut acc = self.b;
for (w, x) in self.w.iter().zip(hidden) {
acc += w * x;
}
if self.with_markov {
let emb = emb.expect("with_markov confidence head scored without the markov embedding");
debug_assert_eq!(hidden.len() + emb.len(), self.in_dim);
for (w, x) in self.w[hidden.len()..].iter().zip(emb) {
acc += w * x;
}
} else {
debug_assert_eq!(hidden.len(), self.in_dim);
}
acc
}
}
pub struct MarkovHead {
pub w1_bf16: CudaSlice<u8>, pub w2: GpuTensor, pub rank: usize,
pub vocab: usize,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum DsparkHarvest {
Dflash,
Dspark,
}
impl DsparkHarvest {
pub fn resolve(cfg: &DflashCfg) -> Self {
Self::resolve_value(
std::env::var("MEMRA_DSPARK_HARVEST").ok().as_deref(),
cfg.strategy_dspark,
)
}
pub fn resolve_value(v: Option<&str>, strategy_dspark: bool) -> Self {
match v {
None | Some("") => {
if strategy_dspark {
DsparkHarvest::Dspark
} else {
DsparkHarvest::Dflash
}
}
set => Self::from_env_value(set),
}
}
pub fn from_env_value(v: Option<&str>) -> Self {
match v {
None | Some("") | Some("dflash") => DsparkHarvest::Dflash,
Some("dspark") => DsparkHarvest::Dspark,
Some(other) => panic!(
"MEMRA_DSPARK_HARVEST={other}: unknown harvest convention (dflash|dspark); \
refusing — a wrong convention verifies every draft slot against a position \
the drafter row was not trained for (DSPARK-POSTMORTEM-20260820.md)"
),
}
}
pub fn name(self) -> &'static str {
match self {
DsparkHarvest::Dflash => "dflash",
DsparkHarvest::Dspark => "dspark",
}
}
pub fn from_name(v: &str) -> Option<Self> {
match v {
"dflash" => Some(DsparkHarvest::Dflash),
"dspark" => Some(DsparkHarvest::Dspark),
_ => None,
}
}
pub fn first_row(self) -> usize {
match self {
DsparkHarvest::Dflash => 1,
DsparkHarvest::Dspark => 0,
}
}
pub fn n_drafts(self, b: usize) -> usize {
match self {
DsparkHarvest::Dflash => b - 1,
DsparkHarvest::Dspark => b,
}
}
pub fn trained_offset_of_row(self, row: usize) -> usize {
match self {
DsparkHarvest::Dflash => row,
DsparkHarvest::Dspark => row + 1,
}
}
}
pub fn dspark_strategy_census(txt: &str) -> bool {
let arch = txt
.find("\"architectures\"")
.and_then(|i| {
let rest = &txt[i..];
let a = rest.find('[')?;
let b = rest.find(']')?;
Some(rest[a..b].contains("DSpark"))
})
.unwrap_or(false);
let proj = txt
.find("\"projector_type\"")
.map(|i| {
let rest = &txt[i..];
let after = rest.find(':').map(|c| &rest[c + 1..]).unwrap_or("");
after.trim_start().starts_with("\"dspark\"")
})
.unwrap_or(false);
arch || proj
}
pub fn dspark_accept_prefix(cand: &[u32], vam: &[u32], vt: usize) -> usize {
let mut m = 0usize;
while m < vt - 1 && cand[m + 1] == vam[m] {
m += 1;
}
m
}
#[derive(Clone, Copy, PartialEq, Debug)]
pub enum DsparkVtPolicy {
Ladder,
Confidence { tau: f32 },
ConfidenceSlot { tau: f32 },
}
impl DsparkVtPolicy {
pub fn resolve(has_confidence_head: bool) -> Self {
Self::resolve_value(
std::env::var("MEMRA_DSPARK_VT").ok().as_deref(),
std::env::var("MEMRA_DSPARK_VT_TAU").ok().as_deref(),
std::env::var("MEMRA_DFLASH_ADAPT").ok().as_deref(),
has_confidence_head,
)
}
pub fn resolve_value(
vt: Option<&str>,
tau: Option<&str>,
adapt: Option<&str>,
has_confidence_head: bool,
) -> Self {
match vt {
None | Some("") => {
if adapt == Some("0") || !has_confidence_head {
DsparkVtPolicy::Ladder
} else {
Self::from_env_value(Some("confidence-slot"), tau, adapt)
}
}
set => Self::from_env_value(set, tau, adapt),
}
}
pub fn from_env_value(vt: Option<&str>, tau: Option<&str>, adapt: Option<&str>) -> Self {
match vt {
None | Some("") | Some("ladder") => DsparkVtPolicy::Ladder,
Some(mode @ ("confidence" | "confidence-slot")) => {
if adapt == Some("0") {
panic!(
"MEMRA_DSPARK_VT={mode} together with MEMRA_DFLASH_ADAPT=0 is \
contradictory (a pinned fixed window vs a per-round confidence \
window); unset one — refuse-on-ambiguity"
);
}
let tau = tau
.map(|t| {
t.parse::<f32>()
.unwrap_or_else(|_| panic!("MEMRA_DSPARK_VT_TAU={t}: not a float"))
})
.unwrap_or(0.5);
assert!(
tau > 0.0 && tau < 1.0,
"MEMRA_DSPARK_VT_TAU={tau}: confidence threshold must be in (0,1)"
);
if mode == "confidence" {
DsparkVtPolicy::Confidence { tau }
} else {
DsparkVtPolicy::ConfidenceSlot { tau }
}
}
Some(other) => panic!(
"MEMRA_DSPARK_VT={other}: unknown verify-window policy \
(ladder|confidence|confidence-slot); refusing — a wrong policy \
silently reverts the H4 arm (DSPARK-POSTMORTEM-20260820.md)"
),
}
}
pub fn is_confidence(&self) -> bool {
!matches!(self, DsparkVtPolicy::Ladder)
}
pub fn size_window(&self, raws: &[f32], vt_cap: usize) -> Option<usize> {
match *self {
DsparkVtPolicy::Ladder => None,
DsparkVtPolicy::Confidence { tau } => Some(dspark_confidence_vt(raws, tau, vt_cap)),
DsparkVtPolicy::ConfidenceSlot { tau } => {
Some(dspark_slot_confidence_vt(raws, tau, vt_cap))
}
}
}
}
pub fn dspark_confidence_vt(raws: &[f32], tau: f32, vt_cap: usize) -> usize {
let mut surv = 1.0f32;
let mut kept = 0usize;
for &r in raws {
surv *= 1.0 / (1.0 + (-r).exp());
if surv < tau {
break;
}
kept += 1;
}
(1 + kept).clamp(2, vt_cap.max(2))
}
pub fn dspark_slot_confidence_vt(raws: &[f32], tau: f32, vt_cap: usize) -> usize {
let mut kept = 0usize;
for &r in raws {
let p = 1.0 / (1.0 + (-r).exp());
if p < tau {
break;
}
kept += 1;
}
(1 + kept).clamp(2, vt_cap.max(2))
}
fn bf16_to_f32(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(2)
.map(|c| f32::from_bits((u16::from_le_bytes([c[0], c[1]]) as u32) << 16))
.collect()
}
fn encode_q8_0(vals: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(vals.len() / 32 * 34);
for blk in vals.chunks_exact(32) {
let amax = blk.iter().fold(0f32, |a, v| a.max(v.abs()));
let d = amax / 127.0;
let id = if d > 0.0 { 1.0 / d } else { 0.0 };
let dh = half_from_f32(d);
out.extend_from_slice(&dh.to_le_bytes());
for &v in blk {
out.push(((v * id).round().clamp(-127.0, 127.0)) as i8 as u8);
}
}
out
}
fn encode_q4_0(vals: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(vals.len() / 32 * 18);
for blk in vals.chunks_exact(32) {
let mut amax = 0f32;
let mut mx = 0f32;
for &v in blk {
if v.abs() > amax {
amax = v.abs();
mx = v;
}
}
let d = mx / -8.0;
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
out.extend_from_slice(&half_from_f32(d).to_le_bytes());
for j in 0..16 {
let x0 = (blk[j] * id + 8.5).clamp(0.0, 15.0) as u8;
let x1 = (blk[j + 16] * id + 8.5).clamp(0.0, 15.0) as u8;
out.push(x0 | (x1 << 4));
}
}
out
}
fn half_from_f32(v: f32) -> u16 {
let b = v.to_bits();
let sign = ((b >> 16) & 0x8000) as u16;
let exp = ((b >> 23) & 0xff) as i32 - 127 + 15;
let man = b & 0x7fffff;
if exp <= 0 {
return sign;
} if exp >= 31 {
return sign | 0x7c00;
} let mut h = sign | ((exp as u16) << 10) | ((man >> 13) as u16);
let rem = man & 0x1fff;
if rem > 0x1000 || (rem == 0x1000 && (h & 1) == 1) {
h += 1;
}
h
}
impl DflashDraft {
pub fn load(e: &Engine, dir: &std::path::Path) -> Result<Self, Box<dyn std::error::Error>> {
let txt = std::fs::read_to_string(dir.join("config.json"))?;
fn num(txt: &str, key: &str) -> Option<f64> {
let i = txt.find(&format!("\"{key}\""))?;
let rest = &txt[i..];
let colon = rest.find(':')?;
let val: String = rest[colon + 1..]
.trim_start()
.chars()
.take_while(|c| {
c.is_ascii_digit()
|| *c == '.'
|| *c == '-'
|| *c == 'e'
|| *c == 'E'
|| *c == '+'
})
.collect();
val.parse().ok()
}
fn num_list(txt: &str, key: &str) -> Vec<usize> {
let Some(i) = txt.find(&format!("\"{key}\"")) else {
return Vec::new();
};
let rest = &txt[i..];
let (Some(a), Some(b)) = (rest.find('['), rest.find(']')) else {
return Vec::new();
};
rest[a + 1..b]
.split(',')
.filter_map(|s| s.trim().parse().ok())
.collect()
}
let g = |k: &str| num(&txt, k).unwrap_or_else(|| panic!("config missing {k}")) as usize;
let layer_sliding: Vec<bool> = {
let i = txt.find("\"layer_types\"").expect("layer_types");
let rest = &txt[i..];
let (a, b) = (rest.find('[').unwrap(), rest.find(']').unwrap());
rest[a + 1..b]
.split(',')
.map(|s| s.contains("sliding_attention"))
.collect()
};
let sliding_window = if layer_sliding.iter().any(|&s| s) {
g("sliding_window")
} else {
num(&txt, "sliding_window")
.map(|v| v as usize)
.unwrap_or(usize::MAX)
};
let cfg = DflashCfg {
hidden: g("hidden_size"),
n_head: g("num_attention_heads"),
n_kv: g("num_key_value_heads"),
head_dim: g("head_dim"),
n_ff: g("intermediate_size"),
n_layer: g("num_hidden_layers"),
eps: num(&txt, "rms_norm_eps").expect("rms_norm_eps") as f32,
rope_theta: num(&txt, "rope_theta").expect("rope_theta") as f32,
block_size: g("block_size"),
mask_token_id: g("mask_token_id") as u32,
target_layer_ids: num_list(&txt, "target_layer_ids"),
sliding_window,
layer_sliding,
strategy_dspark: dspark_strategy_census(&txt),
};
let st = memra_gguf::safetensors::StModel::open(&dir.join("model.safetensors"))?;
let up = |name: &str| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (_info, bytes) = st
.raw(name)
.ok_or_else(|| format!("missing tensor {name}"))?;
Ok(e.htod(&bf16_to_f32(bytes))?)
};
let prec = std::env::var("MEMRA_DFLASH_PREC").unwrap_or_else(|_| "q8".into());
let upw = |name: &str| -> Result<GpuTensor, Box<dyn std::error::Error>> {
let (info, bytes) = st
.raw(name)
.ok_or_else(|| format!("missing tensor {name}"))?;
let shape = info.ne(); let in_f = shape[0] as usize;
let is_ffn = name.contains(".mlp.");
let bf16 = prec == "bf16"
|| (prec == "mixed" && !is_ffn)
|| (prec == "fc" && name == "fc.weight");
if bf16 {
return Ok(GpuTensor::FloatBf16 {
data: e.upload_u8(bytes)?,
ne: shape.to_vec(),
});
}
let f32s = bf16_to_f32(bytes);
if prec == "q4" {
let q = encode_q4_0(&f32s);
return Ok(GpuTensor::Quant {
bytes: e.upload_u8(&q)?,
qtype: crate::QT_Q4_0,
row_bytes: in_f / 32 * 18,
ne: shape.to_vec(),
scale: 1.0,
rp: false,
#[cfg(memra_cutlass)]
cutlass: None,
fp8: None,
blk: None,
rp4: None,
f16: None,
});
}
let q = encode_q8_0(&f32s);
Ok(GpuTensor::Quant {
bytes: e.upload_u8(&q)?,
qtype: crate::QT_Q8_0,
row_bytes: in_f / 32 * 34,
ne: shape.to_vec(),
scale: 1.0,
rp: false,
#[cfg(memra_cutlass)]
cutlass: None,
fp8: None,
blk: None,
rp4: None,
f16: None,
})
};
let mut layers = Vec::with_capacity(cfg.n_layer);
for i in 0..cfg.n_layer {
let p = |s: &str| format!("layers.{i}.{s}");
layers.push(DflashLayer {
wq: upw(&p("self_attn.q_proj.weight"))?,
wk: upw(&p("self_attn.k_proj.weight"))?,
wv: upw(&p("self_attn.v_proj.weight"))?,
wo: upw(&p("self_attn.o_proj.weight"))?,
w_gate: upw(&p("mlp.gate_proj.weight"))?,
w_up: upw(&p("mlp.up_proj.weight"))?,
w_down: upw(&p("mlp.down_proj.weight"))?,
ln_in: up(&p("input_layernorm.weight"))?,
ln_post: up(&p("post_attention_layernorm.weight"))?,
q_norm: up(&p("self_attn.q_norm.weight"))?,
k_norm: up(&p("self_attn.k_norm.weight"))?,
});
}
let markov = if let Some((info, bytes)) = st.raw("markov_head.markov_w1.weight") {
let sh = info.ne(); let (rank, vocab) = (sh[0] as usize, sh[1] as usize);
let (i2, b2) = st
.raw("markov_head.markov_w2.weight")
.ok_or("markov_w2 missing beside markov_w1")?;
let w2 = if prec == "bf16" {
GpuTensor::FloatBf16 {
data: e.upload_u8(b2)?,
ne: i2.ne().to_vec(),
}
} else {
let w2f = bf16_to_f32(b2);
let w2q = encode_q8_0(&w2f);
GpuTensor::Quant {
bytes: e.upload_u8(&w2q)?,
qtype: crate::QT_Q8_0,
row_bytes: rank / 32 * 34,
ne: vec![rank as u64, vocab as u64],
scale: 1.0,
rp: false,
#[cfg(memra_cutlass)]
cutlass: None,
fp8: None,
blk: None,
rp4: None,
f16: None,
}
};
Some(MarkovHead {
w1_bf16: e.upload_u8(bytes)?,
w2,
rank,
vocab,
})
} else {
None
};
let confidence = if let Some((info, bytes)) = st.raw("confidence_head.proj.weight") {
let sh = info.ne(); let in_dim = sh[0] as usize;
let (_bi, bb) = st
.raw("confidence_head.proj.bias")
.ok_or("confidence bias missing beside weight")?;
let with_markov = markov
.as_ref()
.map(|m| in_dim == cfg.hidden + m.rank)
.unwrap_or(false);
if !with_markov && in_dim != cfg.hidden {
panic!(
"confidence_head in_dim {in_dim} matches neither hidden {} nor hidden+rank",
cfg.hidden
);
}
Some(ConfidenceHead {
w: bf16_to_f32(bytes),
b: bf16_to_f32(bb)[0],
in_dim,
with_markov,
})
} else {
None
};
{
let mut consumed: std::collections::HashSet<String> = std::collections::HashSet::new();
for i in 0..cfg.n_layer {
for s in [
"self_attn.q_proj.weight",
"self_attn.k_proj.weight",
"self_attn.v_proj.weight",
"self_attn.o_proj.weight",
"self_attn.q_norm.weight",
"self_attn.k_norm.weight",
"input_layernorm.weight",
"post_attention_layernorm.weight",
"mlp.gate_proj.weight",
"mlp.up_proj.weight",
"mlp.down_proj.weight",
] {
consumed.insert(format!("layers.{i}.{s}"));
}
}
for s in [
"fc.weight",
"hidden_norm.weight",
"norm.weight",
"markov_head.markov_w1.weight",
"markov_head.markov_w2.weight",
"confidence_head.proj.weight",
"confidence_head.proj.bias",
] {
consumed.insert(s.into());
}
let leftovers: Vec<&String> = st.names().filter(|n| !consumed.contains(*n)).collect();
if !leftovers.is_empty() {
if markov.is_some() {
panic!("dspark census: unrecognized tensors {leftovers:?}");
}
eprintln!("[dflash census] unmapped tensors (ignored): {leftovers:?}");
}
}
let rope_yarn =
if txt.contains("\"rope_type\": \"yarn\"") || txt.contains("\"rope_type\":\"yarn\"") {
let factor = num(&txt, "factor").expect("yarn factor") as f64;
let orig = num(&txt, "original_max_position_embeddings").expect("yarn orig");
let beta_fast = num(&txt, "beta_fast").expect("beta_fast");
let beta_slow = num(&txt, "beta_slow").expect("beta_slow");
let base = cfg.rope_theta as f64;
let d = cfg.head_dim as f64;
let corr =
|r: f64| d * (orig / (r * 2.0 * std::f64::consts::PI)).ln() / (2.0 * base.ln());
let low = corr(beta_fast).floor().max(0.0);
let high = corr(beta_slow).ceil().min(d - 1.0);
let half = cfg.head_dim / 2;
let mut ff = Vec::with_capacity(half);
for j in 0..half {
let base_inv = base.powf(-2.0 * j as f64 / d);
let ramp = (((j as f64) - low) / (high - low)).clamp(0.0, 1.0);
let ex = 1.0 - ramp; let yarn_inv = (base_inv / factor) * (1.0 - ex) + base_inv * ex;
ff.push((base_inv / yarn_inv) as f32);
}
let mscale = (0.1 * factor.ln() + 1.0) as f32;
Some((e.htod(&ff)?, mscale))
} else {
None
};
eprintln!(
"[dspark] harvest={} (checkpoint census strategy_dspark={}, MEMRA_DSPARK_HARVEST {})",
DsparkHarvest::resolve_value(
std::env::var("MEMRA_DSPARK_HARVEST").ok().as_deref(),
cfg.strategy_dspark,
)
.name(),
cfg.strategy_dspark,
match std::env::var("MEMRA_DSPARK_HARVEST") {
Ok(v) if !v.is_empty() => "set",
_ => "unset",
},
);
eprintln!(
"[dspark] verify-window={:?} (accept-rate head {}, MEMRA_DSPARK_VT {})",
DsparkVtPolicy::resolve(confidence.is_some()),
if confidence.is_some() {
"present"
} else {
"ABSENT -> ladder"
},
match std::env::var("MEMRA_DSPARK_VT") {
Ok(v) if !v.is_empty() => "set",
_ => "unset",
},
);
Ok(Self {
fc: upw("fc.weight")?,
hidden_norm: up("hidden_norm.weight")?,
norm: up("norm.weight")?,
cfg,
layers,
markov,
confidence,
rope_yarn,
})
}
fn rope_rows(
&self,
e: &Engine,
x: &mut CudaSlice<f32>,
pos_d: &CudaSlice<i32>,
n_heads: usize,
n_tokens: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let c = &self.cfg;
match &self.rope_yarn {
Some((ff, mscale)) => {
e.rope_neox_ff(
x,
pos_d,
c.head_dim,
c.head_dim,
n_heads,
n_tokens,
c.rope_theta,
1.0,
ff,
)?;
e.scale_inplace(x, *mscale, n_tokens * n_heads * c.head_dim)?;
}
None => {
e.rope_neox(
x,
pos_d,
c.head_dim,
c.head_dim,
n_heads,
n_tokens,
c.rope_theta,
1.0,
)?;
}
}
Ok(())
}
fn mm(
&self,
e: &Engine,
w: &GpuTensor,
x: &CudaSlice<f32>,
t: usize,
_in_f: usize,
_out_f: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
Ok(e.matmul(w, x, t)?)
}
pub fn ctx_features(
&self,
e: &Engine,
taps: &CudaSlice<f32>,
t: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let c = &self.cfg;
let n_taps = c.target_layer_ids.len();
let fc_out = self.mm(e, &self.fc, taps, t, n_taps * c.hidden, c.hidden)?;
let mut out = e.uninit(t * c.hidden)?;
e.rms_norm(&fc_out, &self.hidden_norm, &mut out, c.hidden, t, c.eps)?;
Ok(out)
}
pub fn forward(
&self,
e: &Engine,
target_hidden: &CudaSlice<f32>,
noise_emb: &CudaSlice<f32>,
pos: &[i32],
ctx: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let ctx_f = self.ctx_features(e, target_hidden, ctx)?;
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&ctx_f)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-ctx_features.f32"), bytes)?;
}
self.forward_block(e, &ctx_f, noise_emb, pos, ctx)
}
pub fn forward_block(
&self,
e: &Engine,
ctx_f: &CudaSlice<f32>,
noise_emb: &CudaSlice<f32>,
pos: &[i32],
ctx: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let c = &self.cfg;
let (h, nh, nkv, hd) = (c.hidden, c.n_head, c.n_kv, c.head_dim);
let b = c.block_size;
assert_eq!(pos.len(), ctx + b, "pos covers ctx rows then block rows");
let pos_blk = e.htod_i32(&pos[ctx..])?;
let mut x = e.clone_dtod(noise_emb)?; for (li, l) in self.layers.iter().enumerate() {
let _ = li;
let mut xn = e.uninit(b * h)?;
e.rms_norm(&x, &l.ln_in, &mut xn, h, b, c.eps)?;
let q0 = self.mm(e, &l.wq, &xn, b, h, nh * hd)?;
let k0c = self.mm(e, &l.wk, ctx_f, ctx, h, nkv * hd)?;
let v0c = self.mm(e, &l.wv, ctx_f, ctx, h, nkv * hd)?;
let k0b = self.mm(e, &l.wk, &xn, b, h, nkv * hd)?;
let v0b = self.mm(e, &l.wv, &xn, b, h, nkv * hd)?;
let mut k0 = e.uninit((ctx + b) * nkv * hd)?;
e.copy_into(&mut k0, 0, &k0c, ctx * nkv * hd)?;
e.copy_into(&mut k0, ctx * nkv * hd, &k0b, b * nkv * hd)?;
let mut v = e.uninit((ctx + b) * nkv * hd)?;
e.copy_into(&mut v, 0, &v0c, ctx * nkv * hd)?;
e.copy_into(&mut v, ctx * nkv * hd, &v0b, b * nkv * hd)?;
if li == 0 {
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&q0)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_q0.f32"), bytes)?;
}
}
let mut q = e.uninit(b * nh * hd)?;
let mut k = e.uninit((ctx + b) * nkv * hd)?;
e.rms_norm(&q0, &l.q_norm, &mut q, hd, b * nh, c.eps)?;
if li == 0 {
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&q)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_qn.f32"), bytes)?;
}
}
e.rms_norm(&k0, &l.k_norm, &mut k, hd, (ctx + b) * nkv, c.eps)?;
let norope = std::env::var("MEMRA_DFLASH_NOROPE").is_ok();
if !norope {
self.rope_rows(e, &mut q, &pos_blk, nh, b)?;
}
if li == 0 {
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let dump = |name: &str,
t: &cudarc::driver::CudaSlice<f32>|
-> Result<(), Box<dyn std::error::Error>> {
let v = e.dtoh(t)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_{name}.f32"), bytes)?;
Ok(())
};
dump("xn", &xn)?;
dump("q_prerope", &q)?;
}
}
let pos_all = e.htod_i32(pos)?;
if !norope {
self.rope_rows(e, &mut k, &pos_all, nkv, ctx + b)?;
}
let mut attn = e.uninit(b * nh * hd)?;
let scale = 1.0f32 / (hd as f32).sqrt();
if std::env::var("MEMRA_DFLASH_FA").is_ok() {
e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, b, ctx + b, scale, false)?;
} else {
e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, b, ctx + b, scale, false)?;
}
let o = self.mm(e, &l.wo, &attn, b, nh * hd, h)?;
let mut x1 = e.uninit(b * h)?;
e.add(&o, &x, &mut x1, b * h)?;
if li == 0 {
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let dump = |name: &str,
t: &cudarc::driver::CudaSlice<f32>|
-> Result<(), Box<dyn std::error::Error>> {
let v = e.dtoh(t)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_{name}.f32"), bytes)?;
Ok(())
};
dump("q", &q)?;
dump("k", &k)?;
dump("attn", &attn)?;
dump("x1", &x1)?;
}
}
let mut x1n = e.uninit(b * h)?;
e.rms_norm(&x1, &l.ln_post, &mut x1n, h, b, c.eps)?;
let gate = self.mm(e, &l.w_gate, &x1n, b, h, c.n_ff)?;
let up_ = self.mm(e, &l.w_up, &x1n, b, h, c.n_ff)?;
let mut act = e.uninit(b * c.n_ff)?;
e.silu_mul(&gate, &up_, &mut act, b * c.n_ff)?;
let down = self.mm(e, &l.w_down, &act, b, c.n_ff, h)?;
let mut x2 = e.uninit(b * h)?;
e.add(&down, &x1, &mut x2, b * h)?;
x = x2;
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&x)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-layer{li}_out.f32"), bytes)?;
}
}
let mut out = e.uninit(b * h)?;
e.rms_norm(&x, &self.norm, &mut out, h, b, c.eps)?;
Ok(out)
}
}
pub struct DflashKv {
pub k: Vec<CudaSlice<f32>>, pub v: Vec<CudaSlice<f32>>,
pub len: usize,
pub cap: usize,
}
impl DflashKv {
pub fn new(
e: &Engine,
cfg: &DflashCfg,
cap: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
let rowsz = cfg.n_kv * cfg.head_dim;
let mut k = Vec::with_capacity(cfg.n_layer);
let mut v = Vec::with_capacity(cfg.n_layer);
for _ in 0..cfg.n_layer {
k.push(e.uninit((cap + cfg.block_size) * rowsz)?);
v.push(e.uninit((cap + cfg.block_size) * rowsz)?);
}
Ok(Self { k, v, len: 0, cap })
}
}
impl DflashDraft {
pub fn ingest_ctx(
&self,
e: &Engine,
kv: &mut DflashKv,
feats: &CudaSlice<f32>,
pos_new: &[i32],
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let c = &self.cfg;
let (h, nkv, hd) = (c.hidden, c.n_kv, c.head_dim);
assert!(kv.len + t <= kv.cap, "draft kv overflow");
let pos_d = e.htod_i32(pos_new)?;
for (li, l) in self.layers.iter().enumerate() {
let k0 = self.mm(e, &l.wk, feats, t, h, nkv * hd)?;
let v0 = self.mm(e, &l.wv, feats, t, h, nkv * hd)?;
let mut kn = e.uninit(t * nkv * hd)?;
e.rms_norm(&k0, &l.k_norm, &mut kn, hd, t * nkv, c.eps)?;
self.rope_rows(e, &mut kn, &pos_d, nkv, t)?;
e.copy_into(&mut kv.k[li], kv.len * nkv * hd, &kn, t * nkv * hd)?;
e.copy_into(&mut kv.v[li], kv.len * nkv * hd, &v0, t * nkv * hd)?;
}
kv.len += t;
Ok(())
}
pub fn forward_round(
&self,
e: &Engine,
kv: &mut DflashKv,
noise_emb: &CudaSlice<f32>,
pos_block: &[i32],
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let c = &self.cfg;
let (h, nh, nkv, hd) = (c.hidden, c.n_head, c.n_kv, c.head_dim);
let b = c.block_size;
assert_eq!(pos_block.len(), b);
let ctx = kv.len;
let pos_blk = e.htod_i32(pos_block)?;
let mut x = e.clone_dtod(noise_emb)?;
for (li, l) in self.layers.iter().enumerate() {
let mut xn = e.uninit(b * h)?;
e.rms_norm(&x, &l.ln_in, &mut xn, h, b, c.eps)?;
let q0 = self.mm(e, &l.wq, &xn, b, h, nh * hd)?;
let k0b = self.mm(e, &l.wk, &xn, b, h, nkv * hd)?;
let v0b = self.mm(e, &l.wv, &xn, b, h, nkv * hd)?;
let mut q = e.uninit(b * nh * hd)?;
let mut kb = e.uninit(b * nkv * hd)?;
e.rms_norm(&q0, &l.q_norm, &mut q, hd, b * nh, c.eps)?;
e.rms_norm(&k0b, &l.k_norm, &mut kb, hd, b * nkv, c.eps)?;
self.rope_rows(e, &mut q, &pos_blk, nh, b)?;
self.rope_rows(e, &mut kb, &pos_blk, nkv, b)?;
e.copy_into(&mut kv.k[li], ctx * nkv * hd, &kb, b * nkv * hd)?;
e.copy_into(&mut kv.v[li], ctx * nkv * hd, &v0b, b * nkv * hd)?;
let mut attn = e.uninit(b * nh * hd)?;
let scale = 1.0f32 / (hd as f32).sqrt();
if std::env::var("MEMRA_DFLASH_FA").is_ok() {
e.fa_prefill(
&q,
&kv.k[li],
&kv.v[li],
&mut attn,
hd,
nh,
nkv,
b,
ctx + b,
scale,
false,
)?;
} else {
e.sdpa_naive(
&q,
&kv.k[li],
&kv.v[li],
&mut attn,
hd,
nh,
nkv,
b,
ctx + b,
scale,
false,
)?;
}
let o = self.mm(e, &l.wo, &attn, b, nh * hd, h)?;
let mut x1 = e.uninit(b * h)?;
e.add(&o, &x, &mut x1, b * h)?;
let mut x1n = e.uninit(b * h)?;
e.rms_norm(&x1, &l.ln_post, &mut x1n, h, b, c.eps)?;
let gate = self.mm(e, &l.w_gate, &x1n, b, h, c.n_ff)?;
let up_ = self.mm(e, &l.w_up, &x1n, b, h, c.n_ff)?;
let mut act = e.uninit(b * c.n_ff)?;
e.silu_mul(&gate, &up_, &mut act, b * c.n_ff)?;
let down = self.mm(e, &l.w_down, &act, b, c.n_ff, h)?;
let mut x2 = e.uninit(b * h)?;
e.add(&down, &x1, &mut x2, b * h)?;
x = x2;
}
let mut out = e.uninit(b * h)?;
e.rms_norm(&x, &self.norm, &mut out, h, b, c.eps)?;
Ok(out)
}
}
impl crate::hybrid::HybridModel {
pub fn generate_spec_dflash(
&self,
e: &Engine,
draft: &DflashDraft,
prompt: &[u32],
max_new: usize,
eos: &[u32],
) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
use crate::cache::{Cache, DflashTapSink};
let n_embd = self.cfg.n_embd as usize;
let c = &draft.cfg;
assert_eq!(n_embd, c.hidden, "draft hidden must match target n_embd");
let b = c.block_size;
let n_taps = c.target_layer_ids.len();
let max_ctx = prompt.len() + max_new + b + 8;
assert!(
max_ctx <= c.sliding_window,
"first-light dflash round is windowless — ctx cap {} exceeds the draft window {}",
max_ctx,
c.sliding_window
);
let mut cache = Cache::new(e, &self.cfg, max_ctx)?;
let tp = prompt.len();
cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: e.uninit(tp * n_taps * n_embd)?,
hidden: n_embd,
t: tp,
base: 0,
});
let t_prime = std::time::Instant::now();
let (logits, _h_seed, _hiddens) = self.prime_cache(e, prompt, &mut cache, 0)?;
let mut last = crate::forward::argmax(&logits) as u32;
let mut dkv = DflashKv::new(e, &draft.cfg, max_ctx)?;
{
let taps = cache.dflash_taps.take().unwrap();
let n_taps_h = n_taps * n_embd;
let mut r0 = 0usize;
while r0 < tp {
let t_c = (tp - r0).min(256);
let tv = e.view(&taps.buf, tp * n_taps_h);
let win = tv.slice(r0 * n_taps_h..(r0 + t_c) * n_taps_h);
let mut chunk = e.uninit(t_c * n_taps_h)?;
e.copy_view_into(&mut chunk, 0, &win, t_c * n_taps_h)?;
let f = draft.ctx_features(e, &chunk, t_c)?;
let pos_c: Vec<i32> = ((r0 as i32)..(r0 + t_c) as i32).collect();
draft.ingest_ctx(e, &mut dkv, &f, &pos_c, t_c)?;
r0 += t_c;
}
}
let mut ctx_len = tp;
e.stream().synchronize()?;
crate::PRIME_NANOS.store(
t_prime.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
let emb_scale = if std::env::var("MEMRA_DFLASH_EMB_SCALE").as_deref() == Ok("1") {
(n_embd as f32).sqrt()
} else {
1.0
};
let mut out = Vec::with_capacity(max_new);
let n_vocab = self.output.out_features();
let vt_cap: usize = std::env::var("MEMRA_DFLASH_VERIFY_T")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(8)
.clamp(2, b);
let adapt = std::env::var("MEMRA_DFLASH_ADAPT").as_deref() != Ok("0");
let mut vt = vt_cap;
let mut attempted = 0usize;
let mut accepted = 0usize;
e.set_verify_exact(true);
'outer: while out.len() < max_new {
let start = cache.pos; let mut block: Vec<u32> = vec![c.mask_token_id; b];
block[0] = last;
let mut noise = e.htod(&self.embd.gather(n_embd, &block))?;
if emb_scale != 1.0 {
e.scale_inplace(&mut noise, emb_scale, b * n_embd)?;
}
if std::env::var("MEMRA_DFLASH_DEBUG").as_deref() == Ok("1") && start == cache.pos {
let nv = e.dtoh(&noise)?;
let r0: f32 = nv[..n_embd].iter().map(|x| x * x).sum::<f32>().sqrt();
let r1: f32 = nv[n_embd..2 * n_embd]
.iter()
.map(|x| x * x)
.sum::<f32>()
.sqrt();
eprintln!(
"[dflash noise] |row0(last)|={r0:.3} |row1(MASK id {})|={r1:.3}",
c.mask_token_id
);
}
let pos_block: Vec<i32> = ((start as i32)..(start + b) as i32).collect();
let dh = draft.forward_round(e, &mut dkv, &noise, &pos_block)?;
let mut rows = e.uninit((b - 1) * n_embd)?;
{
let dv = e.view(&dh, b * n_embd);
let tail = dv.slice(n_embd..b * n_embd);
e.copy_view_into(&mut rows, 0, &tail, (b - 1) * n_embd)?;
}
let mut dl = e.matmul(&self.output, &rows, b - 1)?;
let markov_on = std::env::var("MEMRA_DFLASH_MARKOV").as_deref() != Ok("0");
let mut chain_d = e.stream().alloc_zeros::<u32>(b)?;
if let (Some(mk), true) = (&draft.markov, markov_on) {
e.set_u32_one(&mut chain_d, last)?;
for k in 0..(b - 1) {
let mut f = e.uninit(mk.rank)?;
e.gather_row_bf16(&mk.w1_bf16, &chain_d, k, &mut f, mk.rank)?;
let bias = e.matmul(&mk.w2, &f, 1)?;
e.add_row_inplace(&mut dl, &bias, n_vocab, k * n_vocab)?;
e.argmax_token_device_col(&dl, k, n_vocab, &mut chain_d, k + 1)?;
}
} else {
for i in 0..(b - 1) {
e.argmax_token_device_col(&dl, i, n_vocab, &mut chain_d, i + 1)?;
}
}
let chain = e.dtoh_u32(&chain_d)?;
let dtoks = &chain[1..];
for (i, &dt) in dtoks.iter().enumerate() {
block[i + 1] = dt;
}
let dbg = std::env::var("MEMRA_DFLASH_DEBUG").as_deref() == Ok("1");
let vblock = &block[..vt];
cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: e.uninit(vt * n_taps * n_embd)?,
hidden: n_embd,
t: vt,
base: 0,
});
let (vam, _vh) = self.gemma4_decode_step_t_am(e, vblock, start, &mut cache)?;
let taps = cache.dflash_taps.take().unwrap();
if dbg {
eprintln!(
"[dflash r] start={start} last={last}\n draft={:?}\n vam ={:?}",
&block[1..],
&vam
);
}
let mut m = 0usize;
while m < vt - 1 && block[m + 1] as usize == vam[m] as usize {
m += 1;
}
attempted += vt - 1;
accepted += m;
out.push(last);
if eos.contains(&last) {
break 'outer;
}
for &dt in &block[1..=m] {
out.push(dt);
if eos.contains(&dt) {
break 'outer;
}
if out.len() >= max_new {
break 'outer;
}
}
let next = vam[m] as u32;
let keep = m + 1;
for kvl in cache.kv.iter_mut().flatten() {
kvl.len -= vt - keep;
e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
}
cache.pos -= vt - keep;
{
let tv = e.view(&taps.buf, vt * n_taps * n_embd);
let keep_view = tv.slice(0..keep * n_taps * n_embd);
let mut kept = e.uninit(keep * n_taps * n_embd)?;
e.copy_view_into(&mut kept, 0, &keep_view, keep * n_taps * n_embd)?;
let f = draft.ctx_features(e, &kept, keep)?;
let pos_k: Vec<i32> = ((ctx_len as i32)..(ctx_len + keep) as i32).collect();
draft.ingest_ctx(e, &mut dkv, &f, &pos_k, keep)?;
ctx_len += keep;
}
last = next;
if adapt {
vt = (m + 2).clamp(3, vt_cap);
}
}
e.set_verify_exact(false);
if std::env::var("MEMRA_SPEC_STATS").as_deref() == Ok("1") {
eprintln!(
"[dflash] acceptance {accepted}/{attempted} = {:.3}",
accepted as f64 / attempted.max(1) as f64
);
}
Ok(out)
}
}
pub(crate) struct DsparkSnapBatch {
pub(crate) snap: crate::cache::CacheSnapshot,
lin: Vec<usize>,
conv_table: CudaSlice<u64>,
ssm_table: CudaSlice<u64>,
host_ssm: Vec<u64>,
conv_words: usize,
ssm_words: usize,
}
impl DsparkSnapBatch {
pub(crate) fn new(
e: &Engine,
cache: &crate::cache::Cache,
) -> Result<Option<Self>, Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
let snap = cache.snapshot(e)?;
let lin: Vec<usize> = (0..cache.recur.len())
.filter(|&il| cache.recur[il].is_some())
.collect();
if lin.is_empty() {
return Ok(None);
}
let first = cache.recur[lin[0]].as_ref().unwrap();
let (conv_words, ssm_words) = (first.conv_state.len(), first.ssm_state.len());
for &il in &lin {
let rl = cache.recur[il].as_ref().unwrap();
if rl.conv_state.len() != conv_words || rl.ssm_state.len() != ssm_words {
return Ok(None);
}
}
let n = lin.len();
let mut host_conv = vec![0u64; 2 * n];
let mut host_ssm = vec![0u64; 2 * n];
{
let s = &e.gpu.stream();
for (k, &il) in lin.iter().enumerate() {
let rl = cache.recur[il].as_ref().unwrap();
let (pc, _g0) = rl.conv_state.device_ptr(s);
let (ps, _g1) = rl.ssm_state.device_ptr(s);
let (dc, _g2) = snap.conv[il].as_ref().unwrap().device_ptr(s);
let (ds, _g3) = snap.ssm[il].as_ref().unwrap().device_ptr(s);
host_conv[k] = pc as u64;
host_conv[n + k] = dc as u64;
host_ssm[k] = ps as u64;
host_ssm[n + k] = ds as u64;
}
}
let conv_table = e.htod_u64(&host_conv)?;
let ssm_table = e.htod_u64(&host_ssm)?;
Ok(Some(Self {
snap,
lin,
conv_table,
ssm_table,
host_ssm,
conv_words,
ssm_words,
}))
}
pub(crate) fn refresh(
&mut self,
e: &Engine,
cache: &crate::cache::Cache,
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
for il in 0..cache.kv.len() {
self.snap.kv_len[il] = cache.kv[il].as_ref().map(|kvl| kvl.len);
}
self.snap.pos = cache.pos;
let n = self.lin.len();
{
let s = &e.gpu.stream();
for (k, &il) in self.lin.iter().enumerate() {
let rl = cache.recur[il].as_ref().unwrap();
let (ps, _g) = rl.ssm_state.device_ptr(s);
self.host_ssm[k] = ps as u64;
}
}
e.htod_u64_into(&self.host_ssm, &mut self.ssm_table)?;
e.copy_batch_uniform_f32(&self.conv_table, n, self.conv_words)?;
e.copy_batch_uniform_f32(&self.ssm_table, n, self.ssm_words)?;
Ok(())
}
}
impl crate::hybrid::HybridModel {
pub fn generate_spec_dspark(
&self,
e: &Engine,
draft: &DflashDraft,
prompt: &[u32],
max_new: usize,
eos: &[u32],
) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
use crate::cache::{Cache, DflashTapSink};
assert!(
self.cfg.gemma4.is_none(),
"gemma4 targets use generate_spec_dflash; this is the qwen-hybrid arm"
);
let n_embd = self.cfg.n_embd as usize;
let c = &draft.cfg;
assert_eq!(n_embd, c.hidden, "draft hidden must match target n_embd");
let b = c.block_size;
let n_taps = c.target_layer_ids.len();
let max_ctx = prompt.len() + max_new + b + 8;
assert!(
max_ctx <= c.sliding_window,
"dspark round is windowless — ctx cap {} exceeds the draft window {}",
max_ctx,
c.sliding_window
);
let mut cache = Cache::new(e, &self.cfg, max_ctx)?;
let tp = prompt.len();
cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: e.uninit(tp * n_taps * n_embd)?,
hidden: n_embd,
t: tp,
base: 0,
});
let t_prime = std::time::Instant::now();
let (logits, _h_seed, _hiddens) = self.prime_cache(e, prompt, &mut cache, 0)?;
let mut last = crate::forward::argmax(&logits) as u32;
let mut dkv = DflashKv::new(e, &draft.cfg, max_ctx)?;
{
let taps = cache.dflash_taps.take().unwrap();
let n_taps_h = n_taps * n_embd;
let mut r0 = 0usize;
while r0 < tp {
let t_c = (tp - r0).min(256);
let tv = e.view(&taps.buf, tp * n_taps_h);
let win = tv.slice(r0 * n_taps_h..(r0 + t_c) * n_taps_h);
let mut chunk = e.uninit(t_c * n_taps_h)?;
e.copy_view_into(&mut chunk, 0, &win, t_c * n_taps_h)?;
let f = draft.ctx_features(e, &chunk, t_c)?;
let pos_c: Vec<i32> = ((r0 as i32)..(r0 + t_c) as i32).collect();
draft.ingest_ctx(e, &mut dkv, &f, &pos_c, t_c)?;
r0 += t_c;
}
}
let mut ctx_len = tp;
e.stream().synchronize()?;
crate::PRIME_NANOS.store(
t_prime.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
let mut out = Vec::with_capacity(max_new);
let n_vocab = self.output.out_features();
let harvest = DsparkHarvest::resolve(&draft.cfg);
let nd = harvest.n_drafts(b);
let r0 = harvest.first_row();
let vt_cap: usize = std::env::var("MEMRA_DFLASH_VERIFY_T")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(nd + 1)
.clamp(2, nd + 1);
let adapt = std::env::var("MEMRA_DFLASH_ADAPT").as_deref() != Ok("0");
let vt_policy = DsparkVtPolicy::resolve(draft.confidence.is_some());
if vt_policy.is_confidence() {
assert!(
draft.confidence.is_some(),
"MEMRA_DSPARK_VT={vt_policy:?} needs a checkpoint with an accept-rate \
head (confidence_head.* absent in this export)"
);
}
let mut vt = vt_cap;
let mut attempted = 0usize;
let mut accepted = 0usize;
let mut snapb: Option<DsparkSnapBatch> = None;
let mut snapb_off = !crate::spec::state_copy_batch_on();
let defer_rb = crate::spec::dspark_defer_readback_on() && !vt_policy.is_confidence();
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let embd_gpu = if !defer_rb || crate::spec::spec_host_embd() {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
let mut vg_guard = self.dspark_vgraphs.lock().unwrap();
if vg_guard.is_none() && embd_gpu.is_some() && crate::spec::dspark_verify_graph_on() {
*vg_guard = crate::spec::DsparkVerifyGraphs::new(e, &cache, vt_cap, n_embd)?;
}
let vgraphs: &mut Option<crate::spec::DsparkVerifyGraphs> = &mut vg_guard;
let (mut ns_draft, mut ns_snap, mut ns_verify, mut ns_roll, mut ns_ingest) =
(0u64, 0u64, 0u64, 0u64, 0u64);
let mut rounds = 0usize;
let stats = std::env::var("MEMRA_SPEC_STATS").as_deref() == Ok("1");
let clock = |on: bool, e: &Engine| -> std::time::Instant {
if on {
let _ = e.stream().synchronize();
}
std::time::Instant::now()
};
'outer: while out.len() < max_new {
rounds += 1;
let start = cache.pos; let t0 = clock(stats, e);
e.set_verify_exact(true);
let mut block: Vec<u32> = vec![c.mask_token_id; b];
block[0] = last;
let noise = e.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(e, &mut dkv, &noise, &pos_block)?;
let mut rows = e.uninit(nd * n_embd)?;
{
let dv = e.view(&dh, b * n_embd);
let src = dv.slice(r0 * n_embd..(r0 + nd) * n_embd);
e.copy_view_into(&mut rows, 0, &src, nd * n_embd)?;
}
let mut dl = e.matmul(&self.output, &rows, nd)?;
let markov_on = std::env::var("MEMRA_DFLASH_MARKOV").as_deref() != Ok("0");
let mut chain_d = e.stream().alloc_zeros::<u32>(nd + 1)?;
let want_conf_emb = vt_policy.is_confidence()
&& draft.confidence.as_ref().is_some_and(|ch| ch.with_markov);
let mut conf_emb: Option<CudaSlice<f32>> = match (&draft.markov, want_conf_emb) {
(Some(mk), true) => Some(e.uninit(nd * mk.rank)?),
(None, true) => unreachable!(
"with_markov confidence head without a markov table — the loader forbids it"
),
_ => None,
};
if let (Some(mk), true) = (&draft.markov, markov_on) {
e.set_u32_one(&mut chain_d, last)?;
for k in 0..nd {
let mut f = e.uninit(mk.rank)?;
e.gather_row_bf16(&mk.w1_bf16, &chain_d, k, &mut f, mk.rank)?;
if let Some(ce) = conf_emb.as_mut() {
let fv = e.view(&f, mk.rank);
e.copy_view_into(ce, k * mk.rank, &fv, mk.rank)?;
}
let bias = e.matmul(&mk.w2, &f, 1)?;
e.add_row_inplace(&mut dl, &bias, n_vocab, k * n_vocab)?;
e.argmax_token_device_col(&dl, k, n_vocab, &mut chain_d, k + 1)?;
}
} else {
if want_conf_emb {
e.set_u32_one(&mut chain_d, last)?;
}
for i in 0..nd {
if let (Some(ce), Some(mk)) = (conf_emb.as_mut(), &draft.markov) {
let mut f = e.uninit(mk.rank)?;
e.gather_row_bf16(&mk.w1_bf16, &chain_d, i, &mut f, mk.rank)?;
let fv = e.view(&f, mk.rank);
e.copy_view_into(ce, i * mk.rank, &fv, mk.rank)?;
}
e.argmax_token_device_col(&dl, i, n_vocab, &mut chain_d, i + 1)?;
}
}
e.set_verify_exact(false);
let ckpt_on = std::env::var("MEMRA_DSPARK_CKPT").as_deref() != Ok("0");
let ckpt_gate = std::env::var("MEMRA_DSPARK_CKPT_GATE").as_deref() == Ok("1");
let deferred = embd_gpu.is_some() && (ckpt_on || ckpt_gate);
if vt_policy.is_confidence() {
let ch = draft.confidence.as_ref().expect("asserted at loop entry");
let (rows_h, emb_h) = match conf_emb.as_ref() {
Some(ce) => {
let (a, b2) = e.dtoh_pair(&rows, ce)?;
(a, Some(b2))
}
None => (e.dtoh(&rows)?, None),
};
let rank = draft.markov.as_ref().map(|m| m.rank).unwrap_or(0);
let mut raws = Vec::with_capacity(nd);
for k in 0..nd {
let hrow = &rows_h[k * n_embd..(k + 1) * n_embd];
let emb = emb_h.as_ref().map(|eh| &eh[k * rank..(k + 1) * rank]);
raws.push(ch.raw_score(hrow, emb));
}
vt = vt_policy
.size_window(&raws, vt_cap)
.expect("confidence policies always size the window");
}
let mut cand: Vec<u32> = Vec::with_capacity(nd + 1);
if !deferred {
let chain = e.dtoh_u32(&chain_d)?;
cand.push(last);
cand.extend_from_slice(&chain[1..]);
}
ns_draft += clock(stats, e).duration_since(t0).as_nanos() as u64;
let t1 = std::time::Instant::now();
let mut snap_legacy: Option<crate::cache::CacheSnapshot> = None;
if !snapb_off && snapb.is_none() {
snapb = DsparkSnapBatch::new(e, &cache)?;
snapb_off = snapb.is_none();
} else if let Some(sb) = snapb.as_mut() {
sb.refresh(e, &cache)?;
}
let snap: &crate::cache::CacheSnapshot = match snapb.as_ref() {
Some(sb) => &sb.snap,
None => {
snap_legacy = Some(cache.snapshot(e)?);
snap_legacy.as_ref().unwrap()
}
};
let _ = &snap_legacy;
ns_snap += clock(stats, e).duration_since(t1).as_nanos() as u64;
let t2 = std::time::Instant::now();
let tap_buf = match vgraphs.as_mut().and_then(|g| g.tap_bufs.remove(&vt)) {
Some(buf) => buf,
None => e.uninit(vt * n_taps * n_embd)?,
};
cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: tap_buf,
hidden: n_embd,
t: vt,
base: 0,
});
let verify_res = (|cache: &mut crate::cache::Cache,
cand: &mut Vec<u32>,
vgraphs: &mut Option<crate::spec::DsparkVerifyGraphs>|
-> Result<
(Vec<u32>, Option<crate::spec::DsparkVerifyCkpt>),
Box<dyn std::error::Error>,
> {
if deferred {
let g = embd_gpu.expect("deferred implies resident embed");
let (am_d, vck) = self.dspark_verify_t_am_ckpt_dev(
e,
&chain_d,
vt,
start,
cache,
(g, embd_qt, embd_rb),
vgraphs.as_mut(),
)?;
let ch = e.stream().clone_dtoh(&chain_d)?;
let am = e.stream().clone_dtoh(&am_d)?;
e.stream().synchronize()?;
cand.push(last);
cand.extend_from_slice(&ch[1..]);
Ok((am, Some(vck)))
} else if ckpt_on || ckpt_gate {
let (vam, vck) = self.dspark_verify_t_am_ckpt(e, &cand[..vt], start, cache)?;
Ok((vam, Some(vck)))
} else {
Ok((self.dspark_verify_t_am(e, &cand[..vt], start, cache)?, None))
}
})(&mut cache, &mut cand, vgraphs);
let (vam, vck) = match verify_res {
Ok(v) => v,
Err(err) => {
if let (Some(g), Some(taps)) = (vgraphs.as_mut(), cache.dflash_taps.take()) {
g.tap_bufs.insert(vt, taps.buf);
}
return Err(err);
}
};
let taps = cache.dflash_taps.take().unwrap();
let tap_local: Option<CudaSlice<f32>> = match vgraphs.as_mut() {
Some(g) => {
g.tap_bufs.insert(vt, taps.buf);
None
}
None => Some(taps.buf),
};
let tap_ref: &CudaSlice<f32> = match &tap_local {
Some(b) => b,
None => &vgraphs.as_ref().expect("ctx present above").tap_bufs[&vt],
};
ns_verify += clock(stats, e).duration_since(t2).as_nanos() as u64;
let m = dspark_accept_prefix(&cand, &vam, vt);
attempted += vt - 1;
accepted += m;
out.push(last);
if eos.contains(&last) {
break 'outer;
}
for &dt in &cand[1..=m] {
if out.len() >= max_new {
break 'outer;
}
out.push(dt);
if eos.contains(&dt) {
break 'outer;
}
}
let next = vam[m];
let keep = m + 1;
let t3 = std::time::Instant::now();
let slab_commit = vgraphs.as_ref().map(|g| g.round_slab).unwrap_or(false);
if keep < vt {
if ckpt_gate {
if slab_commit {
self.dspark_commit_prefix_slab(
e,
&mut cache,
snap,
vgraphs.as_ref().expect("slab_commit implies ctx"),
keep,
)?;
} else {
let vck = vck.as_ref().expect("gate arm always fills the ckpt");
self.dspark_commit_prefix(e, &mut cache, snap, vck, keep)?;
}
let capture = |cache: &Cache| -> Result<
(usize, Vec<Option<usize>>, Vec<(Vec<f32>, Vec<f32>)>),
Box<dyn std::error::Error>,
> {
let mut lens = Vec::new();
let mut states = Vec::new();
for il in 0..cache.kv.len() {
lens.push(cache.kv[il].as_ref().map(|k| k.len));
if let Some(rl) = &cache.recur[il] {
states.push((e.dtoh(&rl.conv_state)?, e.dtoh(&rl.ssm_state)?));
}
}
Ok((cache.pos, lens, states))
};
let (p1, l1, st1) = capture(&cache)?;
crate::pp::restore_cache_checkpoint(e, &self.cfg, None, &mut cache, snap)?;
let ram = self.dspark_verify_t_am(e, &cand[..keep], start, &mut cache)?;
assert_eq!(
&ram[..],
&vam[..keep],
"prefix replay must reproduce the verify argmaxes"
);
let (p2, l2, st2) = capture(&cache)?;
assert_eq!(p1, p2, "ckpt-gate: pos mismatch");
assert_eq!(l1, l2, "ckpt-gate: kv_len mismatch");
for (il, ((c1, s1v), (c2, s2v))) in st1.iter().zip(&st2).enumerate() {
let bits = |a: &[f32], b: &[f32]| {
a.iter().zip(b).all(|(x, y)| x.to_bits() == y.to_bits())
};
assert!(
bits(c1, c2),
"ckpt-gate: linear layer {il} conv state differs"
);
assert!(
bits(s1v, s2v),
"ckpt-gate: linear layer {il} ssm state differs"
);
}
} else if slab_commit {
self.dspark_commit_prefix_slab(
e,
&mut cache,
snap,
vgraphs.as_ref().expect("slab_commit implies ctx"),
keep,
)?;
} else if let Some(vck) = vck.as_ref() {
self.dspark_commit_prefix(e, &mut cache, snap, vck, keep)?;
} else {
crate::pp::restore_cache_checkpoint(e, &self.cfg, None, &mut cache, snap)?;
debug_assert_eq!(cache.pos, start, "rollback landed off the round start");
let ram = self.dspark_verify_t_am(e, &cand[..keep], start, &mut cache)?;
debug_assert_eq!(
&ram[..],
&vam[..keep],
"prefix replay must reproduce the verify argmaxes"
);
}
}
ns_roll += clock(stats, e).duration_since(t3).as_nanos() as u64;
let t4 = std::time::Instant::now();
{
let tv = e.view(tap_ref, vt * n_taps * n_embd);
let keep_view = tv.slice(0..keep * n_taps * n_embd);
let mut kept = e.uninit(keep * n_taps * n_embd)?;
e.copy_view_into(&mut kept, 0, &keep_view, keep * n_taps * n_embd)?;
let f = draft.ctx_features(e, &kept, keep)?;
let pos_k: Vec<i32> = ((ctx_len as i32)..(ctx_len + keep) as i32).collect();
draft.ingest_ctx(e, &mut dkv, &f, &pos_k, keep)?;
ctx_len += keep;
}
ns_ingest += clock(stats, e).duration_since(t4).as_nanos() as u64;
last = next;
if !vt_policy.is_confidence() && adapt {
vt = (m + 2).clamp(3, vt_cap);
}
}
if stats {
let ms = |n: u64| n as f64 / 1e6;
eprintln!(
"[dspark-q38] acceptance {accepted}/{attempted} = {:.3} rounds={rounds} \
draft={:.1}ms snap={:.1}ms verify={:.1}ms rollback+replay={:.1}ms ingest={:.1}ms",
accepted as f64 / attempted.max(1) as f64,
ms(ns_draft),
ms(ns_snap),
ms(ns_verify),
ms(ns_roll),
ms(ns_ingest)
);
}
Ok(out)
}
}
pub struct DsparkSpecSession {
pub cache: crate::cache::Cache,
dkv: DflashKv,
last: u32,
ctx_len: usize,
vt: usize,
pub rounds: usize,
max_ctx: usize,
done: bool,
snapb: Option<DsparkSnapBatch>,
snapb_off: bool,
}
impl DsparkSpecSession {
pub fn cache_max_ctx(&self) -> usize {
self.max_ctx
}
pub fn finished(&self) -> bool {
self.done
}
pub fn pos(&self) -> usize {
self.cache.pos
}
}
impl crate::hybrid::HybridModel {
pub fn dspark_spec_session_new(
&self,
e: &Engine,
draft: &DflashDraft,
prompt: &[u32],
ctx_cap: usize,
) -> Result<DsparkSpecSession, Box<dyn std::error::Error>> {
use crate::cache::{Cache, DflashTapSink};
assert!(
self.cfg.gemma4.is_none(),
"gemma4 targets use the assistant-drafter route; dspark is the qwen-hybrid arm"
);
let n_embd = self.cfg.n_embd as usize;
let c = &draft.cfg;
assert_eq!(n_embd, c.hidden, "draft hidden must match target n_embd");
let b = c.block_size;
let n_taps = c.target_layer_ids.len();
let max_ctx = ctx_cap.min(c.sliding_window);
if prompt.len() + b + 8 > max_ctx {
return Err(format!(
"dspark session needs {} ctx (prompt {} + block {b} + 8), cap {max_ctx}",
prompt.len() + b + 8,
prompt.len()
)
.into());
}
let mut cache = Cache::new(e, &self.cfg, max_ctx)?;
let tp = prompt.len();
cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: e.uninit(tp * n_taps * n_embd)?,
hidden: n_embd,
t: tp,
base: 0,
});
let (logits, _h_seed, _hiddens) = self.prime_cache(e, prompt, &mut cache, 0)?;
let last = crate::forward::argmax(&logits) as u32;
let mut dkv = DflashKv::new(e, &draft.cfg, max_ctx)?;
{
let taps = cache.dflash_taps.take().unwrap();
let n_taps_h = n_taps * n_embd;
let mut r0 = 0usize;
while r0 < tp {
let t_c = (tp - r0).min(256);
let tv = e.view(&taps.buf, tp * n_taps_h);
let win = tv.slice(r0 * n_taps_h..(r0 + t_c) * n_taps_h);
let mut chunk = e.uninit(t_c * n_taps_h)?;
e.copy_view_into(&mut chunk, 0, &win, t_c * n_taps_h)?;
let f = draft.ctx_features(e, &chunk, t_c)?;
let pos_c: Vec<i32> = ((r0 as i32)..(r0 + t_c) as i32).collect();
draft.ingest_ctx(e, &mut dkv, &f, &pos_c, t_c)?;
r0 += t_c;
}
}
e.stream().synchronize()?;
let nd = DsparkHarvest::resolve(&draft.cfg).n_drafts(b);
let vt_cap: usize = std::env::var("MEMRA_DFLASH_VERIFY_T")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(nd + 1)
.clamp(2, nd + 1);
Ok(DsparkSpecSession {
cache,
dkv,
last,
ctx_len: tp,
vt: vt_cap,
rounds: 0,
max_ctx,
done: false,
snapb: None,
snapb_off: !crate::spec::state_copy_batch_on(),
})
}
pub fn dspark_spec_session_burst(
&self,
e: &Engine,
draft: &DflashDraft,
sess: &mut DsparkSpecSession,
burst_target: usize,
eos: &[u32],
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
use crate::cache::DflashTapSink;
let n_embd = self.cfg.n_embd as usize;
let c = &draft.cfg;
let b = c.block_size;
let n_taps = c.target_layer_ids.len();
let n_vocab = self.output.out_features();
let harvest = DsparkHarvest::resolve(&draft.cfg);
let nd = harvest.n_drafts(b);
let r0 = harvest.first_row();
let vt_cap: usize = std::env::var("MEMRA_DFLASH_VERIFY_T")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(nd + 1)
.clamp(2, nd + 1);
let adapt = std::env::var("MEMRA_DFLASH_ADAPT").as_deref() != Ok("0");
let vt_policy = DsparkVtPolicy::resolve(draft.confidence.is_some());
if vt_policy.is_confidence() {
assert!(
draft.confidence.is_some(),
"MEMRA_DSPARK_VT={vt_policy:?} needs a checkpoint with an accept-rate \
head (confidence_head.* absent in this export)"
);
}
let mut out: Vec<u32> = Vec::with_capacity(burst_target + b);
let mut drafted = 0usize;
let mut accepted_n = 0usize;
let defer_rb = crate::spec::dspark_defer_readback_on() && !vt_policy.is_confidence();
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let embd_gpu = if !defer_rb || crate::spec::spec_host_embd() {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
'outer: while out.len() < burst_target && !sess.done {
let start = sess.cache.pos;
if start + nd + 1 > sess.max_ctx {
sess.done = true;
break;
}
sess.rounds += 1;
let mut vt = sess.vt;
e.set_verify_exact(true);
let mut block: Vec<u32> = vec![c.mask_token_id; b];
block[0] = sess.last;
let noise = e.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(e, &mut sess.dkv, &noise, &pos_block)?;
let mut rows = e.uninit(nd * n_embd)?;
{
let dv = e.view(&dh, b * n_embd);
let src = dv.slice(r0 * n_embd..(r0 + nd) * n_embd);
e.copy_view_into(&mut rows, 0, &src, nd * n_embd)?;
}
let mut dl = e.matmul(&self.output, &rows, nd)?;
let markov_on = std::env::var("MEMRA_DFLASH_MARKOV").as_deref() != Ok("0");
let mut chain_d = e.stream().alloc_zeros::<u32>(nd + 1)?;
let want_conf_emb = vt_policy.is_confidence()
&& draft.confidence.as_ref().is_some_and(|ch| ch.with_markov);
let mut conf_emb: Option<CudaSlice<f32>> = match (&draft.markov, want_conf_emb) {
(Some(mk), true) => Some(e.uninit(nd * mk.rank)?),
(None, true) => unreachable!(
"with_markov confidence head without a markov table — the loader forbids it"
),
_ => None,
};
if let (Some(mk), true) = (&draft.markov, markov_on) {
e.set_u32_one(&mut chain_d, sess.last)?;
for k in 0..nd {
let mut f = e.uninit(mk.rank)?;
e.gather_row_bf16(&mk.w1_bf16, &chain_d, k, &mut f, mk.rank)?;
if let Some(ce) = conf_emb.as_mut() {
let fv = e.view(&f, mk.rank);
e.copy_view_into(ce, k * mk.rank, &fv, mk.rank)?;
}
let bias = e.matmul(&mk.w2, &f, 1)?;
e.add_row_inplace(&mut dl, &bias, n_vocab, k * n_vocab)?;
e.argmax_token_device_col(&dl, k, n_vocab, &mut chain_d, k + 1)?;
}
} else {
if want_conf_emb {
e.set_u32_one(&mut chain_d, sess.last)?;
}
for i in 0..nd {
if let (Some(ce), Some(mk)) = (conf_emb.as_mut(), &draft.markov) {
let mut f = e.uninit(mk.rank)?;
e.gather_row_bf16(&mk.w1_bf16, &chain_d, i, &mut f, mk.rank)?;
let fv = e.view(&f, mk.rank);
e.copy_view_into(ce, i * mk.rank, &fv, mk.rank)?;
}
e.argmax_token_device_col(&dl, i, n_vocab, &mut chain_d, i + 1)?;
}
}
e.set_verify_exact(false);
let ckpt_on = std::env::var("MEMRA_DSPARK_CKPT").as_deref() != Ok("0");
let deferred = embd_gpu.is_some() && ckpt_on;
if vt_policy.is_confidence() {
let ch = draft.confidence.as_ref().expect("asserted at burst entry");
let (rows_h, emb_h) = match conf_emb.as_ref() {
Some(ce) => {
let (a, b2) = e.dtoh_pair(&rows, ce)?;
(a, Some(b2))
}
None => (e.dtoh(&rows)?, None),
};
let rank = draft.markov.as_ref().map(|m| m.rank).unwrap_or(0);
let mut raws = Vec::with_capacity(nd);
for k in 0..nd {
let hrow = &rows_h[k * n_embd..(k + 1) * n_embd];
let emb = emb_h.as_ref().map(|eh| &eh[k * rank..(k + 1) * rank]);
raws.push(ch.raw_score(hrow, emb));
}
vt = vt_policy
.size_window(&raws, vt_cap)
.expect("confidence policies always size the window");
}
let mut cand: Vec<u32> = Vec::with_capacity(nd + 1);
if !deferred {
let chain = e.dtoh_u32(&chain_d)?;
cand.push(sess.last);
cand.extend_from_slice(&chain[1..]);
}
let mut snap_legacy: Option<crate::cache::CacheSnapshot> = None;
if !sess.snapb_off && sess.snapb.is_none() {
sess.snapb = DsparkSnapBatch::new(e, &sess.cache)?;
sess.snapb_off = sess.snapb.is_none();
} else if let Some(sb) = sess.snapb.as_mut() {
sb.refresh(e, &sess.cache)?;
}
let snap: &crate::cache::CacheSnapshot = match sess.snapb.as_ref() {
Some(sb) => &sb.snap,
None => {
snap_legacy = Some(sess.cache.snapshot(e)?);
snap_legacy.as_ref().unwrap()
}
};
let _ = &snap_legacy;
sess.cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: e.uninit(vt * n_taps * n_embd)?,
hidden: n_embd,
t: vt,
base: 0,
});
let (vam, vck) = if deferred {
let g = embd_gpu.expect("deferred implies resident embed");
let (am_d, vck) = self.dspark_verify_t_am_ckpt_dev(
e,
&chain_d,
vt,
start,
&mut sess.cache,
(g, embd_qt, embd_rb),
None,
)?;
let ch = e.stream().clone_dtoh(&chain_d)?;
let am = e.stream().clone_dtoh(&am_d)?;
e.stream().synchronize()?;
cand.push(sess.last);
cand.extend_from_slice(&ch[1..]);
(am, Some(vck))
} else if ckpt_on {
let (vam, vck) =
self.dspark_verify_t_am_ckpt(e, &cand[..vt], start, &mut sess.cache)?;
(vam, Some(vck))
} else {
(
self.dspark_verify_t_am(e, &cand[..vt], start, &mut sess.cache)?,
None,
)
};
let taps = sess.cache.dflash_taps.take().unwrap();
let m = dspark_accept_prefix(&cand, &vam, vt);
drafted += vt - 1;
accepted_n += m;
out.push(sess.last);
if eos.contains(&sess.last) {
sess.done = true;
break 'outer;
}
for &dt in &cand[1..=m] {
out.push(dt);
if eos.contains(&dt) {
sess.done = true;
break 'outer;
}
}
let next = vam[m];
let keep = m + 1;
if keep < vt {
if let Some(vck) = vck.as_ref() {
self.dspark_commit_prefix(e, &mut sess.cache, snap, vck, keep)?;
} else {
crate::pp::restore_cache_checkpoint(e, &self.cfg, None, &mut sess.cache, snap)?;
debug_assert_eq!(sess.cache.pos, start, "rollback landed off the round start");
let ram = self.dspark_verify_t_am(e, &cand[..keep], start, &mut sess.cache)?;
debug_assert_eq!(
&ram[..],
&vam[..keep],
"prefix replay must reproduce the verify argmaxes"
);
}
}
{
let tv = e.view(&taps.buf, vt * n_taps * n_embd);
let keep_view = tv.slice(0..keep * n_taps * n_embd);
let mut kept = e.uninit(keep * n_taps * n_embd)?;
e.copy_view_into(&mut kept, 0, &keep_view, keep * n_taps * n_embd)?;
let f = draft.ctx_features(e, &kept, keep)?;
let pos_k: Vec<i32> =
((sess.ctx_len as i32)..(sess.ctx_len + keep) as i32).collect();
draft.ingest_ctx(e, &mut sess.dkv, &f, &pos_k, keep)?;
sess.ctx_len += keep;
}
sess.last = next;
if vt_policy.is_confidence() {
sess.vt = vt;
} else if adapt {
sess.vt = (m + 2).clamp(3, vt_cap);
}
}
Ok((out, drafted, accepted_n))
}
}
#[cfg(test)]
mod dspark_harvest_tests {
use super::{DsparkHarvest, DsparkVtPolicy, dspark_accept_prefix, dspark_strategy_census};
const B: usize = 7;
#[test]
fn dspark_strategy_requires_shifted_harvest() {
let h = DsparkHarvest::Dspark;
assert_eq!(
h.first_row(),
0,
"DSPARK-strategy checkpoints (SpecForge OnlineDSparkModel, \
training.strategy=dspark — the q38 arm-a export) supervise ALL rows with \
SHIFTED labels: label_offsets = arange(1, block_size+1), i.e. the ANCHOR \
row's output is draft 1 (specforge/algorithms/common/\
dflash_family_model.py:816; sglang v0.5.17 dspark_draft.py:248,260). \
Harvesting from row 1 re-opens the DSPARK-POSTMORTEM-20260820 slot \
misalignment (accept 2.9 -> 1.43)."
);
assert_eq!(
h.n_drafts(B),
B,
"DSpark harvests gamma = block_size drafts per round (sglang \
dspark_config.py:269, verify_num_draft_tokens = gamma+1); b-1 is the \
DFlash mask-fill count and drops the best-trained slot \
(DSPARK-POSTMORTEM-20260820.md §3-H1)."
);
for row in 0..B {
assert_eq!(
h.trained_offset_of_row(row),
row + 1,
"OnlineDSparkModel trains row k to predict anchor+k+1 \
(dflash_family_model.py:816); a same-position (mask-fill) mapping \
here verifies every slot one position early — the postmortem's \
collapse."
);
}
}
#[test]
fn dflash_strategy_keeps_mask_fill_harvest() {
let h = DsparkHarvest::Dflash;
assert_eq!(h.first_row(), 1, "DFlash drafts start at mask row 1");
assert_eq!(h.n_drafts(B), B - 1, "DFlash harvests block_size-1 drafts");
for row in 1..B {
assert_eq!(h.trained_offset_of_row(row), row);
}
}
#[test]
fn every_candidate_verifies_the_position_its_row_was_trained_for() {
for h in [DsparkHarvest::Dflash, DsparkHarvest::Dspark] {
for i in 1..=h.n_drafts(B) {
let row = h.first_row() + i - 1;
assert_eq!(
h.trained_offset_of_row(row),
i,
"{h:?}: candidate {i} rides row {row}, which is trained for \
offset {} — harvest misaligned",
h.trained_offset_of_row(row)
);
}
}
}
#[test]
fn env_seam_parses_and_refuses() {
assert_eq!(
DsparkHarvest::from_env_value(None),
DsparkHarvest::Dflash,
"the ENV-ONLY parser keeps the historical arm; the ratified strategy-keyed \
default lives in resolve_value (checkpoint census), not here"
);
assert_eq!(
DsparkHarvest::from_env_value(Some("dspark")),
DsparkHarvest::Dspark
);
assert_eq!(
DsparkHarvest::from_env_value(Some("dflash")),
DsparkHarvest::Dflash
);
assert!(
std::panic::catch_unwind(|| DsparkHarvest::from_env_value(Some("shifted"))).is_err(),
"unknown harvest values must REFUSE, not default"
);
assert_eq!(
DsparkHarvest::from_name("dspark"),
Some(DsparkHarvest::Dspark)
);
assert_eq!(
DsparkHarvest::from_name("dflash"),
Some(DsparkHarvest::Dflash)
);
assert_eq!(DsparkHarvest::from_name("mask-fill"), None);
}
#[test]
fn ratified_default_harvest_is_strategy_keyed() {
assert_eq!(
DsparkHarvest::resolve_value(None, true),
DsparkHarvest::Dspark,
"owner-ratified 2026-08-20: unset env defaults a DSPARK-strategy \
checkpoint to the shifted harvest (DSPARK-POSTMORTEM-20260820.md B1)"
);
assert_eq!(
DsparkHarvest::resolve_value(None, false),
DsparkHarvest::Dflash
);
assert_eq!(
DsparkHarvest::resolve_value(Some(""), false),
DsparkHarvest::Dflash
);
assert_eq!(
DsparkHarvest::resolve_value(Some("dflash"), true),
DsparkHarvest::Dflash
);
assert_eq!(
DsparkHarvest::resolve_value(Some("dspark"), false),
DsparkHarvest::Dspark
);
assert!(
std::panic::catch_unwind(|| DsparkHarvest::resolve_value(Some("shifted"), true))
.is_err()
);
}
#[test]
fn strategy_census_reads_the_checkpoint_not_the_env() {
let q38 = r#"{"architectures": ["Qwen3DSparkModel"], "block_size": 7,
"dflash_config": {"projector_type": "dspark", "markov_rank": 256}}"#;
assert!(dspark_strategy_census(q38));
assert!(dspark_strategy_census(
r#"{"architectures": ["Qwen3DSparkModel"]}"#
));
assert!(dspark_strategy_census(
r#"{"dflash_config": {"projector_type": "dspark"}}"#
));
let dflash = r#"{"architectures": ["Qwen3DFlashModel"],
"dflash_config": {"attention_mode": "gqa"}}"#;
assert!(!dspark_strategy_census(dflash));
assert!(!dspark_strategy_census("{}"));
}
#[test]
fn ratified_default_vt_is_confidence_slot_tau_half() {
assert_eq!(
DsparkVtPolicy::resolve_value(None, None, None, true),
DsparkVtPolicy::ConfidenceSlot { tau: 0.5 },
"owner-ratified 2026-08-20: unset MEMRA_DSPARK_VT defaults to \
confidence-slot tau=.5 on a head-carrying checkpoint (H4 cells 2-3)"
);
assert_eq!(
DsparkVtPolicy::resolve_value(None, Some("0.35"), None, true),
DsparkVtPolicy::ConfidenceSlot { tau: 0.35 }
);
assert!(
std::panic::catch_unwind(|| DsparkVtPolicy::resolve_value(
None,
Some("nan-ish"),
None,
true
))
.is_err()
);
assert_eq!(
DsparkVtPolicy::resolve_value(None, None, None, false),
DsparkVtPolicy::Ladder
);
assert_eq!(
DsparkVtPolicy::resolve_value(None, None, Some("0"), true),
DsparkVtPolicy::Ladder
);
assert_eq!(
DsparkVtPolicy::resolve_value(Some("ladder"), None, None, true),
DsparkVtPolicy::Ladder
);
assert_eq!(
DsparkVtPolicy::resolve_value(Some("confidence"), Some("0.35"), None, true),
DsparkVtPolicy::Confidence { tau: 0.35 }
);
assert!(
std::panic::catch_unwind(|| DsparkVtPolicy::resolve_value(
Some("confidence-slot"),
None,
Some("0"),
true
))
.is_err()
);
}
#[test]
fn dspark_trained_rows_through_mask_fill_harvest_accept_nothing() {
const BASE: u32 = 1000;
let anchor: u32 = BASE; let vam: Vec<u32> = (1..=B as u32 + 1).map(|j| BASE + j).collect();
let dspark_trained_row_argmax =
|r: usize| BASE + DsparkHarvest::Dspark.trained_offset_of_row(r) as u32;
let h = DsparkHarvest::Dspark;
let mut cand = vec![anchor];
for i in 1..=h.n_drafts(B) {
cand.push(dspark_trained_row_argmax(h.first_row() + i - 1));
}
let vt = h.n_drafts(B) + 1;
assert_eq!(
dspark_accept_prefix(&cand, &vam, vt),
vt - 1,
"aligned harvest must accept the full block"
);
let wrong = DsparkHarvest::Dflash;
let mut cand_wrong = vec![anchor];
for i in 1..=wrong.n_drafts(B) {
cand_wrong.push(dspark_trained_row_argmax(wrong.first_row() + i - 1));
}
let vt_wrong = wrong.n_drafts(B) + 1;
assert_eq!(
dspark_accept_prefix(&cand_wrong, &vam, vt_wrong),
0,
"mask-fill harvest of a dspark-trained drafter verifies every slot against \
a position the row was not trained for (DSPARK-POSTMORTEM-20260820.md)"
);
}
}
#[cfg(test)]
mod dspark_vt_tests {
use super::{ConfidenceHead, DsparkVtPolicy, dspark_confidence_vt, dspark_slot_confidence_vt};
fn logit(p: f32) -> f32 {
(p / (1.0 - p)).ln()
}
#[test]
fn confidence_vt_is_cumprod_survival_not_per_slot_threshold() {
let raws: Vec<f32> = [0.9, 0.8, 0.9, 0.9, 0.9, 0.9, 0.9]
.iter()
.map(|&p| logit(p))
.collect();
assert_eq!(
dspark_confidence_vt(&raws, 0.5, 8),
6,
"keeps 5 drafts + anchor"
);
assert_eq!(
dspark_confidence_vt(&raws, 0.7, 8),
3,
"tau=0.7 keeps 2 drafts"
);
assert_eq!(
dspark_confidence_vt(&raws, 0.05, 8),
8,
"tau→0 = full block"
);
}
#[test]
fn slot_arm_truncates_at_first_low_confidence_slot() {
let raws: Vec<f32> = [0.9, 0.8, 0.9, 0.9, 0.9, 0.9, 0.9]
.iter()
.map(|&p| logit(p))
.collect();
assert_eq!(dspark_slot_confidence_vt(&raws, 0.5, 8), 8);
assert_eq!(dspark_confidence_vt(&raws, 0.5, 8), 6);
let tail: Vec<f32> = [0.9, 0.9, 0.3, 0.9, 0.9, 0.9, 0.9]
.iter()
.map(|&p| logit(p))
.collect();
assert_eq!(
dspark_slot_confidence_vt(&tail, 0.5, 8),
3,
"2 drafts + anchor"
);
assert_eq!(
dspark_slot_confidence_vt(&tail, 0.95, 8),
2,
"floor at tau=0.95"
);
}
#[test]
fn confidence_vt_floor_and_cap() {
let cold: Vec<f32> = [0.1f32, 0.1, 0.1].iter().map(|&p| logit(p)).collect();
assert_eq!(
dspark_confidence_vt(&cold, 0.5, 8),
2,
"floor = anchor + 1 draft"
);
assert_eq!(
dspark_slot_confidence_vt(&cold, 0.5, 8),
2,
"slot arm same floor"
);
let hot: Vec<f32> = vec![logit(0.99); 7];
assert_eq!(dspark_confidence_vt(&hot, 0.5, 5), 5, "vt_cap binds");
assert_eq!(
dspark_confidence_vt(&hot, 0.5, 8),
8,
"full block when confident"
);
assert_eq!(
dspark_slot_confidence_vt(&hot, 0.5, 5),
5,
"slot arm same cap"
);
assert_eq!(dspark_confidence_vt(&[], 0.5, 8), 2);
assert_eq!(dspark_slot_confidence_vt(&[], 0.5, 8), 2);
}
#[test]
fn vt_policy_env_seam_parses_and_refuses() {
assert_eq!(
DsparkVtPolicy::from_env_value(None, None, None),
DsparkVtPolicy::Ladder,
"default stays the shipped ladder — the H4 arm is opt-in"
);
assert_eq!(
DsparkVtPolicy::from_env_value(Some(""), None, None),
DsparkVtPolicy::Ladder
);
assert_eq!(
DsparkVtPolicy::from_env_value(Some("ladder"), None, Some("0")),
DsparkVtPolicy::Ladder,
"ladder + ADAPT=0 = the fixed-window arm, untouched"
);
assert_eq!(
DsparkVtPolicy::from_env_value(Some("confidence"), None, None),
DsparkVtPolicy::Confidence { tau: 0.5 },
"tau defaults to 0.5 (raw sigmoid, no STS sidecar — postmortem §3-H4)"
);
assert_eq!(
DsparkVtPolicy::from_env_value(Some("confidence"), Some("0.35"), Some("1")),
DsparkVtPolicy::Confidence { tau: 0.35 }
);
assert_eq!(
DsparkVtPolicy::from_env_value(Some("confidence-slot"), Some("0.6"), None),
DsparkVtPolicy::ConfidenceSlot { tau: 0.6 },
"the owner-directive per-slot arm parses with the same tau env"
);
assert!(
std::panic::catch_unwind(|| DsparkVtPolicy::from_env_value(
Some("confidence-slot"),
None,
Some("0")
))
.is_err(),
"confidence-slot + MEMRA_DFLASH_ADAPT=0 must REFUSE like confidence"
);
assert!(
std::panic::catch_unwind(|| DsparkVtPolicy::from_env_value(Some("static"), None, None))
.is_err(),
"unknown policy values must REFUSE, not default — a typo silently \
reverting the window policy invalidates an A/B"
);
assert!(
std::panic::catch_unwind(|| DsparkVtPolicy::from_env_value(
Some("confidence"),
None,
Some("0")
))
.is_err(),
"confidence + MEMRA_DFLASH_ADAPT=0 is contradictory and must REFUSE"
);
for bad in ["0", "1", "1.5", "-0.1", "nan"] {
assert!(
std::panic::catch_unwind(|| DsparkVtPolicy::from_env_value(
Some("confidence"),
Some(bad),
None
))
.is_err(),
"tau={bad} must REFUSE (survival threshold lives in (0,1))"
);
}
}
#[test]
fn raw_score_matches_the_parity_gate_dot() {
let ch = ConfidenceHead {
w: vec![0.5, -1.0, 2.0, 0.25, -0.5],
b: 0.125,
in_dim: 5,
with_markov: true,
};
let hidden = [1.0f32, 2.0, 3.0];
let emb = [4.0f32, 8.0];
let want = 0.125 + 0.5 * 1.0 - 1.0 * 2.0 + 2.0 * 3.0 + 0.25 * 4.0 - 0.5 * 8.0;
assert_eq!(ch.raw_score(&hidden, Some(&emb)), want);
let ch_plain = ConfidenceHead {
w: vec![0.5, -1.0, 2.0],
b: -0.25,
in_dim: 3,
with_markov: false,
};
let want_plain = -0.25 + 0.5 * 1.0 - 1.0 * 2.0 + 2.0 * 3.0;
assert_eq!(ch_plain.raw_score(&hidden, None), want_plain);
}
}