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 is_causal: Option<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 dflash2: Option<Dflash2Head>,
}
pub struct Dflash2Conv {
pub base: CudaSlice<f32>,
pub proj: GpuTensor,
}
pub struct Dflash2Head {
pub attn_conv: Vec<Dflash2Conv>, pub mlp_conv: Vec<Dflash2Conv>, pub hidden_proj: GpuTensor,
pub pred_codebook: Vec<u8>,
pub succ_codebook: Vec<u8>,
pub rank: usize, pub top_k: usize, pub conv_k: usize, pub group_size: usize, pub vocab: usize, }
fn cb_row(cb: &[u8], tok: usize, rank: usize) -> Vec<f32> {
bf16_to_f32(&cb[tok * rank * 2..(tok + 1) * rank * 2])
}
#[allow(clippy::too_many_arguments)]
pub fn dflash2_walk_greedy(
pred_codebook: &[u8],
succ_codebook: &[u8],
vocab: usize,
rank: usize,
top_k: usize,
unary: &[f32],
cand: &[u32],
hproj: &[f32],
anchor: u32,
nd: usize,
) -> Vec<u32> {
let (kk, r) = (top_k, rank);
assert_eq!(unary.len(), nd * kk, "walk: unary shape");
assert_eq!(cand.len(), nd * kk, "walk: candidate shape");
assert_eq!(hproj.len(), nd * r, "walk: hidden-projection shape");
let mut path = Vec::with_capacity(nd);
let mut prev = anchor;
for p in 0..nd {
assert!(
(prev as usize) < vocab,
"walk: predecessor token {prev} outside codebook vocab {vocab}"
);
let pr = cb_row(pred_codebook, prev as usize, r);
let hp = &hproj[p * r..(p + 1) * r];
let gate: Vec<f32> = pr.iter().zip(hp).map(|(a, b)| a * b).collect();
let (mut best, mut bi) = (f32::NEG_INFINITY, 0usize);
for k in 0..kk {
let c = cand[p * kk + k] as usize;
assert!(c < vocab, "walk: candidate {c} outside codebook vocab");
let sr = cb_row(succ_codebook, c, r);
let mut s = unary[p * kk + k];
for j in 0..r {
s += gate[j] * sr[j];
}
if s > best {
best = s;
bi = k;
}
}
prev = cand[p * kk + bi];
path.push(prev);
}
path
}
impl Dflash2Head {
pub fn walk_greedy(
&self,
unary: &[f32],
cand: &[u32],
hproj: &[f32],
anchor: u32,
nd: usize,
) -> Vec<u32> {
dflash2_walk_greedy(
&self.pred_codebook,
&self.succ_codebook,
self.vocab,
self.rank,
self.top_k,
unary,
cand,
hproj,
anchor,
nd,
)
}
#[allow(clippy::too_many_arguments)]
pub fn walk_sampled(
&self,
unary: &[f32],
cand: &[u32],
hproj: &[f32],
anchor: u32,
nd: usize,
temp: f32,
uniforms: &mut dyn FnMut() -> f32,
) -> (Vec<u32>, Vec<f32>, Vec<f32>) {
dflash2_walk_sampled(
&self.pred_codebook,
&self.succ_codebook,
self.vocab,
self.rank,
self.top_k,
unary,
cand,
hproj,
anchor,
nd,
temp,
uniforms,
)
}
}
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 for_draft(draft: &DflashDraft) -> Self {
Self::for_family_value(
draft.dflash2.is_some(),
std::env::var("MEMRA_DSPARK_HARVEST").ok().as_deref(),
draft.cfg.strategy_dspark,
)
}
pub fn for_family_value(is_dflash2: bool, env: Option<&str>, strategy_dspark: bool) -> Self {
if is_dflash2 {
if env == Some("dspark") {
panic!(
"MEMRA_DSPARK_HARVEST=dspark with a DFlash2 checkpoint: DFlash2 \
is mask-fill (b-1 drafts, anchor row is not a draft — reference \
dflash_generate rows 1-verify_size:); the shifted harvest would \
verify every slot one position early. Refusing (census-keyed, \
not env-keyed)."
);
}
return DsparkHarvest::Dflash;
}
Self::resolve_value(env, strategy_dspark)
}
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))
}
pub fn rejection_accept_len(p: &[f32], q: &[f32], u: &[f32]) -> usize {
assert!(
q.len() >= p.len() && u.len() >= p.len(),
"accept walk shape"
);
let mut m = 0usize;
while m < p.len() && (u[m] as f64) * (q[m] as f64) < p[m] as f64 {
m += 1;
}
m
}
#[allow(clippy::too_many_arguments)]
pub fn dflash2_walk_sampled(
pred_codebook: &[u8],
succ_codebook: &[u8],
vocab: usize,
rank: usize,
top_k: usize,
unary: &[f32],
cand: &[u32],
hproj: &[f32],
anchor: u32,
nd: usize,
temp: f32,
uniforms: &mut dyn FnMut() -> f32,
) -> (Vec<u32>, Vec<f32>, Vec<f32>) {
assert!(
temp > 0.0,
"sampled walk is the T>0 arm; T=0 is walk_greedy"
);
let (kk, r) = (top_k, rank);
assert_eq!(unary.len(), nd * kk, "walk: unary shape");
assert_eq!(cand.len(), nd * kk, "walk: candidate shape");
assert_eq!(hproj.len(), nd * r, "walk: hidden-projection shape");
let mut path = Vec::with_capacity(nd);
let mut q_chosen = Vec::with_capacity(nd);
let mut q_rows = Vec::with_capacity(nd * kk);
let mut prev = anchor;
for p in 0..nd {
assert!(
(prev as usize) < vocab,
"walk: predecessor token {prev} outside codebook vocab {vocab}"
);
let pr = cb_row(pred_codebook, prev as usize, r);
let hp = &hproj[p * r..(p + 1) * r];
let gate: Vec<f32> = pr.iter().zip(hp).map(|(a, b)| a * b).collect();
let mut scores = vec![0f32; kk];
for (k, s) in scores.iter_mut().enumerate() {
let c = cand[p * kk + k] as usize;
assert!(c < vocab, "walk: candidate {c} outside codebook vocab");
let sr = cb_row(succ_codebook, c, r);
let mut acc = unary[p * kk + k];
for j in 0..r {
acc += gate[j] * sr[j];
}
*s = acc;
}
let mx = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut z = 0f64;
let ex: Vec<f64> = scores
.iter()
.map(|&s| {
let e0 = (((s - mx) / temp) as f64).exp();
z += e0;
e0
})
.collect();
let probs: Vec<f32> = ex.iter().map(|&e0| (e0 / z) as f32).collect();
let u = uniforms() as f64;
let mut acc = 0f64;
let mut bi = probs
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.map(|(k, _)| k)
.unwrap_or(0);
for (k, &pk) in probs.iter().enumerate() {
acc += pk as f64;
if u < acc {
bi = k;
break;
}
}
prev = cand[p * kk + bi];
path.push(prev);
q_chosen.push(probs[bi]);
q_rows.extend_from_slice(&probs);
}
(path, q_chosen, q_rows)
}
pub(crate) type Dflash2SampledProposal = (Vec<u32>, Vec<f32>, Vec<u32>, Vec<f32>);
pub(crate) enum DsparkDraftSample {
Rows {
th: CudaSlice<f32>, z: CudaSlice<f32>, stats: Vec<(f32, f32, f32)>, },
Selector {
cand: Vec<u32>, q_rows: Vec<f32>, q_chosen: Vec<f32>, top_k: usize,
},
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dspark_accept_sampled(
e: &Engine,
tlogits: &CudaSlice<f32>,
cand: &[u32],
vt: usize,
n_vocab: usize,
dl: &CudaSlice<f32>,
prop: &DsparkDraftSample,
sp: &crate::spec::SpecSampling,
sctr: &mut u32,
uctr: &mut u32,
) -> Result<(usize, u32), Box<dyn std::error::Error>> {
let nq = vt - 1; debug_assert!(nq >= 1 && cand.len() > nq, "sampled accept shape");
let rows: Vec<i32> = (0..nq as i32).collect();
let ids: Vec<u32> = cand[1..=nq].to_vec();
let rowsd = e.htod_i32(&rows)?;
let idsd = e.htod_u32_v(&ids)?;
let (mut pth, mut pz, mut pmx) = (e.zeros(nq)?, e.zeros(nq)?, e.zeros(nq)?);
e.filter_stats(
tlogits, n_vocab, &rowsd, &mut pth, &mut pz, &mut pmx, n_vocab, nq, sp.temp, sp.top_k,
sp.top_p, sp.min_p,
)?;
let mut pj_d = e.zeros(nq)?;
e.softmax_gather_filtered(
tlogits, n_vocab, &idsd, &rowsd, &pth, &pz, &mut pj_d, n_vocab, nq, sp.temp,
)?;
let pj = e.dtoh(&pj_d)?;
let (pthv, pzv, pmxv) = (e.dtoh(&pth)?, e.dtoh(&pz)?, e.dtoh(&pmx)?);
let qj: Vec<f32> = match prop {
DsparkDraftSample::Rows { th, z, .. } => {
let mut qd = e.zeros(nq)?;
e.softmax_gather_filtered(
dl, n_vocab, &idsd, &rowsd, th, z, &mut qd, n_vocab, nq, sp.temp,
)?;
e.dtoh(&qd)?
}
DsparkDraftSample::Selector { q_chosen, .. } => q_chosen[..nq].to_vec(),
};
let mut us = Vec::with_capacity(nq);
for _ in 0..nq {
us.push(crate::spec::host_u01(sp.seed, *uctr));
*uctr = uctr.wrapping_add(1);
}
let m = rejection_accept_len(&pj[..nq], &qj[..nq], &us);
let next = if m == nq {
let rows_l = e.htod_i32(&[(vt - 1) as i32])?;
let (mut th1, mut z1, mut mx1) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
tlogits, n_vocab, &rows_l, &mut th1, &mut z1, &mut mx1, n_vocab, 1, sp.temp, sp.top_k,
sp.top_p, sp.min_p,
)?;
let mut pb = e.zeros(n_vocab)?;
e.gumbel_perturb_filtered_col(
tlogits,
vt - 1,
&mut pb,
n_vocab,
sp.seed,
*sctr,
sp.temp,
&mx1,
&th1,
0,
)?;
*sctr = sctr.wrapping_add(1);
let td = e.argmax_token_device(&pb, n_vocab)?;
e.dtoh_u32_one(&td)?
} else {
let mut col = e.zeros(n_vocab)?;
e.copy_view_into(
&mut col,
0,
&tlogits.slice(m * n_vocab..(m + 1) * n_vocab),
n_vocab,
)?;
let p_stats = (pmxv[m], pthv[m], pzv[m]);
let mut tok_d = e.alloc_u32_zeroed(1)?;
let sc = *sctr;
*sctr = sctr.wrapping_add(1);
match prop {
DsparkDraftSample::Rows { stats, .. } => {
let mut qbuf = e.zeros(n_vocab)?;
e.copy_view_into(
&mut qbuf,
0,
&dl.slice(m * n_vocab..(m + 1) * n_vocab),
n_vocab,
)?;
e.residual_sample_filtered(
&col,
Some(&qbuf),
n_vocab,
sp.temp,
sp.seed,
sc,
p_stats,
stats[m],
&mut tok_d,
)?;
}
DsparkDraftSample::Selector {
cand: cids,
q_rows,
top_k,
..
} => {
let k = *top_k;
let ids_m = e.htod_u32_v(&cids[m * k..(m + 1) * k])?;
let qs_m = e.htod(&q_rows[m * k..(m + 1) * k])?;
e.residual_sample_sparse_q(
&col, &ids_m, &qs_m, k, n_vocab, sp.temp, sp.seed, sc, p_stats, &mut tok_d,
)?;
}
}
e.dtoh_u32(&tok_d)?[0]
};
Ok((m, next))
}
fn dflash2_sdpa_clip_on() -> bool {
std::env::var("MEMRA_DFLASH2_SDPA_CLIP")
.map(|v| v != "0")
.unwrap_or(true)
}
#[allow(clippy::too_many_arguments)]
fn d2_windowed_attn(
e: &Engine,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
attn: &mut CudaSlice<f32>,
hd: usize,
nh: usize,
nkv: usize,
t: usize,
t_kv: usize,
scale: f32,
c: &DflashCfg,
) -> Result<(), Box<dyn std::error::Error>> {
if dflash2_sdpa_clip_on() {
e.sdpa_naive_w_lo(
q,
k,
v,
attn,
hd,
nh,
nkv,
t,
t_kv,
scale,
false,
c.sliding_window,
)
} else {
e.sdpa_naive_w(
q,
k,
v,
attn,
hd,
nh,
nkv,
t,
t_kv,
scale,
false,
c.sliding_window,
)
}
}
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()
}
fn scope<'a>(txt: &'a str, key: &str) -> Option<&'a str> {
let i = txt.find(&format!("\"{key}\""))?;
let rest = &txt[i..];
let open = rest.find('{')?;
let mut depth = 0usize;
for (j, ch) in rest[open..].char_indices() {
match ch {
'{' => depth += 1,
'}' => {
depth -= 1;
if depth == 0 {
return Some(&rest[open..open + j + 1]);
}
}
_ => {}
}
}
None
}
let is_dflash2 = {
let arch = scope_list(&txt, "architectures");
arch.contains("DFlash2DraftModel")
};
fn scope_list(txt: &str, key: &str) -> String {
let Some(i) = txt.find(&format!("\"{key}\"")) else {
return String::new();
};
let rest = &txt[i..];
match (rest.find('['), rest.find(']')) {
(Some(a), Some(b)) if a < b => rest[a + 1..b].to_string(),
_ => String::new(),
}
}
let d2_cfg_txt: Option<&str> = if is_dflash2 {
Some(scope(&txt, "dflash_config").unwrap_or_else(|| {
panic!("DFlash2DraftModel config.json has no dflash_config object — refusing")
}))
} else {
None
};
let g = |k: &str| num(&txt, k).unwrap_or_else(|| panic!("config missing {k}")) as usize;
let g2 = |k: &str| -> usize {
let t = d2_cfg_txt.expect("dflash2 scope");
num(t, k).unwrap_or_else(|| panic!("dflash_config missing {k} — refusing")) 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 is_causal = txt
.find("\"is_causal\"")
.and_then(|i| txt[i..].find(':').map(|c| i + c + 1))
.map(|v| txt[v..].trim_start().starts_with("true"));
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: if is_dflash2 {
let rp = scope(&txt, "rope_parameters")
.unwrap_or_else(|| panic!("DFlash2 config has no rope_parameters — refusing"));
assert!(
rp.contains("\"default\""),
"DFlash2 rope_parameters rope_type is not \"default\" — the port \
implements plain neox rope only; refusing ({rp})"
);
num(rp, "rope_theta").expect("rope_parameters.rope_theta") as f32
} else {
num(&txt, "rope_theta").expect("rope_theta") as f32
},
block_size: if is_dflash2 {
g2("block_size")
} else {
g("block_size")
},
mask_token_id: if is_dflash2 {
g2("mask_token_id")
} else {
g("mask_token_id")
} as u32,
target_layer_ids: if is_dflash2 {
num_list(d2_cfg_txt.expect("dflash2 scope"), "target_layer_ids")
} else {
num_list(&txt, "target_layer_ids")
},
sliding_window,
layer_sliding,
strategy_dspark: dspark_strategy_census(&txt),
is_causal,
};
if is_dflash2 {
assert_eq!(
cfg.is_causal,
Some(false),
"DFlash2 port requires explicit config is_causal=false \
(non-causal symmetric sliding window); got {is_causal:?} — refusing"
);
assert!(
cfg.layer_sliding.iter().all(|&s| s),
"DFlash2 port expects all layers sliding_attention (q38 export); \
got {:?} — refusing (unverified mask program)",
cfg.layer_sliding
);
assert!(
cfg.block_size <= cfg.sliding_window,
"DFlash2 block {} exceeds the sliding window {} — the windowed SDPA \
omits the future-side mask because block rows stay within the window",
cfg.block_size,
cfg.sliding_window
);
}
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 dflash2 = if is_dflash2 {
assert!(
markov.is_none() && confidence.is_none(),
"DFlash2 checkpoint carries markov/confidence tensors — no such \
variant exists in the family (census refuses the ambiguity)"
);
let rank = g2("selector_rank");
let top_k = g2("selector_top_k");
let conv_k = g2("conv_kernel_size");
let group_size = g2("conv_group_size");
let groups = cfg.hidden / group_size;
let load_conv = |name: &str| -> Result<Dflash2Conv, Box<dyn std::error::Error>> {
let (bi, bb) = st
.raw(&format!("{name}.base_kernel"))
.ok_or_else(|| format!("DFlash2 census: missing {name}.base_kernel"))?;
let bne = bi.ne();
assert_eq!(
(bne[0] as usize, bne[1] as usize, bne[2] as usize),
(cfg.hidden, conv_k, 2),
"{name}.base_kernel shape != [2, conv_kernel_size, hidden]"
);
let pname = format!("{name}.kernel_projection.weight");
let (pi, _pb) = st
.raw(&pname)
.ok_or_else(|| format!("DFlash2 census: missing {pname}"))?;
let pne = pi.ne(); assert_eq!(
(pne[0] as usize, pne[1] as usize),
(cfg.hidden, 2 * conv_k * groups),
"{pname} shape != [2*conv_kernel_size*groups, hidden]"
);
Ok(Dflash2Conv {
base: e.htod(&bf16_to_f32(bb))?,
proj: upw(&pname)?,
})
};
let mut attn_conv = Vec::with_capacity(cfg.n_layer);
let mut mlp_conv = Vec::with_capacity(cfg.n_layer);
for i in 0..cfg.n_layer {
attn_conv.push(load_conv(&format!("layers.{i}.attention_conv"))?);
mlp_conv.push(load_conv(&format!("layers.{i}.mlp_conv"))?);
}
let cb = |name: &str| -> Result<(Vec<u8>, usize), Box<dyn std::error::Error>> {
let (ci, cbytes) = st
.raw(&format!("candidate_selector.{name}"))
.ok_or_else(|| format!("DFlash2 census: missing candidate_selector.{name}"))?;
let ne = ci.ne(); assert_eq!(ne[0] as usize, rank, "candidate_selector.{name} rank");
Ok((cbytes.to_vec(), ne[1] as usize))
};
let (pred_codebook, v1) = cb("predecessor_codebook")?;
let (succ_codebook, v2) = cb("successor_codebook")?;
assert_eq!(v1, v2, "codebook vocab mismatch");
let hp_name = "candidate_selector.hidden_projection.weight";
let (hi, _hb) = st
.raw(hp_name)
.ok_or_else(|| format!("DFlash2 census: missing {hp_name}"))?;
assert_eq!(
(hi.ne()[0] as usize, hi.ne()[1] as usize),
(cfg.hidden, rank),
"{hp_name} shape != [rank, hidden]"
);
Some(Dflash2Head {
attn_conv,
mlp_conv,
hidden_proj: upw(hp_name)?,
pred_codebook,
succ_codebook,
rank,
top_k,
conv_k,
group_size,
vocab: v1,
})
} 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}"));
}
if dflash2.is_some() {
for s in [
"attention_conv.base_kernel",
"attention_conv.kernel_projection.weight",
"mlp_conv.base_kernel",
"mlp_conv.kernel_projection.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());
}
if dflash2.is_some() {
for s in [
"candidate_selector.hidden_projection.weight",
"candidate_selector.predecessor_codebook",
"candidate_selector.successor_codebook",
] {
consumed.insert(s.into());
}
}
let leftovers: Vec<&String> = st.names().filter(|n| !consumed.contains(*n)).collect();
if !leftovers.is_empty() {
if markov.is_some() || dflash2.is_some() {
panic!("dspark/dflash2 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 dflash2={} strategy_dspark={}, \
MEMRA_DSPARK_HARVEST {})",
DsparkHarvest::for_family_value(
dflash2.is_some(),
std::env::var("MEMRA_DSPARK_HARVEST").ok().as_deref(),
cfg.strategy_dspark,
)
.name(),
dflash2.is_some(),
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,
dflash2,
})
}
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 d2_conv_prepare(
&self,
e: &Engine,
conv: &Dflash2Conv,
xn: &CudaSlice<f32>,
rows: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let d2 = self
.dflash2
.as_ref()
.expect("d2_conv on a non-dflash2 draft");
let h = self.cfg.hidden;
let groups = h / d2.group_size;
let dyn_ = self.mm(e, &conv.proj, xn, rows, h, 2 * d2.conv_k * groups)?;
let mut out = e.uninit(rows * h)?;
e.dflash2_dynconv(
xn,
&dyn_,
&conv.base,
&mut out,
rows,
h,
d2.group_size,
d2.conv_k,
0,
)?;
Ok((out, dyn_))
}
pub fn d2_conv_finish(
&self,
e: &Engine,
conv: &Dflash2Conv,
y: &CudaSlice<f32>,
dyn_: &CudaSlice<f32>,
rows: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let d2 = self
.dflash2
.as_ref()
.expect("d2_conv on a non-dflash2 draft");
let h = self.cfg.hidden;
let mut out = e.uninit(rows * h)?;
e.dflash2_dynconv(
y,
dyn_,
&conv.base,
&mut out,
rows,
h,
d2.group_size,
d2.conv_k,
1,
)?;
Ok(out)
}
pub fn dflash2_propose_greedy(
&self,
e: &Engine,
dl: &CudaSlice<f32>,
rows: &CudaSlice<f32>,
nd: usize,
n_vocab: usize,
anchor: u32,
) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
let d2 = self
.dflash2
.as_ref()
.expect("dflash2_propose on a non-dflash2 draft");
assert!(
n_vocab <= d2.vocab,
"target head vocab {n_vocab} exceeds the selector codebooks ({})",
d2.vocab
);
let (vals_d, idx_d) = e.topk_rows(dl, nd, n_vocab, d2.top_k)?;
let hproj_d = e.matmul(&d2.hidden_proj, rows, nd)?;
let unary = e.dtoh(&vals_d)?;
let cand = e.dtoh_u32(&idx_d)?;
let hproj = e.dtoh(&hproj_d)?;
Ok(d2.walk_greedy(&unary, &cand, &hproj, anchor, nd))
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dflash2_propose_sampled(
&self,
e: &Engine,
dl: &CudaSlice<f32>,
rows: &CudaSlice<f32>,
nd: usize,
n_vocab: usize,
anchor: u32,
temp: f32,
seed: u64,
uctr: &mut u32,
) -> Result<Dflash2SampledProposal, Box<dyn std::error::Error>> {
let d2 = self
.dflash2
.as_ref()
.expect("dflash2_propose on a non-dflash2 draft");
assert!(
n_vocab <= d2.vocab,
"target head vocab {n_vocab} exceeds the selector codebooks ({})",
d2.vocab
);
let (vals_d, idx_d) = e.topk_rows(dl, nd, n_vocab, d2.top_k)?;
let hproj_d = e.matmul(&d2.hidden_proj, rows, nd)?;
let unary = e.dtoh(&vals_d)?;
let cand = e.dtoh_u32(&idx_d)?;
let hproj = e.dtoh(&hproj_d)?;
let mut draw = || {
let u = crate::spec::host_u01(seed, *uctr);
*uctr = uctr.wrapping_add(1);
u
};
let (path, q_chosen, q_rows) =
d2.walk_sampled(&unary, &cand, &hproj, anchor, nd, temp, &mut draw);
Ok((path, q_chosen, cand, q_rows))
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dspark_chain_sampled(
&self,
e: &Engine,
dl: &mut CudaSlice<f32>,
nd: usize,
n_vocab: usize,
anchor: u32,
sp: &crate::spec::SpecSampling,
sctr: &mut u32,
mut conf_emb: Option<&mut CudaSlice<f32>>,
) -> Result<(Vec<u32>, DsparkDraftSample), Box<dyn std::error::Error>> {
let markov_on = std::env::var("MEMRA_DFLASH_MARKOV").as_deref() != Ok("0");
let mut chain_d = e.stream().alloc_zeros::<u32>(nd + 1)?;
e.set_u32_one(&mut chain_d, anchor)?;
let mut th_all = e.zeros(nd)?;
let mut z_all = e.zeros(nd)?;
let mut mx_all = e.zeros(nd)?;
let mut pb = e.zeros(n_vocab)?;
for k in 0..nd {
if let (Some(mk), true) = (&self.markov, markov_on) {
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_deref_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(dl, &bias, n_vocab, k * n_vocab)?;
} else if let (Some(ce), Some(mk)) = (conf_emb.as_deref_mut(), &self.markov) {
let mut f = e.uninit(mk.rank)?;
e.gather_row_bf16(&mk.w1_bf16, &chain_d, k, &mut f, mk.rank)?;
let fv = e.view(&f, mk.rank);
e.copy_view_into(ce, k * mk.rank, &fv, mk.rank)?;
}
let rows_k = e.htod_i32(&[k as i32])?;
let (mut th1, mut z1, mut mx1) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
dl, n_vocab, &rows_k, &mut th1, &mut z1, &mut mx1, n_vocab, 1, sp.temp, sp.top_k,
sp.top_p, sp.min_p,
)?;
e.gumbel_perturb_filtered_col(
dl, k, &mut pb, n_vocab, sp.seed, *sctr, sp.temp, &mx1, &th1, 0,
)?;
*sctr = sctr.wrapping_add(1);
e.argmax_token_device_col(&pb, 0, n_vocab, &mut chain_d, k + 1)?;
e.copy_into(&mut th_all, k, &th1, 1)?;
e.copy_into(&mut z_all, k, &z1, 1)?;
e.copy_into(&mut mx_all, k, &mx1, 1)?;
}
let chain = e.dtoh_u32(&chain_d)?;
let (thv, zv, mxv) = (e.dtoh(&th_all)?, e.dtoh(&z_all)?, e.dtoh(&mx_all)?);
let stats = (0..nd).map(|i| (mxv[i], thv[i], zv[i])).collect();
Ok((
chain[1..].to_vec(),
DsparkDraftSample::Rows {
th: th_all,
z: z_all,
stats,
},
))
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dspark_propose_sampled(
&self,
e: &Engine,
dl: &mut CudaSlice<f32>,
rows: &CudaSlice<f32>,
nd: usize,
n_vocab: usize,
anchor: u32,
sp: &crate::spec::SpecSampling,
sctr: &mut u32,
uctr: &mut u32,
conf_emb: Option<&mut CudaSlice<f32>>,
) -> Result<(Vec<u32>, DsparkDraftSample), Box<dyn std::error::Error>> {
if let Some(d2) = self.dflash2.as_ref() {
debug_assert!(
conf_emb.is_none(),
"conf_emb stash requested on a DFlash2 selector proposal"
);
let (path, q_chosen, cand, q_rows) = self.dflash2_propose_sampled(
e, dl, rows, nd, n_vocab, anchor, sp.temp, sp.seed, uctr,
)?;
Ok((
path,
DsparkDraftSample::Selector {
cand,
q_rows,
q_chosen,
top_k: d2.top_k,
},
))
} else {
self.dspark_chain_sampled(e, dl, nd, n_vocab, anchor, sp, sctr, conf_emb)
}
}
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 mut xn = e.uninit(b * h)?;
e.rms_norm(&x, &l.ln_in, &mut xn, h, b, c.eps)?;
let mut attn_dyn: Option<CudaSlice<f32>> = None;
if let Some(d2) = &self.dflash2 {
let (xc, dyn_) = self.d2_conv_prepare(e, &d2.attn_conv[li], &xn, b)?;
xn = xc;
attn_dyn = Some(dyn_);
}
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 self.dflash2.is_some() && c.layer_sliding[li] {
debug_assert!(pos.windows(2).all(|w| w[1] == w[0] + 1));
d2_windowed_attn(e, &q, &k, &v, &mut attn, hd, nh, nkv, b, ctx + b, scale, c)?;
} else 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 mut o = self.mm(e, &l.wo, &attn, b, nh * hd, h)?;
if let (Some(d2), Some(dyn_)) = (&self.dflash2, &attn_dyn) {
o = self.d2_conv_finish(e, &d2.attn_conv[li], &o, dyn_, b)?;
}
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 mut mlp_dyn: Option<CudaSlice<f32>> = None;
if let Some(d2) = &self.dflash2 {
let (xc, dyn_) = self.d2_conv_prepare(e, &d2.mlp_conv[li], &x1n, b)?;
x1n = xc;
mlp_dyn = Some(dyn_);
}
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 mut down = self.mm(e, &l.w_down, &act, b, c.n_ff, h)?;
if let (Some(d2), Some(dyn_)) = (&self.dflash2, &mlp_dyn) {
down = self.d2_conv_finish(e, &d2.mlp_conv[li], &down, dyn_, b)?;
}
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 mut attn_dyn: Option<CudaSlice<f32>> = None;
if let Some(d2) = &self.dflash2 {
let (xc, dyn_) = self.d2_conv_prepare(e, &d2.attn_conv[li], &xn, b)?;
xn = xc;
attn_dyn = Some(dyn_);
}
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 self.dflash2.is_some() && c.layer_sliding[li] {
d2_windowed_attn(
e,
&q,
&kv.k[li],
&kv.v[li],
&mut attn,
hd,
nh,
nkv,
b,
ctx + b,
scale,
c,
)?;
} else 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 mut o = self.mm(e, &l.wo, &attn, b, nh * hd, h)?;
if let (Some(d2), Some(dyn_)) = (&self.dflash2, &attn_dyn) {
o = self.d2_conv_finish(e, &d2.attn_conv[li], &o, dyn_, b)?;
}
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 mut mlp_dyn: Option<CudaSlice<f32>> = None;
if let Some(d2) = &self.dflash2 {
let (xc, dyn_) = self.d2_conv_prepare(e, &d2.mlp_conv[li], &x1n, b)?;
x1n = xc;
mlp_dyn = Some(dyn_);
}
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 mut down = self.mm(e, &l.w_down, &act, b, c.n_ff, h)?;
if let (Some(d2), Some(dyn_)) = (&self.dflash2, &mlp_dyn) {
down = self.d2_conv_finish(e, &d2.mlp_conv[li], &down, dyn_, b)?;
}
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!(
draft.dflash2.is_none(),
"DFlash2 drafters ride the qwen-hybrid dspark round (selector + windowed \
attention); the gemma arm has no consumer for the family's ops"
);
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],
sampling: Option<&crate::spec::SpecSampling>,
) -> 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 sp_on: Option<&crate::spec::SpecSampling> = sampling.filter(|s| s.temp > 0.0);
let (mut sctr, mut uctr) = (0u32, 0u32);
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!(
draft.dflash2.is_some() || 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 = match sp_on {
Some(sp) => {
crate::spec::sample_boundary_token(e, &logits, sp, &[], &mut sctr, "dspark-prime")?
}
None => 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::for_draft(draft);
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 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,
};
let mut cand: Vec<u32> = Vec::with_capacity(nd + 1);
let mut prop: Option<DsparkDraftSample> = None;
let mut chain_dev: Option<CudaSlice<u32>> = None;
if let Some(sp) = sp_on {
let (tail, ds) = draft.dspark_propose_sampled(
e,
&mut dl,
&rows,
nd,
n_vocab,
last,
sp,
&mut sctr,
&mut uctr,
conf_emb.as_mut(),
)?;
e.set_verify_exact(false);
cand.push(last);
cand.extend_from_slice(&tail);
prop = Some(ds);
} else if draft.dflash2.is_some() {
let path = draft.dflash2_propose_greedy(e, &dl, &rows, nd, n_vocab, last)?;
e.set_verify_exact(false);
cand.push(last);
cand.extend_from_slice(&path);
} else {
let markov_on = std::env::var("MEMRA_DFLASH_MARKOV").as_deref() != Ok("0");
let mut chain_d = e.stream().alloc_zeros::<u32>(nd + 1)?;
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);
chain_dev = Some(chain_d);
}
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");
if sp_on.is_some() && ckpt_gate {
return Err(
"MEMRA_DSPARK_CKPT_GATE compares verify argmaxes across a replay \
— a greedy-exactness instrument; unset it for T>0 dspark rounds"
.into(),
);
}
let deferred = chain_dev.is_some() && 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");
}
if let Some(chain_d) = chain_dev.as_ref() {
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<CudaSlice<f32>>,
Option<crate::spec::DsparkVerifyCkpt>,
),
Box<dyn std::error::Error>,
> {
if sp_on.is_some() {
if ckpt_on {
let (tl, vck) =
self.dspark_verify_t_logits_ckpt(e, &cand[..vt], start, cache)?;
Ok((Vec::new(), Some(tl), Some(vck)))
} else {
Ok((
Vec::new(),
Some(self.dspark_verify_t_logits(e, &cand[..vt], start, cache)?),
None,
))
}
} else if deferred {
let chain_d = chain_dev.as_ref().expect("deferred implies greedy chain");
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, None, Some(vck)))
} else if ckpt_on || ckpt_gate {
let (vam, vck) = self.dspark_verify_t_am_ckpt(e, &cand[..vt], start, cache)?;
Ok((vam, None, Some(vck)))
} else {
Ok((
self.dspark_verify_t_am(e, &cand[..vt], start, cache)?,
None,
None,
))
}
})(&mut cache, &mut cand, vgraphs);
let (vam, tl, 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, next) = match (sp_on, tl.as_ref()) {
(Some(sp), Some(tl)) => dspark_accept_sampled(
e,
tl,
&cand,
vt,
n_vocab,
&dl,
prop.as_ref()
.expect("sampled round without a proposal record"),
sp,
&mut sctr,
&mut uctr,
)?,
_ => {
let m = dspark_accept_prefix(&cand, &vam, vt);
(m, vam[m])
}
};
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 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)?;
if sp_on.is_none() {
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,
sampling: Option<crate::spec::SpecSampling>,
sctr: u32,
uctr: u32,
}
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,
sampling: Option<crate::spec::SpecSampling>,
) -> 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"
);
if let Some(sp) = sampling.as_ref().filter(|s| s.temp > 0.0) {
if sp.penalty_repeat != 1.0 || sp.penalty_freq != 0.0 || sp.penalty_present != 0.0 {
return Err(
"dspark sampled admission excludes penalties (admission gate \
keeps penalized requests on the plain path)"
.into(),
);
}
}
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 = if draft.dflash2.is_some() {
ctx_cap
} else {
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 mut sctr0 = 0u32;
let last = match sampling.as_ref().filter(|s| s.temp > 0.0) {
Some(sp) => {
crate::spec::sample_boundary_token(e, &logits, sp, &[], &mut sctr0, "dspark-prime")?
}
None => 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::for_draft(draft).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(),
sampling,
sctr: sctr0,
uctr: 0,
})
}
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::for_draft(draft);
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 sp_on: Option<crate::spec::SpecSampling> = sess.sampling.filter(|s| s.temp > 0.0);
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 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,
};
let mut cand: Vec<u32> = Vec::with_capacity(nd + 1);
let mut prop: Option<DsparkDraftSample> = None;
let mut chain_dev: Option<CudaSlice<u32>> = None;
let ckpt_on = std::env::var("MEMRA_DSPARK_CKPT").as_deref() != Ok("0");
let mut deferred = false;
if let Some(sp) = sp_on.as_ref() {
let (tail, ds) = draft.dspark_propose_sampled(
e,
&mut dl,
&rows,
nd,
n_vocab,
sess.last,
sp,
&mut sess.sctr,
&mut sess.uctr,
conf_emb.as_mut(),
)?;
e.set_verify_exact(false);
cand.push(sess.last);
cand.extend_from_slice(&tail);
prop = Some(ds);
} else if draft.dflash2.is_some() {
let path = draft.dflash2_propose_greedy(e, &dl, &rows, nd, n_vocab, sess.last)?;
e.set_verify_exact(false);
cand.push(sess.last);
cand.extend_from_slice(&path);
} else {
let markov_on = std::env::var("MEMRA_DFLASH_MARKOV").as_deref() != Ok("0");
let mut chain_d = e.stream().alloc_zeros::<u32>(nd + 1)?;
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);
deferred = embd_gpu.is_some() && ckpt_on;
chain_dev = Some(chain_d);
}
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");
}
if let Some(chain_d) = chain_dev.as_ref() {
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, tl, vck) = if sp_on.is_some() {
if ckpt_on {
let (tl, vck) =
self.dspark_verify_t_logits_ckpt(e, &cand[..vt], start, &mut sess.cache)?;
(Vec::new(), Some(tl), Some(vck))
} else {
(
Vec::new(),
Some(self.dspark_verify_t_logits(
e,
&cand[..vt],
start,
&mut sess.cache,
)?),
None,
)
}
} else if deferred {
let chain_d = chain_dev.as_ref().expect("deferred implies greedy chain");
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, None, Some(vck))
} else if ckpt_on {
let (vam, vck) =
self.dspark_verify_t_am_ckpt(e, &cand[..vt], start, &mut sess.cache)?;
(vam, None, Some(vck))
} else {
(
self.dspark_verify_t_am(e, &cand[..vt], start, &mut sess.cache)?,
None,
None,
)
};
let taps = sess.cache.dflash_taps.take().unwrap();
let (m, next) = match (sp_on.as_ref(), tl.as_ref()) {
(Some(sp), Some(tl)) => dspark_accept_sampled(
e,
tl,
&cand,
vt,
n_vocab,
&dl,
prop.as_ref()
.expect("sampled round without a proposal record"),
sp,
&mut sess.sctr,
&mut sess.uctr,
)?,
_ => {
let m = dspark_accept_prefix(&cand, &vam, vt);
(m, vam[m])
}
};
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 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)?;
if sp_on.is_none() {
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 dflash2_tests {
use super::{DsparkHarvest, dflash2_walk_greedy, dflash2_walk_sampled, rejection_accept_len};
fn bf16(vals: &[f32]) -> Vec<u8> {
vals.iter()
.flat_map(|v| ((v.to_bits() >> 16) as u16).to_le_bytes())
.collect()
}
const V: usize = 8; const R: usize = 2; const K: usize = 2;
fn books() -> (Vec<u8>, Vec<u8>) {
let mut pred = vec![0f32; V * R];
pred[0] = 1.0; pred[1 * R + 1] = 1.0; pred[2 * R] = 1.0; let mut succ = vec![0f32; V * R];
succ[1 * R] = 2.0; succ[2 * R + 1] = 5.0; succ[3 * R + 1] = 3.0; succ[4 * R] = 10.0; (bf16(&pred), bf16(&succ))
}
#[test]
fn selector_walk_is_a_chain_not_per_slot_argmax() {
let (pred, succ) = books();
let cand: Vec<u32> = vec![1, 2, 3, 4];
let hproj = vec![1.0f32; 2 * R];
let path = dflash2_walk_greedy(&pred, &succ, V, R, K, &[0.0; 4], &cand, &hproj, 0, 2);
assert_eq!(
path,
vec![1, 3],
"walk must seed slot p from slot p-1's CHOSEN candidate \
(z-lab model.py CandidateSelector.select)"
);
}
#[test]
fn selector_walk_unary_term_participates() {
let (pred, succ) = books();
let cand: Vec<u32> = vec![1, 2, 3, 4];
let hproj = vec![1.0f32; 2 * R];
let path = dflash2_walk_greedy(
&pred,
&succ,
V,
R,
K,
&[0.0, 10.0, 0.0, 0.0],
&cand,
&hproj,
0,
2,
);
assert_eq!(
path,
vec![2, 4],
"score = unary + bilinear (reference: `unary[:, position] + einsum(...)`); \
dropping the unary term picks tok 1 here"
);
}
#[test]
fn selector_walk_hidden_gate_participates() {
let (pred, succ) = books();
let cand: Vec<u32> = vec![1, 2, 3, 4];
let hproj = vec![0.0f32, 1.0, 1.0, 1.0];
let path = dflash2_walk_greedy(
&pred,
&succ,
V,
R,
K,
&[0.0, 1.0, 0.0, 0.0],
&cand,
&hproj,
0,
2,
);
assert_eq!(
path,
vec![2, 4],
"the bilinear gate is pred_row .* HIDDEN_PROJECTION (reference: \
`predecessor_codebook(predecessor) * hidden[:, position]`); ignoring \
hproj leaves tok1's margin standing"
);
}
#[test]
fn dflash2_harvest_is_census_keyed() {
assert_eq!(
DsparkHarvest::for_family_value(true, None, false),
DsparkHarvest::Dflash
);
assert_eq!(
DsparkHarvest::for_family_value(true, Some("dflash"), false),
DsparkHarvest::Dflash
);
assert_eq!(
DsparkHarvest::for_family_value(true, None, true),
DsparkHarvest::Dflash
);
assert!(
std::panic::catch_unwind(|| DsparkHarvest::for_family_value(
true,
Some("dspark"),
false
))
.is_err(),
"MEMRA_DSPARK_HARVEST=dspark on a DFlash2 checkpoint must refuse"
);
assert_eq!(
DsparkHarvest::for_family_value(false, Some("dspark"), false),
DsparkHarvest::Dspark
);
assert_eq!(
DsparkHarvest::for_family_value(false, None, false),
DsparkHarvest::Dflash
);
assert_eq!(
DsparkHarvest::for_family_value(false, None, true),
DsparkHarvest::Dspark,
"unset env on a DSPARK-strategy export must keep the ratified census flip"
);
}
#[test]
fn sampled_walk_tiny_temp_matches_greedy() {
let (pred, succ) = books();
let cand: Vec<u32> = vec![1, 2, 3, 4];
let hproj = vec![1.0f32; 2 * R];
let greedy = dflash2_walk_greedy(&pred, &succ, V, R, K, &[0.0; 4], &cand, &hproj, 0, 2);
let mut u = || 0.5f32;
let (path, q_chosen, q_rows) = dflash2_walk_sampled(
&pred, &succ, V, R, K, &[0.0; 4], &cand, &hproj, 0, 2, 1e-6, &mut u,
);
assert_eq!(
path, greedy,
"tiny-T sampled walk must equal the greedy chain"
);
assert_eq!(q_rows.len(), 2 * K);
for (p, &q) in path.iter().zip(&q_chosen) {
let _ = p;
assert!(
q > 0.999,
"tiny-T chosen-candidate prob must be ~1, got {q}"
);
}
}
#[test]
fn sampled_walk_records_the_distribution_it_samples() {
let (pred, succ) = books();
let cand: Vec<u32> = vec![1, 2, 3, 4];
let hproj = vec![1.0f32; 2 * R];
let q1 = (1f64.exp() / (1f64.exp() + 1.0)) as f32;
for (u0, want0) in [(q1 - 0.01, 1u32), (q1 + 0.01, 2u32)] {
let mut seq = vec![u0, 0.0f32].into_iter();
let mut u = move || seq.next().unwrap();
let (path, q_chosen, q_rows) = dflash2_walk_sampled(
&pred, &succ, V, R, K, &[0.0; 4], &cand, &hproj, 0, 2, 2.0, &mut u,
);
assert_eq!(
path[0], want0,
"CDF walk must place u={u0} in the right candidate bracket"
);
let row0: f32 = q_rows[..K].iter().sum();
assert!(
(row0 - 1.0).abs() < 1e-5,
"slot-0 q must sum to 1, got {row0}"
);
let ci = cand[..K].iter().position(|&c| c == path[0]).unwrap();
assert_eq!(
q_chosen[0], q_rows[ci],
"q_chosen must be the recorded row prob of the drawn candidate"
);
assert!(
(q_rows[0] - q1).abs() < 1e-4,
"slot-0 tok1 prob must be softmax(scores/T), got {} want {q1}",
q_rows[0]
);
}
}
#[test]
fn sampled_walk_chains_the_drawn_candidate() {
let (pred, succ) = books();
let cand: Vec<u32> = vec![1, 2, 3, 4];
let hproj = vec![1.0f32; 2 * R];
let mut seq = vec![0.99f32, 0.01].into_iter();
let mut u = move || seq.next().unwrap();
let (path, _, _) = dflash2_walk_sampled(
&pred, &succ, V, R, K, &[0.0; 4], &cand, &hproj, 0, 2, 2.0, &mut u,
);
assert_eq!(path[0], 2, "u=0.99 must draw the low-prob candidate");
assert_eq!(
path[1], 4,
"slot 1 must walk from pred[2] (the DRAWN token), which scores tok4 at 10 \
— chaining from the anchor or the argmax picks tok3"
);
}
#[test]
fn rejection_accept_walk_is_the_leviathan_rule() {
assert_eq!(
rejection_accept_len(&[0.5, 0.5], &[0.5, 0.5], &[0.9, 0.9]),
2
);
assert_eq!(
rejection_accept_len(&[0.5, 0.5], &[0.5, 0.5], &[1.0, 0.0]),
0
);
assert_eq!(rejection_accept_len(&[0.25], &[0.5], &[0.5]), 0);
assert_eq!(rejection_accept_len(&[1e-6], &[0.0], &[0.999]), 1);
assert_eq!(
rejection_accept_len(&[0.9, 0.0, 0.9], &[0.1, 0.9, 0.1], &[0.5, 0.5, 0.5]),
1
);
}
fn tv(a: &[f64], b: &[f64]) -> f64 {
a.iter().zip(b).map(|(x, y)| (x - y).abs()).sum::<f64>() / 2.0
}
fn compose_once(
p: &[f32],
q: &[f32],
u_draw: f32,
u_accept: f32,
u_resid: f32,
invert_accept: bool,
skip_q_in_residual: bool,
) -> usize {
let n = p.len();
let mut acc = 0f64;
let mut x = n - 1;
for (i, &qi) in q.iter().enumerate() {
acc += qi as f64;
if (u_draw as f64) < acc {
x = i;
break;
}
}
let accepted = if invert_accept {
!((u_accept as f64) * (q[x] as f64) < p[x] as f64)
} else {
rejection_accept_len(&p[x..=x], &q[x..=x], &[u_accept]) == 1
};
if accepted {
return x;
}
let r: Vec<f64> = p
.iter()
.zip(q)
.map(|(&pi, &qi)| {
let qq = if skip_q_in_residual { 0.0 } else { qi as f64 };
(pi as f64 - qq).max(0.0)
})
.collect();
let total: f64 = r.iter().sum();
let mut acc = 0f64;
let target = u_resid as f64 * total;
for (i, &ri) in r.iter().enumerate() {
acc += ri;
if acc >= target && ri > 0.0 {
return i;
}
}
n - 1
}
fn compose_tv(q: &[f32], invert_accept: bool, skip_q_in_residual: bool) -> f64 {
let p: Vec<f32> = vec![0.30, 0.22, 0.15, 0.12, 0.09, 0.06, 0.04, 0.02];
let trials = 200_000usize;
let mut counts = vec![0f64; V];
for t in 0..trials {
let u_draw = crate::spec::host_u01(7, (t * 3) as u32);
let u_accept = crate::spec::host_u01(7, (t * 3 + 1) as u32);
let u_resid = crate::spec::host_u01(7, (t * 3 + 2) as u32);
counts[compose_once(
&p,
q,
u_draw,
u_accept,
u_resid,
invert_accept,
skip_q_in_residual,
)] += 1.0;
}
let emp: Vec<f64> = counts.iter().map(|c| c / trials as f64).collect();
let pf: Vec<f64> = p.iter().map(|&v| v as f64).collect();
tv(&emp, &pf)
}
#[test]
fn sampled_round_composition_matches_the_target() {
let q_rows: Vec<f32> = vec![0.02, 0.04, 0.06, 0.09, 0.12, 0.15, 0.22, 0.30];
let q_sparse: Vec<f32> = vec![0.0, 0.7, 0.0, 0.3, 0.0, 0.0, 0.0, 0.0];
for (name, q) in [("rows", &q_rows), ("sparse", &q_sparse)] {
let d = compose_tv(q, false, false);
assert!(
d < 0.01,
"composition[{name}]: committed-token distribution must equal p \
(TV {d:.4} >= 0.01)"
);
}
}
#[test]
fn composition_teeth_inverted_accept_fails() {
let q: Vec<f32> = vec![0.02, 0.04, 0.06, 0.09, 0.12, 0.15, 0.22, 0.30];
let d = compose_tv(&q, true, false);
assert!(
d > 0.05,
"inverted accept rule must fail the composition bound (TV {d:.4})"
);
}
#[test]
fn composition_teeth_residual_without_q_fails() {
let q: Vec<f32> = vec![0.0, 0.7, 0.0, 0.3, 0.0, 0.0, 0.0, 0.0];
let d = compose_tv(&q, false, true);
assert!(
d > 0.05,
"residual that skips the q subtraction must fail the bound (TV {d:.4})"
);
}
}
#[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);
}
}