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 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 struct MarkovHead {
pub w1_bf16: CudaSlice<u8>, pub w2: GpuTensor, pub rank: usize,
pub vocab: usize,
}
fn bf16_to_f32(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(2)
.map(|c| f32::from_bits((u16::from_le_bytes([c[0], c[1]]) as u32) << 16))
.collect()
}
fn encode_q8_0(vals: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(vals.len() / 32 * 34);
for blk in vals.chunks_exact(32) {
let amax = blk.iter().fold(0f32, |a, v| a.max(v.abs()));
let d = amax / 127.0;
let id = if d > 0.0 { 1.0 / d } else { 0.0 };
let dh = half_from_f32(d);
out.extend_from_slice(&dh.to_le_bytes());
for &v in blk {
out.push(((v * id).round().clamp(-127.0, 127.0)) as i8 as u8);
}
}
out
}
fn encode_q4_0(vals: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(vals.len() / 32 * 18);
for blk in vals.chunks_exact(32) {
let mut amax = 0f32; let mut mx = 0f32;
for &v in blk { if v.abs() > amax { amax = v.abs(); mx = v; } }
let d = mx / -8.0;
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
out.extend_from_slice(&half_from_f32(d).to_le_bytes());
for j in 0..16 {
let x0 = (blk[j] * id + 8.5).clamp(0.0, 15.0) as u8;
let x1 = (blk[j + 16] * id + 8.5).clamp(0.0, 15.0) as u8;
out.push(x0 | (x1 << 4));
}
}
out
}
fn half_from_f32(v: f32) -> u16 {
let b = v.to_bits();
let sign = ((b >> 16) & 0x8000) as u16;
let exp = ((b >> 23) & 0xff) as i32 - 127 + 15;
let man = b & 0x7fffff;
if exp <= 0 { return sign; } if exp >= 31 { return sign | 0x7c00; } let mut h = sign | ((exp as u16) << 10) | ((man >> 13) as u16);
let rem = man & 0x1fff;
if rem > 0x1000 || (rem == 0x1000 && (h & 1) == 1) { h += 1; }
h
}
impl DflashDraft {
pub fn load(e: &Engine, dir: &std::path::Path) -> Result<Self, Box<dyn std::error::Error>> {
let txt = std::fs::read_to_string(dir.join("config.json"))?;
fn num(txt: &str, key: &str) -> Option<f64> {
let i = txt.find(&format!("\"{key}\""))?;
let rest = &txt[i..];
let colon = rest.find(':')?;
let val: String = rest[colon + 1..].trim_start().chars()
.take_while(|c| c.is_ascii_digit() || *c == '.' || *c == '-' || *c == 'e' || *c == 'E' || *c == '+')
.collect();
val.parse().ok()
}
fn num_list(txt: &str, key: &str) -> Vec<usize> {
let Some(i) = txt.find(&format!("\"{key}\"")) else { return Vec::new() };
let rest = &txt[i..];
let (Some(a), Some(b)) = (rest.find('['), rest.find(']')) else { return Vec::new() };
rest[a + 1..b].split(',').filter_map(|s| s.trim().parse().ok()).collect()
}
let g = |k: &str| num(&txt, k).unwrap_or_else(|| panic!("config missing {k}")) as usize;
let layer_sliding: Vec<bool> = {
let i = txt.find("\"layer_types\"").expect("layer_types");
let rest = &txt[i..];
let (a, b) = (rest.find('[').unwrap(), rest.find(']').unwrap());
rest[a + 1..b].split(',').map(|s| s.contains("sliding_attention")).collect()
};
let cfg = DflashCfg {
hidden: g("hidden_size"),
n_head: g("num_attention_heads"),
n_kv: g("num_key_value_heads"),
head_dim: g("head_dim"),
n_ff: g("intermediate_size"),
n_layer: g("num_hidden_layers"),
eps: num(&txt, "rms_norm_eps").expect("rms_norm_eps") as f32,
rope_theta: num(&txt, "rope_theta").expect("rope_theta") as f32,
block_size: g("block_size"),
mask_token_id: g("mask_token_id") as u32,
target_layer_ids: num_list(&txt, "target_layer_ids"),
sliding_window: g("sliding_window"),
layer_sliding,
};
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, 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, 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 w2f = bf16_to_f32(b2);
let w2q = encode_q8_0(&w2f);
Some(MarkovHead {
w1_bf16: e.upload_u8(bytes)?,
w2: 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, rp4: None, f16: None,
},
rank, vocab,
})
} else { None };
Ok(Self {
fc: upw("fc.weight")?,
hidden_norm: up("hidden_norm.weight")?,
norm: up("norm.weight")?,
cfg,
layers,
markov,
})
}
fn mm(&self, e: &Engine, w: &GpuTensor, x: &CudaSlice<f32>, t: usize, _in_f: usize,
_out_f: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
Ok(e.matmul(w, x, t)?)
}
pub fn ctx_features(&self, e: &Engine, taps: &CudaSlice<f32>, t: usize)
-> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let c = &self.cfg;
let n_taps = c.target_layer_ids.len();
let fc_out = self.mm(e, &self.fc, taps, t, n_taps * c.hidden, c.hidden)?;
let mut out = e.uninit(t * c.hidden)?;
e.rms_norm(&fc_out, &self.hidden_norm, &mut out, c.hidden, t, c.eps)?;
Ok(out)
}
pub fn forward(
&self, e: &Engine, target_hidden: &CudaSlice<f32>, noise_emb: &CudaSlice<f32>,
pos: &[i32], ctx: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let ctx_f = self.ctx_features(e, target_hidden, ctx)?;
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&ctx_f)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-ctx_features.f32"), bytes)?;
}
self.forward_block(e, &ctx_f, noise_emb, pos, ctx)
}
pub fn forward_block(
&self, e: &Engine, ctx_f: &CudaSlice<f32>, noise_emb: &CudaSlice<f32>,
pos: &[i32], ctx: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let c = &self.cfg;
let (h, nh, nkv, hd) = (c.hidden, c.n_head, c.n_kv, c.head_dim);
let b = c.block_size;
assert_eq!(pos.len(), ctx + b, "pos covers ctx rows then block rows");
let pos_blk = e.htod_i32(&pos[ctx..])?;
let mut x = e.clone_dtod(noise_emb)?; for (li, l) in self.layers.iter().enumerate() {
let _ = li;
let mut xn = e.uninit(b * h)?;
e.rms_norm(&x, &l.ln_in, &mut xn, h, b, c.eps)?;
let q0 = self.mm(e, &l.wq, &xn, b, h, nh * hd)?;
let k0c = self.mm(e, &l.wk, ctx_f, ctx, h, nkv * hd)?;
let v0c = self.mm(e, &l.wv, ctx_f, ctx, h, nkv * hd)?;
let k0b = self.mm(e, &l.wk, &xn, b, h, nkv * hd)?;
let v0b = self.mm(e, &l.wv, &xn, b, h, nkv * hd)?;
let mut k0 = e.uninit((ctx + b) * nkv * hd)?;
e.copy_into(&mut k0, 0, &k0c, ctx * nkv * hd)?;
e.copy_into(&mut k0, ctx * nkv * hd, &k0b, b * nkv * hd)?;
let mut v = e.uninit((ctx + b) * nkv * hd)?;
e.copy_into(&mut v, 0, &v0c, ctx * nkv * hd)?;
e.copy_into(&mut v, ctx * nkv * hd, &v0b, b * nkv * hd)?;
if li == 0 { if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&q0)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_q0.f32"), bytes)?;
}}
let mut q = e.uninit(b * nh * hd)?;
let mut k = e.uninit((ctx + b) * nkv * hd)?;
e.rms_norm(&q0, &l.q_norm, &mut q, hd, b * nh, c.eps)?;
if li == 0 { if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&q)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_qn.f32"), bytes)?;
}}
e.rms_norm(&k0, &l.k_norm, &mut k, hd, (ctx + b) * nkv, c.eps)?;
let norope = std::env::var("MEMRA_DFLASH_NOROPE").is_ok();
if !norope { e.rope_neox(&mut q, &pos_blk, hd, hd, nh, b, c.rope_theta, 1.0)?; }
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 { e.rope_neox(&mut k, &pos_all, hd, hd, nkv, ctx + b, c.rope_theta, 1.0)?; }
let mut attn = e.uninit(b * nh * hd)?;
let scale = 1.0f32 / (hd as f32).sqrt();
if std::env::var("MEMRA_DFLASH_FA").is_ok() {
e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, b, ctx + b, scale, false)?;
} else {
e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, b, ctx + b, scale, false)?;
}
let o = self.mm(e, &l.wo, &attn, b, nh * hd, h)?;
let mut x1 = e.uninit(b * h)?;
e.add(&o, &x, &mut x1, b * h)?;
if li == 0 { if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let dump = |name: &str, t: &cudarc::driver::CudaSlice<f32>| -> Result<(), Box<dyn std::error::Error>> {
let v = e.dtoh(t)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-l0_{name}.f32"), bytes)?;
Ok(())
};
dump("q", &q)?; dump("k", &k)?; dump("attn", &attn)?; dump("x1", &x1)?;
}}
let mut x1n = e.uninit(b * h)?;
e.rms_norm(&x1, &l.ln_post, &mut x1n, h, b, c.eps)?;
let gate = self.mm(e, &l.w_gate, &x1n, b, h, c.n_ff)?;
let up_ = self.mm(e, &l.w_up, &x1n, b, h, c.n_ff)?;
let mut act = e.uninit(b * c.n_ff)?;
e.silu_mul(&gate, &up_, &mut act, b * c.n_ff)?;
let down = self.mm(e, &l.w_down, &act, b, c.n_ff, h)?;
let mut x2 = e.uninit(b * h)?;
e.add(&down, &x1, &mut x2, b * h)?;
x = x2;
if let Ok(dir) = std::env::var("MEMRA_DFLASH_DUMP") {
let v = e.dtoh(&x)?;
let bytes: Vec<u8> = v.iter().flat_map(|f| f.to_le_bytes()).collect();
std::fs::write(format!("{dir}/memra-layer{li}_out.f32"), bytes)?;
}
}
let mut out = e.uninit(b * h)?;
e.rms_norm(&x, &self.norm, &mut out, h, b, c.eps)?;
Ok(out)
}
}
pub struct DflashKv {
pub k: Vec<CudaSlice<f32>>, pub v: Vec<CudaSlice<f32>>,
pub len: usize,
pub cap: usize,
}
impl DflashKv {
pub fn new(e: &Engine, cfg: &DflashCfg, cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
let rowsz = cfg.n_kv * cfg.head_dim;
let mut k = Vec::with_capacity(cfg.n_layer);
let mut v = Vec::with_capacity(cfg.n_layer);
for _ in 0..cfg.n_layer {
k.push(e.uninit((cap + cfg.block_size) * rowsz)?);
v.push(e.uninit((cap + cfg.block_size) * rowsz)?);
}
Ok(Self { k, v, len: 0, cap })
}
}
impl DflashDraft {
pub fn ingest_ctx(&self, e: &Engine, kv: &mut DflashKv, feats: &CudaSlice<f32>,
pos_new: &[i32], t: usize) -> Result<(), Box<dyn std::error::Error>> {
let c = &self.cfg;
let (h, nkv, hd) = (c.hidden, c.n_kv, c.head_dim);
assert!(kv.len + t <= kv.cap, "draft kv overflow");
let pos_d = e.htod_i32(pos_new)?;
for (li, l) in self.layers.iter().enumerate() {
let k0 = self.mm(e, &l.wk, feats, t, h, nkv * hd)?;
let v0 = self.mm(e, &l.wv, feats, t, h, nkv * hd)?;
let mut kn = e.uninit(t * nkv * hd)?;
e.rms_norm(&k0, &l.k_norm, &mut kn, hd, t * nkv, c.eps)?;
e.rope_neox(&mut kn, &pos_d, hd, hd, nkv, t, c.rope_theta, 1.0)?;
e.copy_into(&mut kv.k[li], kv.len * nkv * hd, &kn, t * nkv * hd)?;
e.copy_into(&mut kv.v[li], kv.len * nkv * hd, &v0, t * nkv * hd)?;
}
kv.len += t;
Ok(())
}
pub fn forward_round(&self, e: &Engine, kv: &mut DflashKv, noise_emb: &CudaSlice<f32>,
pos_block: &[i32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let c = &self.cfg;
let (h, nh, nkv, hd) = (c.hidden, c.n_head, c.n_kv, c.head_dim);
let b = c.block_size;
assert_eq!(pos_block.len(), b);
let ctx = kv.len;
let pos_blk = e.htod_i32(pos_block)?;
let mut x = e.clone_dtod(noise_emb)?;
for (li, l) in self.layers.iter().enumerate() {
let mut xn = e.uninit(b * h)?;
e.rms_norm(&x, &l.ln_in, &mut xn, h, b, c.eps)?;
let q0 = self.mm(e, &l.wq, &xn, b, h, nh * hd)?;
let k0b = self.mm(e, &l.wk, &xn, b, h, nkv * hd)?;
let v0b = self.mm(e, &l.wv, &xn, b, h, nkv * hd)?;
let mut q = e.uninit(b * nh * hd)?;
let mut kb = e.uninit(b * nkv * hd)?;
e.rms_norm(&q0, &l.q_norm, &mut q, hd, b * nh, c.eps)?;
e.rms_norm(&k0b, &l.k_norm, &mut kb, hd, b * nkv, c.eps)?;
e.rope_neox(&mut q, &pos_blk, hd, hd, nh, b, c.rope_theta, 1.0)?;
e.rope_neox(&mut kb, &pos_blk, hd, hd, nkv, b, c.rope_theta, 1.0)?;
e.copy_into(&mut kv.k[li], ctx * nkv * hd, &kb, b * nkv * hd)?;
e.copy_into(&mut kv.v[li], ctx * nkv * hd, &v0b, b * nkv * hd)?;
let mut attn = e.uninit(b * nh * hd)?;
let scale = 1.0f32 / (hd as f32).sqrt();
if std::env::var("MEMRA_DFLASH_FA").is_ok() {
e.fa_prefill(&q, &kv.k[li], &kv.v[li], &mut attn, hd, nh, nkv, b, ctx + b,
scale, false)?;
} else {
e.sdpa_naive(&q, &kv.k[li], &kv.v[li], &mut attn, hd, nh, nkv, b, ctx + b,
scale, false)?;
}
let o = self.mm(e, &l.wo, &attn, b, nh * hd, h)?;
let mut x1 = e.uninit(b * h)?;
e.add(&o, &x, &mut x1, b * h)?;
let mut x1n = e.uninit(b * h)?;
e.rms_norm(&x1, &l.ln_post, &mut x1n, h, b, c.eps)?;
let gate = self.mm(e, &l.w_gate, &x1n, b, h, c.n_ff)?;
let up_ = self.mm(e, &l.w_up, &x1n, b, h, c.n_ff)?;
let mut act = e.uninit(b * c.n_ff)?;
e.silu_mul(&gate, &up_, &mut act, b * c.n_ff)?;
let down = self.mm(e, &l.w_down, &act, b, c.n_ff, h)?;
let mut x2 = e.uninit(b * h)?;
e.add(&down, &x1, &mut x2, b * h)?;
x = x2;
}
let mut out = e.uninit(b * h)?;
e.rms_norm(&x, &self.norm, &mut out, h, b, c.eps)?;
Ok(out)
}
}
impl crate::hybrid::HybridModel {
pub fn generate_spec_dflash(
&self, e: &Engine, draft: &DflashDraft, prompt: &[u32], max_new: usize, eos: &[u32],
) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
use crate::cache::{Cache, DflashTapSink};
let n_embd = self.cfg.n_embd as usize;
let c = &draft.cfg;
assert_eq!(n_embd, c.hidden, "draft hidden must match target n_embd");
let b = c.block_size;
let n_taps = c.target_layer_ids.len();
let max_ctx = prompt.len() + max_new + b + 8;
assert!(max_ctx <= c.sliding_window,
"first-light dflash round is windowless — ctx cap {} exceeds the draft window {}",
max_ctx, c.sliding_window);
let mut cache = Cache::new(e, &self.cfg, max_ctx)?;
let tp = prompt.len();
cache.dflash_taps = Some(DflashTapSink {
layer_ids: c.target_layer_ids.clone(),
buf: e.uninit(tp * n_taps * n_embd)?,
hidden: n_embd, t: tp,
});
let t_prime = std::time::Instant::now();
let (logits, _h_seed, _hiddens) = self.prime_cache(e, prompt, &mut cache)?;
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,
});
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)
}
}