use std::path::Path;
use anyhow::{Context, Result};
use crate::cpu_gemm::{PackedWeight, gemm_packed};
use crate::weights::LazySt;
pub(crate) struct Lin {
pub(crate) w: Vec<f32>,
pub(crate) n: usize,
pub(crate) k: usize,
packed: std::sync::OnceLock<PackedWeight>,
}
impl Lin {
fn load(st: &LazySt, name: &str, n: usize, k: usize) -> Result<Self> {
let w = st.tensor_f32(name)?;
anyhow::ensure!(w.len() == n * k, "{name} {} != {n}x{k}", w.len());
Ok(Self {
w,
n,
k,
packed: std::sync::OnceLock::new(),
})
}
fn forward(&self, x: &[f32]) -> Vec<f32> {
let m = x.len() / self.k;
let mut out = vec![0f32; m * self.n];
let packed = self
.packed
.get_or_init(|| PackedWeight::new(&self.w, self.n, self.k));
gemm_packed(&mut out, x, packed, m, None);
out
}
}
#[derive(Clone, Copy)]
pub(crate) struct RopeParams {
theta: f64,
rope_angles: usize,
factor: f64,
}
pub struct DgConfig {
pub vocab: usize,
pub hidden: usize,
pub inter: usize,
pub n_layers: usize,
pub n_heads: usize,
pub n_kv: usize,
pub head_dim: usize,
pub n_kv_global: usize,
pub head_dim_global: usize,
pub num_experts: usize,
pub top_k: usize,
pub moe_inter: usize,
pub eps: f32,
pub layer_sliding: Vec<bool>,
pub sliding_window: usize,
pub(crate) rope_sliding: RopeParams,
pub(crate) rope_full: RopeParams,
pub softcap: f32,
pub canvas_length: usize,
}
impl DgConfig {
pub fn load(dir: &Path) -> Result<Self> {
let v: serde_json::Value = serde_json::from_slice(
&std::fs::read(dir.join("config.json")).context("config.json")?,
)?;
let t = v.get("text_config").unwrap_or(&v);
let g = |k: &str| -> Result<usize> {
t.get(k)
.and_then(|x| x.as_u64())
.map(|u| u as usize)
.with_context(|| format!("text_config.{k}"))
};
let n_layers = g("num_hidden_layers")?;
let n_heads = g("num_attention_heads")?;
let hidden = g("hidden_size")?;
let head_dim = g("head_dim").unwrap_or(hidden / n_heads);
let n_kv = g("num_key_value_heads").unwrap_or(n_heads);
let layer_sliding: Vec<bool> = t
.get("layer_types")
.and_then(|x| x.as_array())
.map(|a| {
a.iter()
.map(|s| s.as_str() == Some("sliding_attention"))
.collect()
})
.with_context(|| "text_config.layer_types")?;
anyhow::ensure!(layer_sliding.len() == n_layers, "layer_types length");
let hd_global = g("global_head_dim").unwrap_or(head_dim);
let rope_for = |lt: &str, hd: usize| -> Result<RopeParams> {
let p = t
.get("rope_parameters")
.and_then(|x| x.get(lt))
.with_context(|| format!("rope_parameters.{lt}"))?;
let theta = p
.get("rope_theta")
.and_then(|x| x.as_f64())
.with_context(|| format!("rope_parameters.{lt}.rope_theta"))?;
let rope_type = p
.get("rope_type")
.and_then(|x| x.as_str())
.unwrap_or("default");
let factor = p.get("factor").and_then(|x| x.as_f64()).unwrap_or(1.0);
let rope_angles = match rope_type {
"default" => hd / 2,
"proportional" => {
let prf = p
.get("partial_rotary_factor")
.and_then(|x| x.as_f64())
.unwrap_or(1.0);
(prf * hd as f64 / 2.0) as usize
}
other => anyhow::bail!("unsupported rope_type {other:?} for {lt}"),
};
Ok(RopeParams {
theta,
rope_angles,
factor,
})
};
Ok(Self {
vocab: g("vocab_size")?,
hidden,
inter: g("intermediate_size")?,
n_layers,
n_heads,
n_kv,
head_dim,
n_kv_global: g("num_global_key_value_heads").unwrap_or(n_kv),
head_dim_global: hd_global,
num_experts: g("num_experts").unwrap_or(0),
top_k: g("top_k_experts").unwrap_or(0),
moe_inter: g("moe_intermediate_size").unwrap_or(0),
eps: t
.get("rms_norm_eps")
.and_then(|x| x.as_f64())
.unwrap_or(1e-6) as f32,
layer_sliding,
sliding_window: g("sliding_window").unwrap_or(512),
rope_sliding: rope_for("sliding_attention", head_dim)?,
rope_full: rope_for("full_attention", hd_global)?,
softcap: t
.get("final_logit_softcapping")
.and_then(|x| x.as_f64())
.unwrap_or(30.0) as f32,
canvas_length: v
.get("canvas_length")
.and_then(|x| x.as_u64())
.unwrap_or(256) as usize,
})
}
}
pub(crate) struct DgLayer {
pub(crate) sliding: bool,
pub(crate) n_kv: usize,
pub(crate) hd: usize,
pub(crate) q: Lin,
pub(crate) k: Lin,
pub(crate) v: Option<Lin>,
pub(crate) o: Lin,
pub(crate) q_norm: Vec<f32>,
pub(crate) k_norm: Vec<f32>,
pub(crate) ln_in: Vec<f32>,
pub(crate) ln_post_attn: Vec<f32>,
pub(crate) ln_pre_ff: Vec<f32>,
pub(crate) ln_post_ff: Vec<f32>,
pub(crate) ln_post_ff1: Vec<f32>,
pub(crate) ln_post_ff2: Vec<f32>,
pub(crate) ln_pre_ff2: Vec<f32>,
pub(crate) layer_scalar_dec: f32,
pub(crate) layer_scalar_enc: f32,
pub(crate) gate: Lin,
pub(crate) up: Lin,
pub(crate) down: Lin,
pub(crate) router_proj: Lin,
pub(crate) router_scale: Vec<f32>,
pub(crate) per_expert_scale: Vec<f32>,
pub(crate) experts_gate_up: Vec<f32>,
pub(crate) experts_down: Vec<f32>,
}
pub(crate) struct SelfCond {
pub(crate) pre_norm: Vec<f32>,
pub(crate) gate: Lin,
pub(crate) up: Lin,
pub(crate) down: Lin,
}
pub struct DgCache {
pub k: Vec<Vec<f32>>,
pub v: Vec<Vec<f32>>,
pub ctx: Vec<usize>,
pub seq_len: usize,
}
impl DgCache {
pub fn empty(n_layers: usize) -> Self {
Self {
k: vec![Vec::new(); n_layers],
v: vec![Vec::new(); n_layers],
ctx: vec![0; n_layers],
seq_len: 0,
}
}
}
pub struct DgDecoder {
pub cfg: DgConfig,
pub(crate) embed: Vec<f32>,
pub(crate) layers: Vec<DgLayer>,
pub(crate) final_norm: Vec<f32>,
pub(crate) sc: SelfCond,
}
fn rmsnorm(x: &mut [f32], w: Option<&[f32]>, n: usize, eps: f32) {
for row in x.chunks_mut(n) {
let ms = row.iter().map(|v| v * v).sum::<f32>() / n as f32 + eps;
let inv = 1.0 / ms.sqrt();
match w {
Some(w) => {
for (v, wi) in row.iter_mut().zip(w) {
*v *= inv * wi;
}
}
None => {
for v in row.iter_mut() {
*v *= inv;
}
}
}
}
}
fn gelu_tanh(x: f32) -> f32 {
0.5 * x * (1.0 + (0.797_884_6 * (x + 0.044_715 * x * x * x)).tanh())
}
pub(crate) fn softmax_f32(row: &mut [f32]) {
let m = row.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut s = 0.0;
for v in row.iter_mut() {
*v = (*v - m).exp();
s += *v;
}
for v in row.iter_mut() {
*v /= s;
}
}
fn rope_apply(x: &mut [f32], cos_sin: &[f32], hd: usize) {
let (cos, sin) = cos_sin.split_at(hd);
let half = hd / 2;
for j in 0..half {
let x1 = x[j];
let x2 = x[j + half];
x[j] = x1 * cos[j] - x2 * sin[j];
x[j + half] = x2 * cos[j + half] + x1 * sin[j + half];
}
}
pub(crate) fn rope_cos_sin(pos: usize, hd: usize, rp: RopeParams) -> Vec<f32> {
let half = hd / 2;
let mut out = vec![0f32; 2 * hd];
for j in 0..half {
let inv = if j < rp.rope_angles {
(1.0 / rp.theta.powf(2.0 * j as f64 / hd as f64) / rp.factor) as f32
} else {
0.0
};
let f = pos as f32 * inv;
let (s, c) = f.sin_cos();
out[j] = c;
out[j + half] = c;
out[hd + j] = s;
out[hd + j + half] = s;
}
out
}
impl DgDecoder {
pub fn load(dir: &Path) -> Result<Self> {
let cfg = DgConfig::load(dir)?;
anyhow::ensure!(
cfg.num_experts > 0 && cfg.top_k > 0,
"DiffusionGemma is MoE; num_experts/top_k_experts missing"
);
let st = LazySt::open(dir)?;
let name = |n: &str| -> String {
for pre in ["", "decoder.", "model.decoder."] {
let c = format!("{pre}{n}");
if st.has(&c) {
return c;
}
}
n.to_string() };
let get = |n: &str| -> Result<Vec<f32>> { st.tensor_f32(&name(n)) };
let (h, e, mi) = (cfg.hidden, cfg.num_experts, cfg.moe_inter);
let mut layers = Vec::with_capacity(cfg.n_layers);
for i in 0..cfg.n_layers {
let p = format!("layers.{i}");
let sliding = cfg.layer_sliding[i];
let (n_kv, hd) = if sliding {
(cfg.n_kv, cfg.head_dim)
} else {
(cfg.n_kv_global, cfg.head_dim_global)
};
let scalar = get(&format!("{p}.layer_scalar"))?;
anyhow::ensure!(scalar.len() == 1, "layer_scalar shape");
let scalar_enc = ["", "model."]
.iter()
.find_map(|pre| {
let n = format!("{pre}encoder.language_model.layers.{i}.layer_scalar");
st.has(&n).then(|| st.tensor_f32(&n))
})
.transpose()?
.map_or(scalar[0], |v| v[0]);
let gu = get(&format!("{p}.experts.gate_up_proj"))?;
anyhow::ensure!(gu.len() == e * 2 * mi * h, "experts.gate_up_proj shape");
let ed = get(&format!("{p}.experts.down_proj"))?;
anyhow::ensure!(ed.len() == e * h * mi, "experts.down_proj shape");
layers.push(DgLayer {
sliding,
n_kv,
hd,
q: Lin::load(
&st,
&name(&format!("{p}.self_attn.q_proj.weight")),
cfg.n_heads * hd,
h,
)?,
k: Lin::load(
&st,
&name(&format!("{p}.self_attn.k_proj.weight")),
n_kv * hd,
h,
)?,
v: if sliding {
Some(Lin::load(
&st,
&name(&format!("{p}.self_attn.v_proj.weight")),
n_kv * hd,
h,
)?)
} else {
None
},
o: Lin::load(
&st,
&name(&format!("{p}.self_attn.o_proj.weight")),
h,
cfg.n_heads * hd,
)?,
q_norm: get(&format!("{p}.self_attn.q_norm.weight"))?,
k_norm: get(&format!("{p}.self_attn.k_norm.weight"))?,
ln_in: get(&format!("{p}.input_layernorm.weight"))?,
ln_post_attn: get(&format!("{p}.post_attention_layernorm.weight"))?,
ln_pre_ff: get(&format!("{p}.pre_feedforward_layernorm.weight"))?,
ln_post_ff: get(&format!("{p}.post_feedforward_layernorm.weight"))?,
ln_post_ff1: get(&format!("{p}.post_feedforward_layernorm_1.weight"))?,
ln_post_ff2: get(&format!("{p}.post_feedforward_layernorm_2.weight"))?,
ln_pre_ff2: get(&format!("{p}.pre_feedforward_layernorm_2.weight"))?,
layer_scalar_dec: scalar[0],
layer_scalar_enc: scalar_enc,
gate: Lin::load(
&st,
&name(&format!("{p}.mlp.gate_proj.weight")),
cfg.inter,
h,
)?,
up: Lin::load(&st, &name(&format!("{p}.mlp.up_proj.weight")), cfg.inter, h)?,
down: Lin::load(
&st,
&name(&format!("{p}.mlp.down_proj.weight")),
h,
cfg.inter,
)?,
router_proj: Lin::load(&st, &name(&format!("{p}.router.proj.weight")), e, h)?,
router_scale: get(&format!("{p}.router.scale"))?,
per_expert_scale: get(&format!("{p}.router.per_expert_scale"))?,
experts_gate_up: gu,
experts_down: ed,
});
}
let embed = get("embed_tokens.weight")?;
anyhow::ensure!(embed.len() == cfg.vocab * h, "embed_tokens shape");
Ok(Self {
embed,
layers,
final_norm: get("norm.weight")?,
sc: SelfCond {
pre_norm: get("self_conditioning.pre_norm.weight")?,
gate: Lin::load(
&st,
&name("self_conditioning.gate_proj.weight"),
cfg.inter,
h,
)?,
up: Lin::load(&st, &name("self_conditioning.up_proj.weight"), cfg.inter, h)?,
down: Lin::load(
&st,
&name("self_conditioning.down_proj.weight"),
h,
cfg.inter,
)?,
},
cfg,
})
}
fn ffn_block(&self, l: &DgLayer, x: &mut [f32], layer_scalar: f32) {
let cfg = &self.cfg;
let h = cfg.hidden;
let t_len = x.len() / h;
let residual = x.to_vec();
let mut hf = residual.clone();
rmsnorm(&mut hf, Some(&l.ln_pre_ff), h, cfg.eps);
let g = l.gate.forward(&hf);
let u = l.up.forward(&hf);
let a: Vec<f32> = g.iter().zip(&u).map(|(&g, &u)| gelu_tanh(g) * u).collect();
let mut h1 = l.down.forward(&a);
rmsnorm(&mut h1, Some(&l.ln_post_ff1), h, cfg.eps);
let mut rn = residual.clone();
rmsnorm(&mut rn, None, h, cfg.eps);
let hscale = (h as f32).powf(-0.5);
for row in rn.chunks_mut(h) {
for (v, s) in row.iter_mut().zip(&l.router_scale) {
*v *= s * hscale;
}
}
let mut probs = l.router_proj.forward(&rn); let mut he = residual.clone();
rmsnorm(&mut he, Some(&l.ln_pre_ff2), h, cfg.eps);
let mut h2 = vec![0f32; t_len * h];
let (e_n, mi) = (cfg.num_experts, cfg.moe_inter);
for t in 0..t_len {
let pr = &mut probs[t * e_n..(t + 1) * e_n];
softmax_f32(pr);
let mut idx: Vec<usize> = (0..e_n).collect();
idx.sort_by(|&a, &b| pr[b].partial_cmp(&pr[a]).unwrap());
let top = &idx[..cfg.top_k];
let wsum: f32 = top.iter().map(|&e| pr[e]).sum();
let te = &he[t * h..(t + 1) * h];
for &ei in top {
let w = pr[ei] / wsum * l.per_expert_scale[ei];
let gu = &l.experts_gate_up[ei * 2 * mi * h..(ei + 1) * 2 * mi * h];
let dn = &l.experts_down[ei * h * mi..(ei + 1) * h * mi];
let mut act = vec![0f32; mi];
for m in 0..mi {
let gr = &gu[m * h..(m + 1) * h];
let ur = &gu[(mi + m) * h..(mi + m + 1) * h];
let gv: f32 = gr.iter().zip(te).map(|(a, b)| a * b).sum();
let uv: f32 = ur.iter().zip(te).map(|(a, b)| a * b).sum();
act[m] = gelu_tanh(gv) * uv;
}
let out = &mut h2[t * h..(t + 1) * h];
for j in 0..h {
let dr = &dn[j * mi..(j + 1) * mi];
out[j] += w * dr.iter().zip(&act).map(|(a, b)| a * b).sum::<f32>();
}
}
}
rmsnorm(&mut h2, Some(&l.ln_post_ff2), h, cfg.eps);
let mut comb: Vec<f32> = h1.iter().zip(&h2).map(|(&a, &b)| a + b).collect();
rmsnorm(&mut comb, Some(&l.ln_post_ff), h, cfg.eps);
for (xi, (&r, &c)) in x.iter_mut().zip(residual.iter().zip(&comb)) {
*xi = (r + c) * layer_scalar;
}
}
pub fn canvas_forward(
&self,
ids: &[u32],
cache: &DgCache,
sc_logits: Option<&[f32]>,
) -> (Vec<f32>, Vec<f32>) {
let cfg = &self.cfg;
let (t_len, h) = (ids.len(), cfg.hidden);
let scale = (h as f32).sqrt();
let mut embeds = vec![0f32; t_len * h];
for (t, &id) in ids.iter().enumerate() {
let row = &self.embed[id as usize * h..(id as usize + 1) * h];
for j in 0..h {
embeds[t * h + j] = row[j] * scale;
}
}
let mut sig = vec![0f32; t_len * h];
if let Some(logits) = sc_logits {
for t in 0..t_len {
let mut p = logits[t * cfg.vocab..(t + 1) * cfg.vocab].to_vec();
softmax_f32(&mut p);
let row = &mut sig[t * h..(t + 1) * h];
for (v, pv) in p.iter().enumerate() {
let er = &self.embed[v * h..(v + 1) * h];
for j in 0..h {
row[j] += pv * er[j];
}
}
for v in row.iter_mut() {
*v *= scale;
}
}
}
let mut normed = sig.clone();
rmsnorm(&mut normed, Some(&self.sc.pre_norm), h, cfg.eps);
let g = self.sc.gate.forward(&normed);
let u = self.sc.up.forward(&normed);
let act: Vec<f32> = g.iter().zip(&u).map(|(&g, &u)| gelu_tanh(g) * u).collect();
let sc_out = self.sc.down.forward(&act);
let mut x: Vec<f32> = embeds.iter().zip(&sc_out).map(|(&e, &s)| e + s).collect();
rmsnorm(&mut x, None, h, cfg.eps);
let cs_sliding: Vec<Vec<f32>> = (0..t_len)
.map(|t| rope_cos_sin(cache.seq_len + t, cfg.head_dim, cfg.rope_sliding))
.collect();
let cs_full: Vec<Vec<f32>> = (0..t_len)
.map(|t| rope_cos_sin(cache.seq_len + t, cfg.head_dim_global, cfg.rope_full))
.collect();
for (li, l) in self.layers.iter().enumerate() {
let (n_kv, hd) = (l.n_kv, l.hd);
let n_rep = cfg.n_heads / n_kv;
let cs = if l.sliding { &cs_sliding } else { &cs_full };
let mut hn = x.clone();
rmsnorm(&mut hn, Some(&l.ln_in), h, cfg.eps);
let mut q = l.q.forward(&hn); let k_raw = l.k.forward(&hn); let mut v = match &l.v {
Some(vp) => vp.forward(&hn),
None => k_raw.clone(),
};
let mut k = k_raw;
for t in 0..t_len {
for qh in 0..cfg.n_heads {
let s = &mut q[(t * cfg.n_heads + qh) * hd..(t * cfg.n_heads + qh + 1) * hd];
rmsnorm(s, Some(&l.q_norm), hd, cfg.eps);
rope_apply(s, &cs[t], hd);
}
for kh in 0..n_kv {
let s = &mut k[(t * n_kv + kh) * hd..(t * n_kv + kh + 1) * hd];
rmsnorm(s, Some(&l.k_norm), hd, cfg.eps);
rope_apply(s, &cs[t], hd);
let sv = &mut v[(t * n_kv + kh) * hd..(t * n_kv + kh + 1) * hd];
rmsnorm(sv, None, hd, cfg.eps);
}
}
let ctx = cache.ctx[li];
let span = ctx + t_len;
let (ck, cv) = (&cache.k[li], &cache.v[li]);
let mut attn = vec![0f32; t_len * cfg.n_heads * hd];
for t in 0..t_len {
for qh in 0..cfg.n_heads {
let kv = qh / n_rep;
let qv = &q[(t * cfg.n_heads + qh) * hd..(t * cfg.n_heads + qh + 1) * hd];
let mut scores = vec![0f32; span];
for (s, sc) in scores.iter_mut().enumerate() {
let krow = if s < ctx {
&ck[(kv * ctx + s) * hd..(kv * ctx + s + 1) * hd]
} else {
&k[((s - ctx) * n_kv + kv) * hd..((s - ctx) * n_kv + kv + 1) * hd]
};
*sc = qv.iter().zip(krow).map(|(a, b)| a * b).sum::<f32>();
}
softmax_f32(&mut scores);
let out =
&mut attn[(t * cfg.n_heads + qh) * hd..(t * cfg.n_heads + qh + 1) * hd];
for (s, &w) in scores.iter().enumerate() {
let vrow = if s < ctx {
&cv[(kv * ctx + s) * hd..(kv * ctx + s + 1) * hd]
} else {
&v[((s - ctx) * n_kv + kv) * hd..((s - ctx) * n_kv + kv + 1) * hd]
};
for j in 0..hd {
out[j] += w * vrow[j];
}
}
}
}
let mut ao = l.o.forward(&attn);
rmsnorm(&mut ao, Some(&l.ln_post_attn), h, cfg.eps);
for (xi, a) in x.iter_mut().zip(&ao) {
*xi += a;
}
self.ffn_block(l, &mut x, l.layer_scalar_dec);
}
rmsnorm(&mut x, Some(&self.final_norm), h, cfg.eps);
let mut logits = vec![0f32; t_len * cfg.vocab];
for t in 0..t_len {
let xt = &x[t * h..(t + 1) * h];
let lt = &mut logits[t * cfg.vocab..(t + 1) * cfg.vocab];
for (vv, l) in lt.iter_mut().enumerate() {
let er = &self.embed[vv * h..(vv + 1) * h];
let raw: f32 = er.iter().zip(xt).map(|(a, b)| a * b).sum();
*l = cfg.softcap * (raw / cfg.softcap).tanh();
}
}
(x, logits)
}
pub fn encode(&self, ids: &[u32], cache: &mut DgCache) -> Vec<f32> {
self.encode_probed(ids, cache, None)
}
pub fn encode_probed(
&self,
ids: &[u32],
cache: &mut DgCache,
mut probe: Option<&mut Vec<Vec<f32>>>,
) -> Vec<f32> {
let cfg = &self.cfg;
let (s_len, h) = (ids.len(), cfg.hidden);
let pos0 = cache.seq_len;
let scale = (h as f32).sqrt();
let mut x = vec![0f32; s_len * h];
for (t, &id) in ids.iter().enumerate() {
let row = &self.embed[id as usize * h..(id as usize + 1) * h];
for j in 0..h {
x[t * h + j] = row[j] * scale;
}
}
let cs_sliding: Vec<Vec<f32>> = (0..s_len)
.map(|t| rope_cos_sin(pos0 + t, cfg.head_dim, cfg.rope_sliding))
.collect();
let cs_full: Vec<Vec<f32>> = (0..s_len)
.map(|t| rope_cos_sin(pos0 + t, cfg.head_dim_global, cfg.rope_full))
.collect();
for (li, l) in self.layers.iter().enumerate() {
let (n_kv, hd) = (l.n_kv, l.hd);
let n_rep = cfg.n_heads / n_kv;
let cs = if l.sliding { &cs_sliding } else { &cs_full };
let mut hn = x.clone();
rmsnorm(&mut hn, Some(&l.ln_in), h, cfg.eps);
let mut q = l.q.forward(&hn);
let k_raw = l.k.forward(&hn);
let mut v = match &l.v {
Some(vp) => vp.forward(&hn),
None => k_raw.clone(),
};
let mut k = k_raw;
for t in 0..s_len {
for qh in 0..cfg.n_heads {
let s = &mut q[(t * cfg.n_heads + qh) * hd..(t * cfg.n_heads + qh + 1) * hd];
rmsnorm(s, Some(&l.q_norm), hd, cfg.eps);
rope_apply(s, &cs[t], hd);
}
for kh in 0..n_kv {
let s = &mut k[(t * n_kv + kh) * hd..(t * n_kv + kh + 1) * hd];
rmsnorm(s, Some(&l.k_norm), hd, cfg.eps);
rope_apply(s, &cs[t], hd);
let sv = &mut v[(t * n_kv + kh) * hd..(t * n_kv + kh + 1) * hd];
rmsnorm(sv, None, hd, cfg.eps);
}
}
let ctx = cache.ctx[li];
let mut attn = vec![0f32; s_len * cfg.n_heads * hd];
let (ck, cv) = (&cache.k[li], &cache.v[li]);
for t in 0..s_len {
let p = pos0 + t;
let span = ctx + t + 1;
let lo_logical = if l.sliding && p + 1 >= cfg.sliding_window {
p + 1 - cfg.sliding_window
} else {
0
};
for qh in 0..cfg.n_heads {
let kv = qh / n_rep;
let qv = &q[(t * cfg.n_heads + qh) * hd..(t * cfg.n_heads + qh + 1) * hd];
let mut scores = vec![f32::NEG_INFINITY; span];
for (s, sc) in scores.iter_mut().enumerate() {
let kpos = if s < ctx {
pos0 - ctx + s
} else {
pos0 + (s - ctx)
};
if kpos < lo_logical {
continue;
}
let krow = if s < ctx {
&ck[(kv * ctx + s) * hd..(kv * ctx + s + 1) * hd]
} else {
let st = s - ctx;
&k[(st * n_kv + kv) * hd..(st * n_kv + kv + 1) * hd]
};
*sc = qv.iter().zip(krow).map(|(a, b)| a * b).sum::<f32>();
}
softmax_f32(&mut scores);
let out =
&mut attn[(t * cfg.n_heads + qh) * hd..(t * cfg.n_heads + qh + 1) * hd];
for (s, &w) in scores.iter().enumerate() {
if w == 0.0 {
continue;
}
let vrow = if s < ctx {
&cv[(kv * ctx + s) * hd..(kv * ctx + s + 1) * hd]
} else {
let st = s - ctx;
&v[(st * n_kv + kv) * hd..(st * n_kv + kv + 1) * hd]
};
for j in 0..hd {
out[j] += w * vrow[j];
}
}
}
}
let new_ctx = ctx + s_len;
let mut nk = vec![0f32; n_kv * new_ctx * hd];
let mut nv = vec![0f32; n_kv * new_ctx * hd];
for kh in 0..n_kv {
for s in 0..ctx {
nk[(kh * new_ctx + s) * hd..(kh * new_ctx + s + 1) * hd]
.copy_from_slice(&ck[(kh * ctx + s) * hd..(kh * ctx + s + 1) * hd]);
nv[(kh * new_ctx + s) * hd..(kh * new_ctx + s + 1) * hd]
.copy_from_slice(&cv[(kh * ctx + s) * hd..(kh * ctx + s + 1) * hd]);
}
for t in 0..s_len {
nk[(kh * new_ctx + ctx + t) * hd..(kh * new_ctx + ctx + t + 1) * hd]
.copy_from_slice(&k[(t * n_kv + kh) * hd..(t * n_kv + kh + 1) * hd]);
nv[(kh * new_ctx + ctx + t) * hd..(kh * new_ctx + ctx + t + 1) * hd]
.copy_from_slice(&v[(t * n_kv + kh) * hd..(t * n_kv + kh + 1) * hd]);
}
}
let keep = if l.sliding {
new_ctx.min(cfg.sliding_window - 1)
} else {
new_ctx
};
if keep < new_ctx {
let drop = new_ctx - keep;
let mut tk = vec![0f32; n_kv * keep * hd];
let mut tv = vec![0f32; n_kv * keep * hd];
for kh in 0..n_kv {
tk[kh * keep * hd..(kh + 1) * keep * hd].copy_from_slice(
&nk[(kh * new_ctx + drop) * hd..(kh * new_ctx + new_ctx) * hd],
);
tv[kh * keep * hd..(kh + 1) * keep * hd].copy_from_slice(
&nv[(kh * new_ctx + drop) * hd..(kh * new_ctx + new_ctx) * hd],
);
}
cache.k[li] = tk;
cache.v[li] = tv;
cache.ctx[li] = keep;
} else {
cache.k[li] = nk;
cache.v[li] = nv;
cache.ctx[li] = new_ctx;
}
let mut ao = l.o.forward(&attn);
rmsnorm(&mut ao, Some(&l.ln_post_attn), h, cfg.eps);
for (xi, a) in x.iter_mut().zip(&ao) {
*xi += a;
}
self.ffn_block(l, &mut x, l.layer_scalar_enc);
if let Some(p) = probe.as_deref_mut() {
p.push(x.clone());
}
}
cache.seq_len += s_len;
rmsnorm(&mut x, Some(&self.final_norm), h, cfg.eps);
x
}
}
#[derive(Clone)]
pub struct DgGenConfig {
pub max_new_tokens: usize,
pub max_denoising_steps: usize,
pub entropy_bound: f32,
pub t_min: f32,
pub t_max: f32,
pub stability_threshold: usize,
pub confidence_threshold: f32,
pub eos_token_id: Option<u32>,
pub pad_token_id: u32,
}
impl Default for DgGenConfig {
fn default() -> Self {
Self {
max_new_tokens: 256,
max_denoising_steps: 48,
entropy_bound: 0.1,
t_min: 0.4,
t_max: 0.8,
stability_threshold: 1,
confidence_threshold: 0.005,
eos_token_id: None,
pad_token_id: 0,
}
}
}
pub trait DgRng {
fn uniform_canvas(&mut self, canvas_len: usize, vocab: usize) -> Vec<u32>;
fn multinomial(&mut self, probs: &[f32]) -> u32;
}
pub struct DgXorShiftRng(pub u64);
impl DgXorShiftRng {
fn next_u64(&mut self) -> u64 {
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 7;
self.0 ^= self.0 << 17;
self.0.wrapping_mul(0x2545_F491_4F6C_DD1D)
}
fn next_f32(&mut self) -> f32 {
((self.next_u64() >> 40) as f32) / (1u64 << 24) as f32
}
}
impl DgRng for DgXorShiftRng {
fn uniform_canvas(&mut self, canvas_len: usize, vocab: usize) -> Vec<u32> {
(0..canvas_len)
.map(|_| (self.next_u64() % vocab as u64) as u32)
.collect()
}
fn multinomial(&mut self, probs: &[f32]) -> u32 {
let u = self.next_f32();
let mut acc = 0f32;
for (i, &p) in probs.iter().enumerate() {
acc += p;
if u < acc {
return i as u32;
}
}
probs.len() as u32 - 1
}
}
pub struct DgStepRecord {
pub current: Vec<u32>,
pub processed_logits: Vec<f32>,
pub denoiser: Vec<u32>,
pub accepted: Vec<u32>,
pub renoised: Vec<u32>,
pub argmax: Vec<u32>,
pub stopped: bool,
}
fn entropy_of_logits(row: &[f32]) -> f32 {
let m = row.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut se = 0f32;
for &v in row {
se += (v - m).exp();
}
let lse = m + se.ln();
let mut ent = 0f32;
for &v in row {
ent += (v - lse).exp() * (lse - v);
}
ent
}
fn accept_mask(entropies: &[f32], bound: f32) -> Vec<bool> {
let t = entropies.len();
let mut idx: Vec<usize> = (0..t).collect();
idx.sort_by(|&a, &b| entropies[a].partial_cmp(&entropies[b]).unwrap());
let mut mask = vec![false; t];
let mut cum = 0f32;
for &i in &idx {
cum += entropies[i];
if cum - entropies[i] <= bound {
mask[i] = true;
}
}
mask
}
struct DgStopping {
history: Vec<Vec<i64>>,
stability_threshold: usize,
confidence_threshold: f32,
}
impl DgStopping {
fn new(stability_threshold: usize, confidence_threshold: f32, canvas: usize) -> Self {
Self {
history: vec![vec![-1; canvas]; stability_threshold],
stability_threshold,
confidence_threshold,
}
}
fn check(&mut self, argmax: &[u32], processed_logits: &[f32], vocab: usize) -> bool {
let stable = if self.stability_threshold == 0 {
true
} else {
let s = self
.history
.iter()
.all(|row| row.iter().zip(argmax).all(|(&h, &a)| h == a as i64));
self.history.rotate_left(1);
*self.history.last_mut().unwrap() = argmax.iter().map(|&a| a as i64).collect();
s
};
let t = argmax.len();
let mean_ent = (0..t)
.map(|i| entropy_of_logits(&processed_logits[i * vocab..(i + 1) * vocab]))
.sum::<f32>()
/ t as f32;
stable && mean_ent < self.confidence_threshold
}
}
impl DgDecoder {
pub fn denoise_block(
&self,
cache: &DgCache,
gc: &DgGenConfig,
rng: &mut dyn DgRng,
on_step: Option<&mut dyn FnMut(&DgStepRecord)>,
) -> (Vec<u32>, usize) {
let fwd = |ids: &[u32], sc: Option<&[f32]>| self.canvas_forward(ids, cache, sc).1;
self.denoise_block_with(gc, rng, on_step, fwd)
}
pub fn denoise_block_with(
&self,
gc: &DgGenConfig,
rng: &mut dyn DgRng,
mut on_step: Option<&mut dyn FnMut(&DgStepRecord)>,
mut fwd: impl FnMut(&[u32], Option<&[f32]>) -> Vec<f32>,
) -> (Vec<u32>, usize) {
let canvas = self.cfg.canvas_length;
let vocab = self.cfg.vocab;
let n = gc.max_denoising_steps;
let mut current = rng.uniform_canvas(canvas, vocab);
let mut sc: Option<Vec<f32>> = None;
let mut stopping = DgStopping::new(gc.stability_threshold, gc.confidence_threshold, canvas);
let mut argmax_canvas = current.clone();
let mut steps = 0;
for cur_step in (1..=n).rev() {
steps += 1;
let raw_logits = fwd(¤t, sc.as_deref());
let t = gc.t_min + (gc.t_max - gc.t_min) * (cur_step as f32 / n as f32);
let mut processed = raw_logits;
for v in processed.iter_mut() {
*v /= t;
}
let mut denoiser = vec![0u32; canvas];
let mut argmax = vec![0u32; canvas];
let mut entropies = vec![0f32; canvas];
for i in 0..canvas {
let row = &processed[i * vocab..(i + 1) * vocab];
let mut probs = row.to_vec();
softmax_f32(&mut probs);
denoiser[i] = rng.multinomial(&probs);
argmax[i] = row
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(j, _)| j as u32)
.unwrap();
entropies[i] = entropy_of_logits(row);
}
let mask = accept_mask(&entropies, gc.entropy_bound);
let accepted: Vec<u32> = (0..canvas)
.map(|i| if mask[i] { denoiser[i] } else { current[i] })
.collect();
let random_canvas = rng.uniform_canvas(canvas, vocab);
let renoised: Vec<u32> = (0..canvas)
.map(|i| {
if mask[i] {
accepted[i]
} else {
random_canvas[i]
}
})
.collect();
let stopped = stopping.check(&argmax, &processed, vocab);
if let Some(f) = on_step.as_deref_mut() {
f(&DgStepRecord {
current: current.clone(),
processed_logits: processed.clone(),
denoiser: denoiser.clone(),
accepted: accepted.clone(),
renoised: renoised.clone(),
argmax: argmax.clone(),
stopped,
});
}
sc = Some(processed);
current = renoised;
argmax_canvas = argmax;
if stopped {
break;
}
}
(argmax_canvas, steps)
}
pub fn generate(&self, prompt: &[u32], gc: &DgGenConfig, rng: &mut dyn DgRng) -> Vec<u32> {
let canvas = self.cfg.canvas_length;
let max_new_canvases = gc.max_new_tokens.div_ceil(canvas);
let mut cache = DgCache::empty(self.cfg.n_layers);
let mut out: Vec<u32> = prompt.to_vec();
let mut to_encode: Vec<u32> = prompt.to_vec();
for _ in 0..max_new_canvases {
self.encode(&to_encode, &mut cache);
let (mut tokens, _) = self.denoise_block(&cache, gc, rng, None);
let mut finished = false;
if let Some(eos) = gc.eos_token_id
&& let Some(p) = tokens.iter().position(|&t| t == eos)
{
for t in tokens[p + 1..].iter_mut() {
*t = gc.pad_token_id;
}
finished = true;
}
out.extend_from_slice(&tokens);
if finished {
break;
}
to_encode = tokens;
}
out
}
}