pub fn hc_split_sinkhorn(
mixes: &[f32],
hc_scale: &[f32; 3],
hc_base: &[f32],
hc: usize,
iters: usize,
eps: f32,
pre: &mut [f32],
post: &mut [f32],
comb: &mut [f32],
) {
debug_assert_eq!(mixes.len(), (2 + hc) * hc);
debug_assert_eq!(comb.len(), hc * hc);
for j in 0..hc {
pre[j] = sigmoid(mixes[j] * hc_scale[0] + hc_base[j]) + eps;
post[j] = 2.0 * sigmoid(mixes[j + hc] * hc_scale[1] + hc_base[j + hc]);
}
for j in 0..hc {
for k in 0..hc {
let idx = j * hc + k + hc * 2;
comb[j * hc + k] = mixes[idx] * hc_scale[2] + hc_base[idx];
}
}
for j in 0..hc {
let row = &mut comb[j * hc..(j + 1) * hc];
let m = row.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut sum = 0.0;
for v in row.iter_mut() {
*v = (*v - m).exp();
sum += *v;
}
for v in row.iter_mut() {
*v = *v / sum + eps;
}
}
normalize_cols(comb, hc, eps);
for _ in 0..iters.saturating_sub(1) {
normalize_rows(comb, hc, eps);
normalize_cols(comb, hc, eps);
}
}
fn normalize_rows(m: &mut [f32], n: usize, eps: f32) {
for j in 0..n {
let s: f32 = m[j * n..(j + 1) * n].iter().sum::<f32>() + eps;
for v in m[j * n..(j + 1) * n].iter_mut() {
*v /= s;
}
}
}
fn normalize_cols(m: &mut [f32], n: usize, eps: f32) {
for k in 0..n {
let mut s = eps;
for j in 0..n {
s += m[j * n + k];
}
for j in 0..n {
m[j * n + k] /= s;
}
}
}
#[inline]
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
pub fn hc_mixes(
x_flat: &[f32],
hc_fn: &[f32],
mix_hc: usize,
eps: f32,
pool: Option<&crate::pool::Pool>,
out: &mut [f32],
) {
let n = x_flat.len();
debug_assert_eq!(hc_fn.len(), mix_hc * n);
debug_assert_eq!(out.len(), mix_hc);
let ms = x_flat.iter().map(|v| v * v).sum::<f32>() / n as f32;
let rsqrt = 1.0 / (ms + eps).sqrt();
match pool {
Some(p) if n >= 4096 => {
let addr = crate::pool::SendMut::new(out.as_mut_ptr());
p.run_rows(mix_hc, &|start, end| {
for i in start..end {
let row = &hc_fn[i * n..(i + 1) * n];
let v = row.iter().zip(x_flat).map(|(a, b)| a * b).sum::<f32>() * rsqrt;
unsafe { *addr.at(i) = v };
}
});
}
_ => {
for (i, o) in out.iter_mut().enumerate() {
let row = &hc_fn[i * n..(i + 1) * n];
*o = row.iter().zip(x_flat).map(|(a, b)| a * b).sum::<f32>() * rsqrt;
}
}
}
}
pub fn hc_fold(x: &[f32], pre: &[f32], hc: usize, dim: usize, out: &mut [f32]) {
debug_assert_eq!(x.len(), hc * dim);
out.fill(0.0);
for j in 0..hc {
let w = pre[j];
let src = &x[j * dim..(j + 1) * dim];
for (o, v) in out.iter_mut().zip(src) {
*o += w * v;
}
}
}
pub fn hc_expand(
block_out: &[f32],
residual: &[f32],
post: &[f32],
comb: &[f32],
hc: usize,
dim: usize,
out: &mut [f32],
) {
debug_assert_eq!(residual.len(), hc * dim);
debug_assert_eq!(out.len(), hc * dim);
for j in 0..hc {
let dst = &mut out[j * dim..(j + 1) * dim];
let p = post[j];
for (d, o) in dst.iter_mut().enumerate() {
*o = p * block_out[d];
}
for k in 0..hc {
let w = comb[k * hc + j];
let src = &residual[k * dim..(k + 1) * dim];
for (o, v) in dst.iter_mut().zip(src) {
*o += w * v;
}
}
}
}
pub fn hc_head_pre(mixes: &[f32], scale: f32, base: &[f32], hc: usize, eps: f32, pre: &mut [f32]) {
for j in 0..hc {
pre[j] = sigmoid(mixes[j] * scale + base[j]) + eps;
}
}
pub fn route(
scores_in: &[f32],
bias: Option<&[f32]>,
top_k: usize,
route_scale: f32,
forced: Option<&[usize]>,
mask: Option<&[bool]>,
indices: &mut Vec<usize>,
weights: &mut Vec<f32>,
) {
let n = scores_in.len();
let mut scores = Vec::with_capacity(n);
for &s in scores_in {
let sp = if s > 20.0 { s } else { (1.0 + s.exp()).ln() };
scores.push(sp.sqrt());
}
indices.clear();
weights.clear();
match forced {
Some(f) => indices.extend(f.iter().copied()),
None => {
let mut shifted: Vec<f32> = match bias {
Some(b) => scores.iter().zip(b).map(|(s, b)| s + b).collect(),
None => scores.clone(),
};
if let Some(m) = mask {
for (i, s) in shifted.iter_mut().enumerate() {
if !m.get(i).copied().unwrap_or(true) {
*s = f32::NEG_INFINITY;
}
}
}
for _ in 0..top_k.min(n) {
let mut best = 0usize;
let mut bv = f32::NEG_INFINITY;
for (i, &v) in shifted.iter().enumerate() {
if v > bv {
bv = v;
best = i;
}
}
if !bv.is_finite() {
break;
}
indices.push(best);
shifted[best] = f32::NEG_INFINITY;
}
}
}
for &i in indices.iter() {
weights.push(scores.get(i).copied().unwrap_or(0.0));
}
let sum: f32 = weights.iter().sum();
if sum > 0.0 {
for w in weights.iter_mut() {
*w = *w / sum * route_scale;
}
}
}
pub fn hash_route(tid2eid: &[f32], vocab: usize, top_k: usize, tid: u32) -> Vec<usize> {
let row = (tid as usize).min(vocab.saturating_sub(1)) * top_k;
(0..top_k)
.map(|k| tid2eid.get(row + k).copied().unwrap_or(0.0) as usize)
.collect()
}
pub fn rope_tail(v: &mut [f32], inv_freq: &[f32], pos: usize, rd: usize, inverse: bool) {
let n = v.len();
debug_assert!(
rd <= n && rd % 2 == 0,
"rope tail {rd} wider than the vector {n}"
);
let rd = rd.min(n) & !1;
let base = n - rd;
for i in 0..rd / 2 {
let theta = pos as f32 * inv_freq[i];
let (s, c) = (theta.sin(), theta.cos());
let s = if inverse { -s } else { s };
let a = v[base + 2 * i];
let b = v[base + 2 * i + 1];
v[base + 2 * i] = a * c - b * s;
v[base + 2 * i + 1] = a * s + b * c;
}
}
pub fn rms_inplace(v: &mut [f32], eps: f32) {
let ms = v.iter().map(|x| x * x).sum::<f32>() / v.len() as f32;
let inv = 1.0 / (ms + eps).sqrt();
for x in v.iter_mut() {
*x *= inv;
}
}
pub fn sparse_attend(
q: &[f32],
kv: &[f32],
idxs: &[usize],
sink: f32,
scale: f32,
head_dim: usize,
out: &mut [f32],
) {
let mut m = sink;
let mut scores = Vec::with_capacity(idxs.len());
for &p in idxs {
if p == usize::MAX {
scores.push(f32::NEG_INFINITY);
continue;
}
let k = &kv[p * head_dim..(p + 1) * head_dim];
let dot: f32 = q.iter().zip(k).map(|(a, b)| a * b).sum::<f32>() * scale;
m = m.max(dot);
scores.push(dot);
}
let mut denom = (sink - m).exp();
out.fill(0.0);
for (&p, &s) in idxs.iter().zip(&scores) {
if p == usize::MAX {
continue;
}
let w = (s - m).exp();
denom += w;
let v = &kv[p * head_dim..(p + 1) * head_dim];
for (o, x) in out.iter_mut().zip(v) {
*o += w * x;
}
}
if std::env::var("CMF_ATTN_DEBUG").is_ok() {
eprintln!(
" [порт] позиций={} score={:?} sink={sink:.4} denom={denom:.4} |q|={:.3}",
idxs.iter().filter(|&&p| p != usize::MAX).count(),
scores
.iter()
.map(|x| (x * 10000.0).round() / 10000.0)
.collect::<Vec<_>>(),
q.iter().map(|x| x * x).sum::<f32>().sqrt()
);
}
let inv = 1.0 / denom;
for o in out.iter_mut() {
*o *= inv;
}
}
pub fn o_project(
attn: &[f32],
wo_a_row: &(dyn Fn(usize, &[f32], &mut [f32]) -> f32 + Sync),
scratch_len: usize,
wo_b: &dyn Fn(&[f32], &mut [f32]),
groups: usize,
lora: usize,
pool: Option<&crate::pool::Pool>,
out: &mut [f32],
) {
let per_group = attn.len() / groups;
let mut mid = vec![0.0f32; groups * lora];
let slice_of = |i: usize| {
let g = i / lora;
&attn[g * per_group..(g + 1) * per_group]
};
match pool {
Some(p) if mid.len() >= 256 => {
let addr = crate::pool::SendMut::new(mid.as_mut_ptr());
p.run_rows(mid.len(), &|start, end| {
let mut sc = vec![0.0f32; scratch_len];
for i in start..end {
let v = wo_a_row(i, slice_of(i), &mut sc);
unsafe { *addr.at(i) = v };
}
});
}
_ => {
let mut sc = vec![0.0f32; scratch_len];
for (i, m) in mid.iter_mut().enumerate() {
*m = wo_a_row(i, slice_of(i), &mut sc);
}
}
}
wo_b(&mid, out);
}
pub fn compress_window(
kv: &[f32],
score: &[f32],
ape: &[f32],
ratio: usize,
width: usize,
out: &mut [f32],
) {
debug_assert_eq!(kv.len(), ratio * width);
debug_assert_eq!(ape.len(), ratio * width);
let biased: Vec<f32> = score.iter().zip(ape).map(|(s, a)| s + a).collect();
pool_by_score(kv, &biased, ratio, width, out);
}
pub fn pool_by_score(kv: &[f32], score: &[f32], slots: usize, width: usize, out: &mut [f32]) {
debug_assert_eq!(kv.len(), slots * width);
debug_assert_eq!(score.len(), slots * width);
out.fill(0.0);
for d in 0..width {
let mut m = f32::NEG_INFINITY;
for t in 0..slots {
m = m.max(score[t * width + d]);
}
if !m.is_finite() {
continue;
}
let mut denom = 0.0;
for t in 0..slots {
denom += (score[t * width + d] - m).exp();
}
if denom <= 0.0 {
continue;
}
for t in 0..slots {
out[d] += ((score[t * width + d] - m).exp() / denom) * kv[t * width + d];
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn compress_window_overlap(
prev_kv: &[f32],
prev_score: &[f32],
cur_kv: &[f32],
cur_score: &[f32],
ratio: usize,
d: usize,
out: &mut [f32],
) {
let slots = 2 * ratio;
let mut kv = vec![0.0f32; slots * d];
let mut sc = vec![f32::NEG_INFINITY; slots * d];
let have_prev = prev_kv.len() == ratio * 2 * d;
for t in 0..ratio {
if have_prev {
kv[t * d..(t + 1) * d].copy_from_slice(&prev_kv[t * 2 * d..t * 2 * d + d]);
sc[t * d..(t + 1) * d].copy_from_slice(&prev_score[t * 2 * d..t * 2 * d + d]);
}
let src = t * 2 * d + d;
let dst = (ratio + t) * d;
kv[dst..dst + d].copy_from_slice(&cur_kv[src..src + d]);
sc[dst..dst + d].copy_from_slice(&cur_score[src..src + d]);
}
pool_by_score(&kv, &sc, slots, d, out);
}
#[allow(clippy::too_many_arguments)]
pub fn index_scores(
q_heads: &[f32],
kv: &[f32],
head_weights: &[f32],
n_heads: usize,
head_dim: usize,
n_pos: usize,
causal_limit: usize,
pool: Option<&crate::pool::Pool>,
out: &mut Vec<f32>,
) {
out.clear();
out.resize(n_pos, 0.0);
let score_at = |t: usize| -> f32 {
if t >= causal_limit {
return f32::NEG_INFINITY;
}
let k = &kv[t * head_dim..(t + 1) * head_dim];
let mut acc = 0.0;
for h in 0..n_heads {
let q = &q_heads[h * head_dim..(h + 1) * head_dim];
let dot: f32 = q.iter().zip(k).map(|(a, b)| a * b).sum();
acc += dot.max(0.0) * head_weights[h];
}
acc
};
match pool {
Some(p) if n_pos >= 64 => {
let addr = crate::pool::SendMut::new(out.as_mut_ptr());
p.run_rows(n_pos, &|start, end| {
for t in start..end {
unsafe { *addr.at(t) = score_at(t) };
}
});
}
_ => {
for (t, o) in out.iter_mut().enumerate() {
*o = score_at(t);
}
}
}
}
pub fn top_k_positions(scores: &[f32], k: usize, out: &mut Vec<usize>) {
out.clear();
if k >= scores.len() {
out.extend(
scores
.iter()
.enumerate()
.filter(|(_, v)| v.is_finite())
.map(|(i, _)| i),
);
return;
}
let mut taken = vec![false; scores.len()];
for _ in 0..k.min(scores.len()) {
let mut best = usize::MAX;
let mut bv = f32::NEG_INFINITY;
for (i, &v) in scores.iter().enumerate() {
if !taken[i] && v > bv && v.is_finite() {
bv = v;
best = i;
}
}
if best == usize::MAX {
break;
}
taken[best] = true;
out.push(best);
}
out.sort_unstable();
}
#[allow(clippy::too_many_arguments)]
pub fn expert_swiglu(
x: &[f32],
w1: &dyn Fn(&[f32], &mut [f32]),
w3: &dyn Fn(&[f32], &mut [f32]),
w2: &dyn Fn(&[f32], &mut [f32]),
inter: usize,
weight: f32,
limit: f32,
out: &mut [f32],
) {
let mut gate = vec![0.0f32; inter];
let mut up = vec![0.0f32; inter];
w1(x, &mut gate);
w3(x, &mut up);
if limit > 0.0 {
for u in up.iter_mut() {
*u = u.clamp(-limit, limit);
}
for g in gate.iter_mut() {
*g = g.min(limit);
}
}
for (g, u) in gate.iter_mut().zip(&up) {
let silu = *g / (1.0 + (-*g).exp());
*g = silu * u * weight;
}
w2(&gate, out);
}
#[derive(Debug, Clone, Copy)]
pub struct Dsv4Cfg {
pub dim: usize,
pub n_heads: usize,
pub head_dim: usize,
pub rope_head_dim: usize,
pub q_lora_rank: usize,
pub o_lora_rank: usize,
pub o_groups: usize,
pub hc_mult: usize,
pub hc_sinkhorn_iters: usize,
pub hc_eps: f32,
pub norm_eps: f32,
pub n_routed_experts: usize,
pub top_k: usize,
pub moe_inter: usize,
pub route_scale: f32,
pub swiglu_limit: f32,
pub window: usize,
pub index_topk: usize,
pub vocab: usize,
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn hc_block<F: FnMut(&[f32], &mut [f32])>(
state: &mut [f32],
hc_fn: &[f32],
hc_scale: &[f32; 3],
hc_base: &[f32],
norm_w: &[f32],
cfg: &Dsv4Cfg,
scratch: &mut HcScratch,
pool: Option<&crate::pool::Pool>,
mut block: F,
) {
let (hc, dim) = (cfg.hc_mult, cfg.dim);
let mix_hc = (2 + hc) * hc;
hc_mixes(state, hc_fn, mix_hc, cfg.norm_eps, pool, &mut scratch.mixes);
hc_split_sinkhorn(
&scratch.mixes,
hc_scale,
hc_base,
hc,
cfg.hc_sinkhorn_iters,
cfg.hc_eps,
&mut scratch.pre,
&mut scratch.post,
&mut scratch.comb,
);
hc_fold(state, &scratch.pre, hc, dim, &mut scratch.folded);
let ms = scratch.folded.iter().map(|v| v * v).sum::<f32>() / dim as f32;
let inv = 1.0 / (ms + cfg.norm_eps).sqrt();
for (v, w) in scratch.folded.iter_mut().zip(norm_w) {
*v = *v * inv * w;
}
block(&scratch.folded, &mut scratch.block_out);
scratch.residual.copy_from_slice(state);
hc_expand(
&scratch.block_out,
&scratch.residual,
&scratch.post,
&scratch.comb,
hc,
dim,
state,
);
}
pub struct HcScratch {
pub mixes: Vec<f32>,
pub pre: Vec<f32>,
pub post: Vec<f32>,
pub comb: Vec<f32>,
pub folded: Vec<f32>,
pub block_out: Vec<f32>,
pub residual: Vec<f32>,
}
impl HcScratch {
pub fn new(cfg: &Dsv4Cfg) -> Self {
let (hc, dim) = (cfg.hc_mult, cfg.dim);
Self {
mixes: vec![0.0; (2 + hc) * hc],
pre: vec![0.0; hc],
post: vec![0.0; hc],
comb: vec![0.0; hc * hc],
folded: vec![0.0; dim],
block_out: vec![0.0; dim],
residual: vec![0.0; hc * dim],
}
}
}
pub fn hc_head_fold(
state: &[f32],
hc_fn: &[f32],
hc_scale: f32,
hc_base: &[f32],
cfg: &Dsv4Cfg,
pool: Option<&crate::pool::Pool>,
out: &mut [f32],
) {
let (hc, dim) = (cfg.hc_mult, cfg.dim);
let mut mixes = vec![0.0f32; hc];
hc_mixes(state, hc_fn, hc, cfg.norm_eps, pool, &mut mixes);
let mut pre = vec![0.0f32; hc];
hc_head_pre(&mixes, hc_scale, hc_base, hc, cfg.hc_eps, &mut pre);
hc_fold(state, &pre, hc, dim, out);
}
pub struct Dsv4Layer {
pub attn_norm: Vec<f32>,
pub ffn_norm: Vec<f32>,
pub wq_a: crate::qtensor::QTensor,
pub q_norm: Vec<f32>,
pub wq_b: crate::qtensor::QTensor,
pub wkv: crate::qtensor::QTensor,
pub kv_norm: Vec<f32>,
pub wo_a: crate::qtensor::QTensor,
pub wo_b: crate::qtensor::QTensor,
pub attn_sink: Vec<f32>,
pub compressor: Option<Dsv4Compressor>,
pub indexer: Option<Dsv4Indexer>,
pub hc_attn_fn: Vec<f32>,
pub hc_attn_base: Vec<f32>,
pub hc_attn_scale: [f32; 3],
pub hc_ffn_fn: Vec<f32>,
pub hc_ffn_base: Vec<f32>,
pub hc_ffn_scale: [f32; 3],
pub gate: crate::qtensor::QTensor,
pub gate_bias: Option<Vec<f32>>,
pub tid2eid: Option<Vec<f32>>,
pub experts: Vec<Dsv4Expert>,
pub shared: Dsv4Expert,
pub mask: Option<Vec<bool>>,
}
pub struct Dsv4Expert {
pub w1: crate::qtensor::QTensor,
pub w2: crate::qtensor::QTensor,
pub w3: crate::qtensor::QTensor,
}
pub struct Dsv4Compressor {
pub wkv: crate::qtensor::QTensor,
pub wgate: crate::qtensor::QTensor,
pub norm: Vec<f32>,
pub ape: Vec<f32>,
pub ratio: usize,
pub overlap: bool,
}
pub struct Dsv4Indexer {
pub wq_b: crate::qtensor::QTensor,
pub weights_proj: crate::qtensor::QTensor,
pub compressor: Dsv4Compressor,
}
pub struct Dsv4Globals {
pub inv_freq_compress: Vec<f32>,
pub inv_freq_window: Vec<f32>,
pub embed: crate::qtensor::QTensor,
pub norm: Vec<f32>,
pub head: crate::qtensor::QTensor,
pub hc_head_fn: Vec<f32>,
pub hc_head_base: Vec<f32>,
pub hc_head_scale: f32,
}
pub struct Dsv4State {
pub window: Vec<Vec<f32>>,
pub compressed: Vec<Vec<f32>>,
pub index_kv: Vec<Vec<f32>>,
pub pending_kv: Vec<Vec<f32>>,
pub pending_score: Vec<Vec<f32>>,
pub prev_kv: Vec<Vec<f32>>,
pub prev_score: Vec<Vec<f32>>,
pub pending_ix_kv: Vec<Vec<f32>>,
pub pending_ix_score: Vec<Vec<f32>>,
pub prev_ix_kv: Vec<Vec<f32>>,
pub prev_ix_score: Vec<Vec<f32>>,
pub pos: usize,
pub kv_id: u64,
}
impl Dsv4State {
pub fn new(layers: usize) -> Self {
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT: AtomicU64 = AtomicU64::new(1);
Self {
kv_id: NEXT.fetch_add(1, Ordering::Relaxed),
window: vec![Vec::new(); layers],
compressed: vec![Vec::new(); layers],
index_kv: vec![Vec::new(); layers],
pending_kv: vec![Vec::new(); layers],
pending_score: vec![Vec::new(); layers],
prev_kv: vec![Vec::new(); layers],
prev_score: vec![Vec::new(); layers],
pending_ix_kv: vec![Vec::new(); layers],
pending_ix_score: vec![Vec::new(); layers],
prev_ix_kv: vec![Vec::new(); layers],
prev_ix_score: vec![Vec::new(); layers],
pos: 0,
}
}
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
fn compressor_step(
cp: &Dsv4Compressor,
hidden: &[f32],
pos: usize,
rd: usize,
norm_eps: f32,
inv_freq: &[f32],
pool: Option<&crate::pool::Pool>,
pending_kv: &mut Vec<f32>,
pending_score: &mut Vec<f32>,
prev_kv: &mut Vec<f32>,
prev_score: &mut Vec<f32>,
) -> Option<Vec<f32>> {
let width = cp.wkv.rows();
let ew = if cp.overlap { width / 2 } else { width };
let mut ckv = vec![0.0f32; width];
let mut cscore = vec![0.0f32; width];
cp.wkv.matvec(hidden, &mut ckv, pool);
cp.wgate.matvec(hidden, &mut cscore, pool);
if cp.overlap {
let slot = pos % cp.ratio;
for (c, a) in cscore
.iter_mut()
.zip(&cp.ape[slot * width..(slot + 1) * width])
{
*c += a;
}
}
pending_kv.extend_from_slice(&ckv);
pending_score.extend_from_slice(&cscore);
if pending_kv.len() / width < cp.ratio {
return None;
}
let mut folded = vec![0.0f32; ew];
if cp.overlap {
compress_window_overlap(
prev_kv,
prev_score,
pending_kv,
pending_score,
cp.ratio,
ew,
&mut folded,
);
*prev_kv = std::mem::take(pending_kv);
*prev_score = std::mem::take(pending_score);
} else {
compress_window(
pending_kv,
pending_score,
&cp.ape,
cp.ratio,
width,
&mut folded,
);
}
rms_weighted(&mut folded, &cp.norm, norm_eps);
rope_tail(&mut folded, inv_freq, pos + 1 - cp.ratio, rd, false);
pending_kv.clear();
pending_score.clear();
Some(folded)
}
pub(crate) mod prof {
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
pub static ATTN_NS: AtomicU64 = AtomicU64::new(0);
pub static MOE_NS: AtomicU64 = AtomicU64::new(0);
pub static CALLS: AtomicU64 = AtomicU64::new(0);
pub static HC_NS: AtomicU64 = AtomicU64::new(0);
pub static HEAD_NS: AtomicU64 = AtomicU64::new(0);
pub static ALL_NS: AtomicU64 = AtomicU64::new(0);
pub static TOKENS: AtomicU64 = AtomicU64::new(0);
pub fn note_layer(li: usize) {
CALLS.fetch_add(1, Ordering::Relaxed);
if li == 0 {
TOKENS.fetch_add(1, Ordering::Relaxed);
}
}
static REPORT: AtomicBool = AtomicBool::new(false);
pub fn on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_DSV4_PROFILE").is_ok_and(|v| v != "0"))
}
pub fn report() {
if !on() || REPORT.swap(true, Ordering::Relaxed) {
return;
}
let calls = CALLS.load(Ordering::Relaxed).max(1);
let toks = TOKENS.load(Ordering::Relaxed).max(1);
let (a, m) = (
ATTN_NS.load(Ordering::Relaxed) as f64 / 1e6,
MOE_NS.load(Ordering::Relaxed) as f64 / 1e6,
);
let all = ALL_NS.load(Ordering::Relaxed) as f64 / 1e6;
let hc = HC_NS.load(Ordering::Relaxed) as f64 / 1e6;
let hd = HEAD_NS.load(Ordering::Relaxed) as f64 / 1e6;
eprintln!(
"[dsv4-профиль] {calls} вызовов слоя за {toks} токенов | \
на токен: внимание {:.0} мс, MoE {:.0} мс, гипер-связи+нормы {:.0} мс, \
голова {:.0} мс | на вызов: внимание {:.2}, MoE {:.2}, связи {:.2}",
a / toks as f64,
m / toks as f64,
hc / toks as f64,
hd / toks as f64,
a / calls as f64,
m / calls as f64,
hc / calls as f64,
);
eprintln!(
"[dsv4-профиль] весь проход {:.0} мс на токен; вне счётчиков {:.0} мс",
all / toks as f64,
(all - a - m - hd) / toks as f64,
);
#[cfg(feature = "gpu")]
{
let ae = crate::gpu_wgpu::ATT_ENC_NS.load(Ordering::Relaxed) as f64 / 1e6;
let aw = crate::gpu_wgpu::ATT_WAIT_NS.load(Ordering::Relaxed) as f64 / 1e6;
if ae + aw > 0.0 {
eprintln!(
"[dsv4-профиль] кадр внимания на вызов: кодирование {:.2} мс, \
отправка и ожидание {:.2} мс",
ae / calls as f64,
aw / calls as f64,
);
}
let e = crate::gpu_wgpu::MOE_ENC_NS.load(Ordering::Relaxed) as f64 / 1e6;
let wt = crate::gpu_wgpu::MOE_WAIT_NS.load(Ordering::Relaxed) as f64 / 1e6;
if e + wt > 0.0 {
eprintln!(
"[dsv4-профиль] кадр MoE на вызов: кодирование {:.2} мс, \
отправка и ожидание {:.2} мс",
e / calls as f64,
wt / calls as f64,
);
}
}
}
}
pub fn profile_report() {
prof::report();
}
fn gpu_attn_enabled() -> bool {
#[cfg(feature = "gpu")]
{
use std::sync::OnceLock;
static ON: OnceLock<bool> = OnceLock::new();
*ON.get_or_init(|| {
let want = std::env::var("CMF_DSV4_GPU_ATTN")
.map(|v| v != "0")
.unwrap_or(false);
let have = want && crate::gpu::backend_available();
if want && !have {
tracing::warn!(
"CMF_DSV4_GPU_ATTN задан, но устройства нет — блок внимания остаётся на CPU. Проверьте CMF_GPU=wgpu и Vulkan-ICD."
);
}
if std::env::var("CMF_DSV4_FRAME_DEBUG").is_ok() {
eprintln!("кадр dsv4: запрошен={want} доступен={have}");
}
have
})
}
#[cfg(not(feature = "gpu"))]
{
false
}
}
#[cfg(feature = "gpu")]
#[allow(clippy::too_many_arguments)]
fn attn_frame(
l: &Dsv4Layer,
cfg: &Dsv4Cfg,
st: &Dsv4State,
li: usize,
qn: &[f32],
idxs: &[usize],
inv_freq: &[f32],
pos: usize,
win_len: usize,
scale: f32,
out: &mut [f32],
) -> bool {
let hd = cfg.head_dim;
let (Some(wq_a), Some(wq_b), Some(wo_a), Some(wo_b)) = (
l.wq_a.model_idx(),
l.wq_b.model_idx(),
l.wo_a.model_idx(),
l.wo_b.model_idx(),
) else {
return false;
};
let Some(model) = l.wq_b.model_arc() else {
return false;
};
let n_comp = st.compressed[li].len() / hd;
let cap = (cfg.window + n_comp.next_power_of_two().max(64)) * hd;
let kv_id = st.kv_id;
if !crate::gpu_wgpu::dsv4_cache_write(kv_id, li, 0, &st.window[li], cap) {
return false;
}
if n_comp > 0
&& !crate::gpu_wgpu::dsv4_cache_write(
kv_id,
li,
cfg.window * hd,
&st.compressed[li],
cap,
)
{
return false;
}
let idx32: Vec<u32> = idxs
.iter()
.map(|&p| {
if p < win_len {
p as u32
} else {
(cfg.window + (p - win_len)) as u32
}
})
.collect();
let w = crate::gpu_wgpu::Dsv4AttnW {
wq_a,
wq_b,
wo_a,
wo_b,
q_norm: &l.q_norm,
sink: &l.attn_sink,
};
let g = crate::gpu_wgpu::Dsv4AttnGeom {
dim: cfg.dim,
nh: cfg.n_heads,
hd,
rd: cfg.rope_head_dim,
q_lora: cfg.q_lora_rank,
o_lora: cfg.o_lora_rank,
o_groups: cfg.o_groups,
eps: cfg.norm_eps,
scale,
};
crate::gpu_wgpu::dsv4_attn_frame(
&model,
&w,
g,
&[],
Some(qn),
kv_id,
li,
&idx32,
inv_freq,
pos,
out,
)
}
#[derive(Default)]
pub struct AttnPrep {
pub qr: Vec<f32>,
pub idxs: Vec<usize>,
pub win_len: usize,
}
#[allow(clippy::too_many_arguments)]
pub fn attention_step(
hidden: &[f32],
l: &Dsv4Layer,
cfg: &Dsv4Cfg,
st: &mut Dsv4State,
li: usize,
inv_freq: &[f32],
pool: Option<&crate::pool::Pool>,
prep_out: Option<&mut AttnPrep>,
out: &mut [f32],
) {
let _t0 = prof::on().then(std::time::Instant::now);
let _guard = scopeguard_attn(_t0);
let (hd, rd) = (cfg.head_dim, cfg.rope_head_dim);
let pos = st.pos;
if std::env::var("CMF_FREQ_DEBUG").is_ok() && li == 0 && pos == 0 {
eprintln!(
" [порт] rd={rd} частот={} inv_freq[0..4]={:?}",
inv_freq.len(),
&inv_freq[..4.min(inv_freq.len())]
);
}
let mut qr = vec![0.0f32; cfg.q_lora_rank];
l.wq_a.matvec(hidden, &mut qr, pool);
rms_weighted(&mut qr, &l.q_norm, cfg.norm_eps);
let on_gpu = gpu_attn_enabled();
let mut kv = vec![0.0f32; hd];
l.wkv.matvec(hidden, &mut kv, pool);
rms_weighted(&mut kv, &l.kv_norm, cfg.norm_eps);
rope_tail(&mut kv, inv_freq, pos, rd, false);
if let Some(cp) = &l.compressor {
let mut pk = std::mem::take(&mut st.pending_kv[li]);
let mut ps = std::mem::take(&mut st.pending_score[li]);
let mut qk = std::mem::take(&mut st.prev_kv[li]);
let mut qs = std::mem::take(&mut st.prev_score[li]);
let entry = compressor_step(
cp,
hidden,
pos,
rd,
cfg.norm_eps,
inv_freq,
pool,
&mut pk,
&mut ps,
&mut qk,
&mut qs,
);
st.pending_kv[li] = pk;
st.pending_score[li] = ps;
st.prev_kv[li] = qk;
st.prev_score[li] = qs;
if let Some(e) = entry {
st.compressed[li].extend_from_slice(&e);
}
}
if let Some(ix) = &l.indexer {
let mut pk = std::mem::take(&mut st.pending_ix_kv[li]);
let mut ps = std::mem::take(&mut st.pending_ix_score[li]);
let mut qk = std::mem::take(&mut st.prev_ix_kv[li]);
let mut qs = std::mem::take(&mut st.prev_ix_score[li]);
let entry = compressor_step(
&ix.compressor,
hidden,
pos,
rd,
cfg.norm_eps,
inv_freq,
pool,
&mut pk,
&mut ps,
&mut qk,
&mut qs,
);
st.pending_ix_kv[li] = pk;
st.pending_ix_score[li] = ps;
st.prev_ix_kv[li] = qk;
st.prev_ix_score[li] = qs;
if let Some(e) = entry {
st.index_kv[li].extend_from_slice(&e);
}
}
st.window[li].extend_from_slice(&kv);
let cap = cfg.window * hd;
if st.window[li].len() > cap {
let drop = st.window[li].len() - cap;
st.window[li].drain(..drop);
}
let win_len = st.window[li].len() / hd;
let n_pos = win_len + st.compressed[li].len() / hd;
let mut idxs: Vec<usize> = (0..win_len).collect();
if !st.compressed[li].is_empty() && !no_compressed() {
let n_comp = st.compressed[li].len() / hd;
match &l.indexer {
Some(ix) => {
let ih = ix.weights_proj.rows();
let idim = ix.wq_b.rows() / ih.max(1);
let mut qi = vec![0.0f32; ix.wq_b.rows()];
ix.wq_b.matvec(&qr, &mut qi, pool);
for h in 0..ih {
rope_tail(&mut qi[h * idim..(h + 1) * idim], inv_freq, pos, rd, false);
}
let mut hw = vec![0.0f32; ih];
ix.weights_proj.matvec(hidden, &mut hw, pool);
let sc_factor = (idim as f32).powf(-0.5) * (ih as f32).powf(-0.5);
for w in hw.iter_mut() {
*w *= sc_factor;
}
let n_ix = st.index_kv[li].len() / idim.max(1);
let mut sc = Vec::new();
index_scores(
&qi,
&st.index_kv[li],
&hw,
ih,
idim,
n_ix.min(n_comp),
n_ix.min(n_comp),
pool,
&mut sc,
);
let mut picked = Vec::new();
top_k_positions(&sc, cfg.index_topk, &mut picked);
idxs.extend(picked.into_iter().map(|p| win_len + p));
}
None => idxs.extend((0..n_comp).map(|p| win_len + p)),
}
}
debug_assert!(idxs.iter().all(|&p| p < n_pos));
if let Some(p) = prep_out {
p.qr = qr;
p.idxs = idxs;
p.win_len = win_len;
return;
}
let scale = (hd as f32).powf(-0.5);
#[cfg(feature = "gpu")]
if on_gpu
&& attn_frame(
l, cfg, st, li, &qr, &idxs, inv_freq, pos, win_len, scale, out,
)
{
return;
}
let mut q = vec![0.0f32; cfg.n_heads * hd];
l.wq_b.matvec(&qr, &mut q, pool);
for h in 0..cfg.n_heads {
let head = &mut q[h * hd..(h + 1) * hd];
rms_inplace(head, cfg.norm_eps);
rope_tail(head, inv_freq, pos, rd, false);
}
let mut cache: Vec<f32> = st.window[li].clone();
cache.extend_from_slice(&st.compressed[li]);
let mut attn = vec![0.0f32; cfg.n_heads * hd];
for h in 0..cfg.n_heads {
let qh = &q[h * hd..(h + 1) * hd];
let mut oh = vec![0.0f32; hd];
sparse_attend(qh, &cache, &idxs, l.attn_sink[h], scale, hd, &mut oh);
rope_tail(&mut oh, inv_freq, pos, rd, true);
attn[h * hd..(h + 1) * hd].copy_from_slice(&oh);
}
o_project(
&attn,
&|r, x, sc| l.wo_a.row_dot(r, x, sc),
l.wo_a.cols(),
&|mid, dst| l.wo_b.matvec(mid, dst, pool),
cfg.o_groups,
cfg.o_lora_rank,
pool,
out,
);
}
pub fn rms_weighted(v: &mut [f32], w: &[f32], eps: f32) {
let ms = v.iter().map(|x| x * x).sum::<f32>() / v.len() as f32;
let inv = 1.0 / (ms + eps).sqrt();
for (x, g) in v.iter_mut().zip(w) {
*x = *x * inv * g;
}
}
thread_local! {
static ROUTE_COUNTS: std::cell::RefCell<Vec<Vec<u64>>> =
const { std::cell::RefCell::new(Vec::new()) };
}
fn route_stats_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_MOE_STATS").is_ok())
}
fn record_route(li: usize, n_layers_hint: usize, n_experts: usize, idx: &[usize]) {
ROUTE_COUNTS.with(|c| {
let mut c = c.borrow_mut();
if c.len() <= li.max(n_layers_hint) {
c.resize(li.max(n_layers_hint) + 1, Vec::new());
}
let row = &mut c[li];
if row.len() < n_experts {
row.resize(n_experts, 0);
}
for &e in idx {
if e < row.len() {
row[e] += 1;
}
}
});
}
pub fn take_route_counts() -> Vec<Vec<u64>> {
ROUTE_COUNTS.with(|c| std::mem::take(&mut *c.borrow_mut()))
}
struct Charge(Option<std::time::Instant>, &'static std::sync::atomic::AtomicU64);
impl Drop for Charge {
fn drop(&mut self) {
if let Some(t) = self.0 {
self.1.fetch_add(
t.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
}
}
}
fn scopeguard_attn(t: Option<std::time::Instant>) -> Charge {
Charge(t, &prof::ATTN_NS)
}
fn scopeguard_moe(t: Option<std::time::Instant>, li: usize) -> Charge {
if t.is_some() {
prof::note_layer(li);
}
Charge(t, &prof::MOE_NS)
}
#[cfg(feature = "gpu")]
#[allow(clippy::too_many_arguments)]
fn dsv4_layer_loop(
state: &mut [f32],
layers: &[Dsv4Layer],
g: &Dsv4Globals,
cfg: &Dsv4Cfg,
st: &mut Dsv4State,
token_id: u32,
inv_freq: &[f32],
pool: Option<&crate::pool::Pool>,
scratch: &mut HcScratch,
) -> bool {
let dim = cfg.dim;
let freqs_of = |l: &Dsv4Layer| -> &[f32] {
let f = if l.compressor.is_some() {
&g.inv_freq_compress
} else {
&g.inv_freq_window
};
if f.is_empty() { inv_freq } else { f.as_slice() }
};
let mut on_dev = vec![false; layers.len()];
for (li, l) in layers.iter().enumerate() {
let Some(pk) = pack_for(l, cfg, li) else {
return false;
};
if l.wq_a.model_idx().is_none()
|| l.wq_b.model_idx().is_none()
|| l.wo_a.model_idx().is_none()
|| l.wo_b.model_idx().is_none()
{
return false;
}
let Some(model) = l.experts.first().and_then(|e| e.w1.model_arc()) else {
return false;
};
let gu_q2 = l
.experts
.first()
.is_some_and(|e| e.w1.model_dtype() == Some(cortiq_core::TensorDtype::Q2TiledP));
let attn_ok = [
l.wq_a.model_idx(),
l.wq_b.model_idx(),
l.wo_a.model_idx(),
l.wo_b.model_idx(),
]
.into_iter()
.flatten()
.all(|i| crate::gpu_wgpu::dsv4_weight_ready(&model, i));
on_dev[li] = attn_ok
&& pk.globals.len() == cfg.n_routed_experts
&& crate::gpu_wgpu::dsv4_experts_ready(&model, &pk.tensors, cfg.moe_inter, dim, gu_q2);
}
if !on_dev.iter().any(|&x| x) {
return false;
}
let (mut folded, post0, comb0) = hc_fold_norm(
state,
&layers[0].hc_attn_fn,
&layers[0].hc_attn_scale,
&layers[0].hc_attn_base,
&layers[0].attn_norm,
cfg,
pool,
);
if !crate::gpu_wgpu::dsv4_state_write(state)
|| !crate::gpu_wgpu::dsv4_hc_write(&post0, &comb0)
{
return false;
}
let mut sink_out = vec![0.0f32; dim];
for (li, l) in layers.iter().enumerate() {
if !on_dev[li] {
if !crate::gpu_wgpu::dsv4_state_read(state) {
return false;
}
let freqs = freqs_of(l);
hc_block(
state,
&l.hc_attn_fn,
&l.hc_attn_scale,
&l.hc_attn_base,
&l.attn_norm,
cfg,
scratch,
pool,
|f, o| attention_step(f, l, cfg, st, li, freqs, pool, None, o),
);
hc_block(
state,
&l.hc_ffn_fn,
&l.hc_ffn_scale,
&l.hc_ffn_base,
&l.ffn_norm,
cfg,
scratch,
pool,
|f, o| moe_step(f, l, cfg, token_id, li, pool, o),
);
if let Some(n) = layers.get(li + 1) {
let (f, p2, c2) = hc_fold_norm(
state,
&n.hc_attn_fn,
&n.hc_attn_scale,
&n.hc_attn_base,
&n.attn_norm,
cfg,
pool,
);
folded = f;
if !crate::gpu_wgpu::dsv4_hc_write(&p2, &c2) {
return false;
}
}
if !crate::gpu_wgpu::dsv4_state_write(state) {
return false;
}
continue;
}
let mut prep = AttnPrep::default();
attention_step(
&folded,
l,
cfg,
st,
li,
freqs_of(l),
pool,
Some(&mut prep),
&mut sink_out,
);
let hd = cfg.head_dim;
let n_comp = st.compressed[li].len() / hd;
let cap = (cfg.window + n_comp.next_power_of_two().max(64)) * hd;
let kv_id = st.kv_id;
if !crate::gpu_wgpu::dsv4_cache_write(kv_id, li, 0, &st.window[li], cap)
|| (n_comp > 0
&& !crate::gpu_wgpu::dsv4_cache_write(
kv_id,
li,
cfg.window * hd,
&st.compressed[li],
cap,
))
{
return false;
}
let idx32: Vec<u32> = prep
.idxs
.iter()
.map(|&p| {
if p < prep.win_len {
p as u32
} else {
(cfg.window + (p - prep.win_len)) as u32
}
})
.collect();
let Some(pk) = pack_for(l, cfg, li) else {
return false;
};
let (Some(wq_a), Some(wq_b), Some(wo_a), Some(wo_b)) = (
l.wq_a.model_idx(),
l.wq_b.model_idx(),
l.wo_a.model_idx(),
l.wo_b.model_idx(),
) else {
return false;
};
let Some(model) = l.experts.first().and_then(|e| e.w1.model_arc()) else {
return false;
};
let forced: Option<Vec<usize>> = l.tid2eid.as_ref().and_then(|tbl| {
let v: Vec<usize> = hash_route(tbl, cfg.vocab, cfg.top_k, token_id)
.into_iter()
.map(|gi| pk.to_slot[gi])
.collect();
if v.iter().any(|&x| x == usize::MAX) {
None
} else {
Some(v)
}
});
if l.tid2eid.is_some() && forced.is_none() {
return false;
}
let bias: Option<Vec<f32>> = l
.gate_bias
.as_deref()
.map(|b| pk.globals.iter().map(|&gi| b[gi]).collect());
let nxt = layers.get(li + 1);
let w = crate::gpu_wgpu::Dsv4LayerW {
attn: crate::gpu_wgpu::Dsv4AttnW {
wq_a,
wq_b,
wo_a,
wo_b,
q_norm: &l.q_norm,
sink: &l.attn_sink,
},
moe: crate::gpu_wgpu::Dsv4MoeW {
experts: &pk.tensors,
logits: &[],
bias: bias.as_deref(),
forced: forced.as_deref(),
remap: None,
},
hc_ffn_fn: &l.hc_ffn_fn,
hc_ffn_scale: &l.hc_ffn_scale,
hc_ffn_base: &l.hc_ffn_base,
hc_next_fn: nxt.map(|n| n.hc_attn_fn.as_slice()),
hc_next_scale: nxt.map_or(&l.hc_attn_scale, |n| &n.hc_attn_scale),
hc_next_base: nxt.map_or(&l.hc_attn_base, |n| n.hc_attn_base.as_slice()),
ffn_norm: &l.ffn_norm,
next_norm: nxt.map_or(&l.attn_norm, |n| n.attn_norm.as_slice()),
next_q_norm: nxt.map_or(&l.q_norm, |n| n.q_norm.as_slice()),
next_wq_a: nxt.and_then(|n| n.wq_a.model_idx()),
router: &pk.router,
};
let geom = crate::gpu_wgpu::Dsv4LayerGeom {
attn: crate::gpu_wgpu::Dsv4AttnGeom {
dim,
nh: cfg.n_heads,
hd,
rd: cfg.rope_head_dim,
q_lora: cfg.q_lora_rank,
o_lora: cfg.o_lora_rank,
o_groups: cfg.o_groups,
eps: cfg.norm_eps,
scale: (hd as f32).powf(-0.5),
},
moe: crate::gpu_wgpu::Dsv4MoeGeom {
hidden: dim,
inter: cfg.moe_inter,
top_k: cfg.top_k,
route_scale: cfg.route_scale,
swiglu_limit: cfg.swiglu_limit,
gu_q2: l.experts.first().is_some_and(|e| {
e.w1.model_dtype() == Some(cortiq_core::TensorDtype::Q2TiledP)
}),
},
hc: cfg.hc_mult,
hc_eps: cfg.hc_eps,
sinkhorn_iters: cfg.hc_sinkhorn_iters,
};
let mut next = vec![0.0f32; dim];
if !crate::gpu_wgpu::dsv4_layer_frame(
&model,
&w,
geom,
kv_id,
li,
Some(&prep.qr),
&idx32,
freqs_of(l),
st.pos,
&mut next,
) {
return false;
}
folded = next;
}
crate::gpu_wgpu::dsv4_state_read(state)
}
#[cfg(feature = "gpu")]
#[allow(clippy::too_many_arguments)]
fn hc_fold_norm(
state: &[f32],
hc_fn: &[f32],
hc_scale: &[f32; 3],
hc_base: &[f32],
norm_w: &[f32],
cfg: &Dsv4Cfg,
pool: Option<&crate::pool::Pool>,
) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let (hc, dim) = (cfg.hc_mult, cfg.dim);
let mix_hc = (2 + hc) * hc;
let mut mixes = vec![0.0f32; mix_hc];
hc_mixes(state, hc_fn, mix_hc, cfg.norm_eps, pool, &mut mixes);
let mut pre = vec![0.0f32; hc];
let mut post = vec![0.0f32; hc];
let mut comb = vec![0.0f32; hc * hc];
hc_split_sinkhorn(
&mixes,
hc_scale,
hc_base,
hc,
cfg.hc_sinkhorn_iters,
cfg.hc_eps,
&mut pre,
&mut post,
&mut comb,
);
let mut folded = vec![0.0f32; dim];
hc_fold(state, &pre, hc, dim, &mut folded);
rms_weighted(&mut folded, norm_w, cfg.norm_eps);
(folded, post, comb)
}
#[cfg(feature = "gpu")]
fn gpu_layer_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("CMF_DSV4_GPU_LAYER").is_ok_and(|v| v != "0")
&& crate::gpu::backend_available()
})
}
#[cfg(feature = "gpu")]
struct Pack {
router: Vec<f32>,
to_slot: Vec<usize>,
remap: Vec<u32>,
globals: Vec<usize>,
tensors: Vec<(usize, usize, usize)>,
}
#[cfg(feature = "gpu")]
fn pack_for(l: &Dsv4Layer, cfg: &Dsv4Cfg, li: usize) -> Option<std::sync::Arc<Pack>> {
use std::collections::HashMap;
use std::sync::{Arc, Mutex, OnceLock};
static CACHE: OnceLock<Mutex<HashMap<usize, Option<Arc<Pack>>>>> = OnceLock::new();
let cache = CACHE.get_or_init(|| Mutex::new(HashMap::new()));
if let Some(v) = cache.lock().unwrap().get(&li) {
return v.clone();
}
let build = || -> Option<Arc<Pack>> {
let mut to_slot = vec![usize::MAX; cfg.n_routed_experts];
let mut globals = Vec::new();
let mut tensors = Vec::new();
let idx3 = |e: &Dsv4Expert| -> Option<(usize, usize, usize)> {
Some((e.w1.model_idx()?, e.w3.model_idx()?, e.w2.model_idx()?))
};
let gu_q2 = l
.experts
.first()
.is_some_and(|e| e.w1.model_dtype() == Some(cortiq_core::TensorDtype::Q2TiledP));
if let Some(n) = std::env::var("CMF_DSV4_PACK_MAX")
.ok()
.and_then(|v| v.parse::<usize>().ok())
{
let mut to_slot = vec![usize::MAX; cfg.n_routed_experts];
let mut globals = Vec::new();
let mut tensors = Vec::new();
for (gi, e) in l.experts.iter().enumerate().take(n) {
to_slot[gi] = globals.len();
globals.push(gi);
tensors.push(idx3(e)?);
}
tensors.push(idx3(&l.shared)?);
let (rows, cols) = (l.gate.rows(), l.gate.cols());
let mut router = vec![0.0f32; rows * cols];
for r in 0..rows {
l.gate.row_f32(r, &mut router[r * cols..(r + 1) * cols]);
}
let remap: Vec<u32> = to_slot
.iter()
.map(|&sl| if sl == usize::MAX { u32::MAX } else { sl as u32 })
.collect();
return Some(Arc::new(Pack {
router,
to_slot,
remap,
globals,
tensors,
}));
}
let room = if std::env::var("CMF_DSV4_COLD_CPU").is_ok_and(|v| v != "0") {
crate::gpu_wgpu::dsv4_experts_fit(cfg.moe_inter, cfg.dim, gu_q2).saturating_sub(1)
} else {
usize::MAX
};
for (gi, e) in l.experts.iter().enumerate() {
if l.mask.as_deref().is_some_and(|m| !m.get(gi).copied().unwrap_or(true)) {
continue;
}
if globals.len() >= room {
break;
}
to_slot[gi] = globals.len();
globals.push(gi);
match idx3(e) {
Some(t) => tensors.push(t),
None => {
if std::env::var("CMF_DSV4_FRAME_DEBUG").is_ok() {
eprintln!("слой {li}: эксперт {gi} без индексов в каталоге");
}
return None;
}
}
}
if globals.is_empty() {
tracing::warn!("слой {li}: маска не оставила ни одного эксперта");
return None;
}
tensors.push(idx3(&l.shared)?); let (rows, cols) = (l.gate.rows(), l.gate.cols());
let mut router = vec![0.0f32; rows * cols];
for r in 0..rows {
l.gate.row_f32(r, &mut router[r * cols..(r + 1) * cols]);
}
let remap: Vec<u32> = to_slot
.iter()
.map(|&sl| if sl == usize::MAX { u32::MAX } else { sl as u32 })
.collect();
Some(Arc::new(Pack {
router,
to_slot,
remap,
globals,
tensors,
}))
};
let v = build();
cache.lock().unwrap().insert(li, v.clone());
v
}
#[cfg(feature = "gpu")]
fn moe_frame(
hidden: &[f32],
l: &Dsv4Layer,
cfg: &Dsv4Cfg,
li: usize,
logits: &[f32],
forced: Option<&[usize]>,
pool: Option<&crate::pool::Pool>,
out: &mut [f32],
) -> bool {
macro_rules! no {
($($t:tt)*) => {{
if std::env::var("CMF_DSV4_FRAME_DEBUG").is_ok() {
eprintln!("кадр MoE отклонён: {}", format_args!($($t)*));
}
return false;
}};
}
let Some(pk) = pack_for(l, cfg, li) else {
no!("слой {li}: упаковка экспертов не построена");
};
let Some(model) = l.experts.first().and_then(|e| e.w1.model_arc()) else {
no!("слой {li}: эксперты не отображены из файла");
};
let fpack: Option<Vec<usize>> = match forced {
Some(f) => {
let v: Vec<usize> = f.iter().map(|&g| pk.to_slot[g]).collect();
if v.iter().any(|&s| s == usize::MAX) {
no!("слой {li}: хеш-слой называет эксперта вне упаковки");
}
Some(v)
}
None => None,
};
let subset = pk.globals.len() < cfg.n_routed_experts;
if std::env::var("CMF_DSV4_MOE_CHECK").is_ok() {
eprintln!(
"[упаковка] слой {li}: globals={} n_routed={} subset={subset} remap.len={}",
pk.globals.len(),
cfg.n_routed_experts,
pk.remap.len()
);
}
let lg: Vec<f32> = if subset {
logits.to_vec()
} else {
pk.globals.iter().map(|&g| logits[g]).collect()
};
let bias: Option<Vec<f32>> = l.gate_bias.as_deref().map(|b| {
if subset {
b.to_vec()
} else {
pk.globals.iter().map(|&g| b[g]).collect()
}
});
let w = crate::gpu_wgpu::Dsv4MoeW {
experts: &pk.tensors,
logits: &lg,
bias: bias.as_deref(),
forced: fpack.as_deref(),
remap: if subset { Some(&pk.remap) } else { None },
};
let g = crate::gpu_wgpu::Dsv4MoeGeom {
hidden: cfg.dim,
inter: cfg.moe_inter,
top_k: cfg.top_k,
route_scale: cfg.route_scale,
swiglu_limit: cfg.swiglu_limit,
gu_q2: l.experts.first().is_some_and(|e| {
e.w1.model_dtype() == Some(cortiq_core::TensorDtype::Q2TiledP)
}),
};
let mut cold = Vec::new();
if !crate::gpu_wgpu::dsv4_moe_frame(&model, &w, g, hidden, &mut cold, out) {
return false;
}
if std::env::var("CMF_DSV4_MOE_CHECK").is_ok() {
let (mut ci, mut cw) = (Vec::new(), Vec::new());
route(
logits,
l.gate_bias.as_deref(),
cfg.top_k,
cfg.route_scale,
forced,
None,
&mut ci,
&mut cw,
);
eprintln!("[выбор CPU] слой {li}: {ci:?} веса {cw:?}");
let csum: f32 = cold.iter().map(|c| c.1).sum();
eprintln!(
"[холодные] слой {li}: вернулось {} из {} | сумма холодных {csum:.4} | \
route_scale {:.4} | {:?}",
cold.len(),
cfg.top_k,
cfg.route_scale,
&cold[..cold.len().min(3)]
);
}
let mut acc = vec![0.0f32; cfg.dim];
for &(gi, wt) in &cold {
let Some(exp) = l.experts.get(gi) else { continue };
run_expert(hidden, exp, cfg, wt, pool, &mut acc);
for (o, a) in out.iter_mut().zip(&acc) {
*o += a;
}
}
true
}
#[cfg(feature = "gpu")]
fn last_grew(now: u64) -> u64 {
use std::sync::atomic::{AtomicU64, Ordering};
static SEEN: AtomicU64 = AtomicU64::new(0);
let was = SEEN.load(Ordering::Relaxed);
if was != now {
SEEN.store(now, Ordering::Relaxed);
compressed_map().lock().unwrap().clear();
return u64::MAX; }
now
}
#[cfg(feature = "gpu")]
fn compressed_map() -> &'static std::sync::Mutex<std::collections::HashMap<(u64, usize), usize>> {
use std::collections::HashMap;
use std::sync::{Mutex, OnceLock};
static W: OnceLock<Mutex<HashMap<(u64, usize), usize>>> = OnceLock::new();
W.get_or_init(|| Mutex::new(HashMap::new()))
}
#[cfg(feature = "gpu")]
fn compressed_written(kv_id: u64, li: usize) -> usize {
compressed_map()
.lock()
.unwrap()
.get(&(kv_id, li))
.copied()
.unwrap_or(0)
}
#[cfg(feature = "gpu")]
fn note_compressed(kv_id: u64, li: usize, n: usize) {
compressed_map().lock().unwrap().insert((kv_id, li), n);
}
#[cfg(feature = "gpu")]
fn gpu_moe2_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("CMF_DSV4_GPU_MOE2").is_ok_and(|v| v != "0")
&& crate::gpu::backend_available()
})
}
pub fn moe_step(
hidden: &[f32],
l: &Dsv4Layer,
cfg: &Dsv4Cfg,
token_id: u32,
li: usize,
pool: Option<&crate::pool::Pool>,
out: &mut [f32],
) {
let _t0 = prof::on().then(std::time::Instant::now);
let _guard = scopeguard_moe(_t0, li);
let mut logits = vec![0.0f32; cfg.n_routed_experts];
l.gate.matvec(hidden, &mut logits, pool);
let (mut idx, mut w) = (Vec::new(), Vec::new());
route(
&logits,
l.gate_bias.as_deref(),
cfg.top_k,
cfg.route_scale,
l.tid2eid
.as_ref()
.map(|tbl| hash_route(tbl, cfg.vocab, cfg.top_k, token_id))
.as_deref(),
l.mask.as_deref(),
&mut idx,
&mut w,
);
if route_stats_on() {
record_route(li, 0, cfg.n_routed_experts, &idx);
}
#[cfg(feature = "gpu")]
if gpu_moe2_enabled() {
let forced = l
.tid2eid
.as_ref()
.map(|tbl| hash_route(tbl, cfg.vocab, cfg.top_k, token_id));
if moe_frame(hidden, l, cfg, li, &logits, forced.as_deref(), pool, out) {
if std::env::var("CMF_DSV4_MOE_CHECK").is_ok() {
let mut want = vec![0.0f32; out.len()];
let mut acc = vec![0.0f32; cfg.dim];
for (e, &ei) in idx.iter().enumerate() {
let Some(exp) = l.experts.get(ei) else { continue };
run_expert(hidden, exp, cfg, w.get(e).copied().unwrap_or(0.0), pool, &mut acc);
for (o, a) in want.iter_mut().zip(&acc) {
*o += a;
}
}
run_expert(hidden, &l.shared, cfg, 1.0, pool, &mut acc);
for (o, a) in want.iter_mut().zip(&acc) {
*o += a;
}
let num: f32 = want.iter().zip(out.iter()).map(|(a, b)| (a - b) * (a - b)).sum();
let den: f32 = want.iter().map(|a| a * a).sum::<f32>().max(1e-20);
let rel = (num / den).sqrt();
if rel > 1e-3 {
let packed = pack_for(l, cfg, li).map_or(0, |p| p.globals.len());
eprintln!(
"[кадр MoE] слой {li}: расхождение {rel:.3e} | выбрано {} | \
упаковано {packed} из {} | хеш={} | смещение={}",
idx.len(),
cfg.n_routed_experts,
l.tid2eid.is_some(),
l.gate_bias.is_some()
);
}
}
return;
}
}
if dump_path().is_some() {
PICKED.with(|p| {
let mut p = p.borrow_mut();
if p.len() <= li {
p.resize(li + 1, Vec::new());
}
p[li] = idx.clone();
});
}
fn gpu_moe_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_DSV4_GPU_MOE").is_ok_and(|v| v != "0"))
}
if gpu_moe_on() && crate::gpu::enabled_here() {
let mut jobs = Vec::with_capacity(idx.len() + 1);
let mut model_ref = None;
let mut ok = true;
for (e, &ei) in idx.iter().enumerate() {
let Some(exp) = l.experts.get(ei) else { continue };
ok &= crate::pipeline::moe_push_job_parts(
&exp.w1,
&exp.w3,
&exp.w2,
hidden,
w.get(e).copied().unwrap_or(0.0),
cfg.swiglu_limit,
&mut jobs,
&mut model_ref,
)
.is_some();
}
ok &= crate::pipeline::moe_push_job_parts(
&l.shared.w1,
&l.shared.w3,
&l.shared.w2,
hidden,
1.0,
cfg.swiglu_limit,
&mut jobs,
&mut model_ref,
)
.is_some();
if ok {
if let Some(m) = model_ref.as_ref() {
if crate::gpu::moe_block(m, &jobs, out) {
if std::env::var("CMF_DSV4_GPU_CHECK").is_ok() {
let mut want = vec![0.0f32; out.len()];
let mut acc = vec![0.0f32; cfg.dim];
for (e, &ei) in idx.iter().enumerate() {
let Some(exp) = l.experts.get(ei) else { continue };
run_expert(
hidden, exp, cfg,
w.get(e).copied().unwrap_or(0.0), pool, &mut acc,
);
for (o, a) in want.iter_mut().zip(&acc) {
*o += a;
}
}
run_expert(hidden, &l.shared, cfg, 1.0, pool, &mut acc);
for (o, a) in want.iter_mut().zip(&acc) {
*o += a;
}
let num: f32 = want
.iter()
.zip(out.iter())
.map(|(a, b)| (a - b) * (a - b))
.sum();
let den: f32 = want.iter().map(|a| a * a).sum::<f32>().max(1e-20);
eprintln!(
"[dsv4-gpu] слой {li}: расхождение {:.3e} | |CPU|={:.5} |GPU|={:.5} | экспертов {}",
(num / den).sqrt(),
den.sqrt(),
out.iter().map(|x| x * x).sum::<f32>().sqrt(),
jobs.len()
);
}
return;
}
}
}
}
out.fill(0.0);
let mut acc = vec![0.0f32; cfg.dim];
for (e, &ei) in idx.iter().enumerate() {
let Some(exp) = l.experts.get(ei) else {
continue;
};
run_expert(
hidden,
exp,
cfg,
w.get(e).copied().unwrap_or(0.0),
pool,
&mut acc,
);
for (o, a) in out.iter_mut().zip(&acc) {
*o += a;
}
}
run_expert(hidden, &l.shared, cfg, 1.0, pool, &mut acc);
for (o, a) in out.iter_mut().zip(&acc) {
*o += a;
}
}
fn run_expert(
x: &[f32],
e: &Dsv4Expert,
cfg: &Dsv4Cfg,
weight: f32,
pool: Option<&crate::pool::Pool>,
out: &mut [f32],
) {
expert_swiglu(
x,
&|src, dst| e.w1.matvec(src, dst, pool),
&|src, dst| e.w3.matvec(src, dst, pool),
&|src, dst| e.w2.matvec(src, dst, pool),
cfg.moe_inter,
weight,
cfg.swiglu_limit,
out,
);
}
fn no_compressed() -> bool {
static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*OFF.get_or_init(|| std::env::var("CMF_DSV4_NO_COMPRESSED").is_ok_and(|v| v != "0"))
}
fn trace_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_DSV4_TRACE").is_ok_and(|v| v != "0"))
}
fn rms_of(v: &[f32]) -> f32 {
if v.is_empty() {
return 0.0;
}
(v.iter().map(|x| x * x).sum::<f32>() / v.len() as f32).sqrt()
}
thread_local! {
static BODY: std::cell::RefCell<Vec<String>> = const { std::cell::RefCell::new(Vec::new()) };
static PICKED: std::cell::RefCell<Vec<Vec<usize>>> =
const { std::cell::RefCell::new(Vec::new()) };
}
fn dump_path() -> Option<&'static str> {
static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
P.get_or_init(|| std::env::var("CMF_DSV4_DUMP").ok())
.as_deref()
}
fn dump_line(json: &str) {
if let Some(p) = dump_path() {
use std::io::Write as _;
if let Ok(mut f) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(p)
{
let _ = writeln!(f, "{json}");
}
}
}
fn vec_json(v: &[f32]) -> String {
let mut s = String::with_capacity(v.len() * 9);
s.push('[');
for (i, x) in v.iter().enumerate() {
if i > 0 {
s.push(',');
}
s.push_str(&format!("{x:.6e}"));
}
s.push(']');
s
}
#[allow(clippy::too_many_arguments)]
pub fn forward_token(
g: &Dsv4Globals,
layers: &[Dsv4Layer],
cfg: &Dsv4Cfg,
st: &mut Dsv4State,
token_id: u32,
inv_freq: &[f32],
pool: Option<&crate::pool::Pool>,
logits: &mut Vec<f32>,
) {
let _t_all = prof::on().then(std::time::Instant::now);
let _all_guard = Charge(_t_all, &prof::ALL_NS);
let (hc, dim) = (cfg.hc_mult, cfg.dim);
let mut emb = vec![0.0f32; dim];
g.embed.row_f32(token_id as usize, &mut emb);
let mut state = vec![0.0f32; hc * dim];
for j in 0..hc {
state[j * dim..(j + 1) * dim].copy_from_slice(&emb);
}
let mut scratch = HcScratch::new(cfg);
let mut dump: Vec<String> = Vec::new();
if dump_path().is_some() {
dump.push(format!("\"embed\":{}", vec_json(&emb)));
PICKED.with(|p| p.borrow_mut().clear());
BODY.with(|b| b.borrow_mut().clear());
dump.push(",\"layers\":[".into());
}
if trace_on() {
eprintln!(
"[dsv4] tok={token_id} pos={} embed rms={:.5}",
st.pos,
rms_of(&emb)
);
}
#[cfg(feature = "gpu")]
let layer_frames = gpu_layer_enabled()
&& dsv4_layer_loop(
&mut state, layers, g, cfg, st, token_id, inv_freq, pool, &mut scratch,
);
#[cfg(not(feature = "gpu"))]
let layer_frames = false;
for (li, l) in layers.iter().enumerate() {
if layer_frames {
break;
}
hc_block(
&mut state,
&l.hc_attn_fn,
&l.hc_attn_scale,
&l.hc_attn_base,
&l.attn_norm,
cfg,
&mut scratch,
pool,
|folded, out| {
if dump_path().is_some() {
BODY.with(|b| b.borrow_mut().push(vec_json(folded)));
}
let freqs = if l.compressor.is_some() {
&g.inv_freq_compress
} else {
&g.inv_freq_window
};
let freqs = if freqs.is_empty() {
inv_freq
} else {
freqs.as_slice()
};
attention_step(folded, l, cfg, st, li, freqs, pool, None, out);
if dump_path().is_some() {
BODY.with(|b| b.borrow_mut().push(vec_json(out)));
}
},
);
if dump_path().is_some() {
dump.push(format!(
"{}{}",
if li == 0 { "" } else { "," },
vec_json(&state)
));
}
let _t_hc2 = prof::on().then(std::time::Instant::now);
hc_block(
&mut state,
&l.hc_ffn_fn,
&l.hc_ffn_scale,
&l.hc_ffn_base,
&l.ffn_norm,
cfg,
&mut scratch,
pool,
|folded, out| moe_step(folded, l, cfg, token_id, li, pool, out),
);
if let Some(t) = _t_hc2 {
prof::HC_NS.fetch_add(
t.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
}
if dump_path().is_some() {
dump.push(format!(",{}", vec_json(&state)));
}
if trace_on() && (st.pos % 64 == 0 || st.pos == 199) {
eprintln!(
"[dsv4] кеши слоя {li}: окно={} сжатых={} индекс={} (ratio={:?})",
st.window[li].len() / cfg.head_dim.max(1),
st.compressed[li].len() / cfg.head_dim.max(1),
st.index_kv[li].len().max(1) / 128,
l.compressor.as_ref().map(|c| (c.ratio, c.overlap)),
);
}
if trace_on() {
let bad = state.iter().filter(|v| !v.is_finite()).count();
eprintln!(
"[dsv4] layer {li:>2}: rms={:.5}{}",
rms_of(&state),
if bad > 0 {
format!(" NON-FINITE x{bad}")
} else {
String::new()
}
);
}
}
st.pos += 1;
let mut h = vec![0.0f32; dim];
hc_head_fold(
&state,
&g.hc_head_fn,
g.hc_head_scale,
&g.hc_head_base,
cfg,
pool,
&mut h,
);
let _t_head = prof::on().then(std::time::Instant::now);
rms_weighted(&mut h, &g.norm, cfg.norm_eps);
logits.clear();
logits.resize(g.head.rows(), 0.0);
g.head.matvec(&h, logits, pool);
if let Some(t) = _t_head {
prof::HEAD_NS.fetch_add(
t.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
}
if dump_path().is_some() {
dump.push("]".into());
let picked = PICKED.with(|p| {
p.borrow()
.iter()
.map(|v| {
format!(
"[{}]",
v.iter()
.map(|e| e.to_string())
.collect::<Vec<_>>()
.join(",")
)
})
.collect::<Vec<_>>()
.join(",")
});
dump.push(format!(",\"experts\":[{picked}]"));
let body = BODY.with(|b| b.borrow().join(","));
dump.push(format!(",\"attn_io\":[{body}]"));
dump_line(&format!(
"{{\"tok\":{token_id},\"pos\":{},{},\"head\":{},\"logits\":{}}}",
st.pos - 1,
dump.join(""),
vec_json(&h),
vec_json(logits)
));
}
if trace_on() {
let (mut top, mut best) = (0usize, f32::NEG_INFINITY);
for (i, &v) in logits.iter().enumerate() {
if v > best {
best = v;
top = i;
}
}
let lo = logits.iter().cloned().fold(f32::MAX, f32::min);
eprintln!(
"[dsv4] head: rms={:.5} logits[{}..{:.3}] argmax={top}",
rms_of(&h),
format_args!("{lo:.3}"),
best
);
}
}
pub fn load(
model: &std::sync::Arc<cortiq_core::CmfModel>,
cfg: &Dsv4Cfg,
n_layers: usize,
) -> Result<(Dsv4Globals, Vec<Dsv4Layer>), String> {
let q = |name: &str| -> Result<crate::qtensor::QTensor, String> {
crate::qtensor::QTensor::from_model(model, name)
};
let f = |name: &str| -> Result<Vec<f32>, String> {
crate::loader::load_f32(model, name, &crate::loader::Overlay::None)
};
let opt_f = |name: &str| -> Option<Vec<f32>> { f(name).ok() };
let rope_of = |base: f32, yarn: bool| -> Vec<f32> {
if yarn {
crate::attention::yarn_inv_freq(cfg.rope_head_dim, base, 16.0, 65536, 32.0, 1.0)
} else {
crate::attention::rope_inv_freq(cfg.rope_head_dim, base)
}
};
let globals = Dsv4Globals {
inv_freq_compress: rope_of(160_000.0, true),
inv_freq_window: rope_of(10_000.0, false),
embed: q("model.embed_tokens.weight")?,
norm: f("model.norm.weight")?,
head: q("lm_head.weight")?,
hc_head_fn: f("model.hc_head_fn")?,
hc_head_base: f("model.hc_head_base")?,
hc_head_scale: *f("model.hc_head_scale")?
.first()
.ok_or("dsv4: empty hc_head_scale")?,
};
let mut layers = Vec::with_capacity(n_layers);
for li in 0..n_layers {
let p = format!("model.layers.{li}");
let scale3 = |name: &str| -> Result<[f32; 3], String> {
let v = f(name)?;
if v.len() < 3 {
return Err(format!("{name}: expected 3 scales, got {}", v.len()));
}
Ok([v[0], v[1], v[2]])
};
let compressor = match q(&format!("{p}.self_attn.compressor.wkv.weight")) {
Ok(wkv) => {
let ape = f(&format!("{p}.self_attn.compressor.ape"))?;
let width = wkv.rows();
let ratio = (ape.len() / width.max(1)).max(1);
Some(Dsv4Compressor {
wkv,
wgate: q(&format!("{p}.self_attn.compressor.wgate.weight"))?,
norm: f(&format!("{p}.self_attn.compressor.norm.weight"))?,
ape,
ratio,
overlap: ratio == 4,
})
}
Err(_) => None,
};
let indexer = match q(&format!("{p}.self_attn.indexer.wq_b.weight")) {
Ok(wq_b) => {
let ape = f(&format!("{p}.self_attn.indexer.compressor.ape"))?;
let cwkv = q(&format!("{p}.self_attn.indexer.compressor.wkv.weight"))?;
let width = cwkv.rows();
let ratio = (ape.len() / width.max(1)).max(1);
Some(Dsv4Indexer {
wq_b,
weights_proj: q(&format!("{p}.self_attn.indexer.weights_proj.weight"))?,
compressor: Dsv4Compressor {
wkv: cwkv,
wgate: q(&format!("{p}.self_attn.indexer.compressor.wgate.weight"))?,
norm: f(&format!("{p}.self_attn.indexer.compressor.norm.weight"))?,
ape,
ratio,
overlap: ratio == 4,
},
})
}
Err(_) => None,
};
let mut experts = Vec::with_capacity(cfg.n_routed_experts);
for e in 0..cfg.n_routed_experts {
let ep = format!("{p}.mlp.experts.{e}");
experts.push(Dsv4Expert {
w1: q(&format!("{ep}.gate_proj.weight"))?,
w2: q(&format!("{ep}.down_proj.weight"))?,
w3: q(&format!("{ep}.up_proj.weight"))?,
});
}
layers.push(Dsv4Layer {
attn_norm: f(&format!("{p}.input_layernorm.weight"))?,
ffn_norm: f(&format!("{p}.post_attention_layernorm.weight"))?,
wq_a: q(&format!("{p}.self_attn.wq_a.weight"))?,
q_norm: f(&format!("{p}.self_attn.q_norm.weight"))?,
wq_b: q(&format!("{p}.self_attn.wq_b.weight"))?,
wkv: q(&format!("{p}.self_attn.wkv.weight"))?,
kv_norm: f(&format!("{p}.self_attn.kv_norm.weight"))?,
wo_a: q(&format!("{p}.self_attn.wo_a.weight"))?,
wo_b: q(&format!("{p}.self_attn.wo_b.weight"))?,
attn_sink: f(&format!("{p}.self_attn.attn_sink"))?,
compressor,
indexer,
hc_attn_fn: f(&format!("{p}.hc_attn_fn"))?,
hc_attn_base: f(&format!("{p}.hc_attn_base"))?,
hc_attn_scale: scale3(&format!("{p}.hc_attn_scale"))?,
hc_ffn_fn: f(&format!("{p}.hc_ffn_fn"))?,
hc_ffn_base: f(&format!("{p}.hc_ffn_base"))?,
hc_ffn_scale: scale3(&format!("{p}.hc_ffn_scale"))?,
gate: q(&format!("{p}.mlp.gate.weight"))?,
gate_bias: opt_f(&format!("{p}.mlp.expert_bias")),
tid2eid: opt_f(&format!("{p}.mlp.tid2eid")),
experts,
mask: if model.tensor(&format!("{p}.mlp.tid2eid")).is_some() {
None
} else {
crate::loader::moe_task_mask(&format!("{p}."), cfg.n_routed_experts)
},
shared: Dsv4Expert {
w1: q(&format!("{p}.mlp.shared_expert.gate_proj.weight"))?,
w2: q(&format!("{p}.mlp.shared_expert.down_proj.weight"))?,
w3: q(&format!("{p}.mlp.shared_expert.up_proj.weight"))?,
},
});
}
Ok((globals, layers))
}
#[cfg(test)]
mod tests {
use super::*;
fn toy() -> (Dsv4Globals, Vec<Dsv4Layer>, Dsv4Cfg) {
use crate::qtensor::QTensor;
let cfg = Dsv4Cfg {
dim: 32,
n_heads: 4,
head_dim: 8,
rope_head_dim: 4,
q_lora_rank: 16,
o_lora_rank: 16,
o_groups: 2,
hc_mult: 4,
hc_sinkhorn_iters: 20,
hc_eps: 1e-6,
norm_eps: 1e-6,
n_routed_experts: 8,
top_k: 2,
moe_inter: 16,
route_scale: 1.0,
swiglu_limit: 10.0,
window: 6,
index_topk: 8,
vocab: 24,
};
let w = |n: usize, seed: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 7 + seed * 13) % 101) as f32 / 101.0 - 0.5) * 0.3)
.collect()
};
let t = |rows: usize, cols: usize, seed: usize| {
QTensor::from_f32(w(rows * cols, seed), rows, cols)
};
let ones = |n: usize| vec![1.0f32; n];
let (dim, hc) = (cfg.dim, cfg.hc_mult);
let q_width = cfg.n_heads * cfg.head_dim;
let kv_width = cfg.head_dim;
let o_per_group = q_width / cfg.o_groups;
let mut layers = Vec::new();
for li in 0..2 {
let experts: Vec<Dsv4Expert> = (0..cfg.n_routed_experts)
.map(|e| Dsv4Expert {
w1: t(cfg.moe_inter, dim, 40 + e + li * 8),
w2: t(dim, cfg.moe_inter, 60 + e + li * 8),
w3: t(cfg.moe_inter, dim, 80 + e + li * 8),
})
.collect();
layers.push(Dsv4Layer {
attn_norm: ones(dim),
ffn_norm: ones(dim),
wq_a: t(cfg.q_lora_rank, dim, 1 + li),
q_norm: ones(cfg.q_lora_rank),
wq_b: t(q_width, cfg.q_lora_rank, 3 + li),
wkv: t(kv_width, dim, 5 + li),
kv_norm: ones(kv_width),
wo_a: t(cfg.o_groups * cfg.o_lora_rank, o_per_group, 7 + li),
wo_b: t(dim, cfg.o_groups * cfg.o_lora_rank, 9 + li),
attn_sink: vec![0.1; cfg.n_heads],
compressor: if li == 1 {
Some(Dsv4Compressor {
wkv: t(2 * kv_width, dim, 11),
wgate: t(2 * kv_width, dim, 13),
norm: ones(kv_width),
ape: vec![0.01; 4 * 2 * kv_width],
ratio: 4,
overlap: true,
})
} else {
None
},
indexer: if li == 1 {
Some(Dsv4Indexer {
wq_b: t(2 * 16, cfg.q_lora_rank, 41),
weights_proj: t(2, dim, 43),
compressor: Dsv4Compressor {
wkv: t(2 * 16, dim, 45),
wgate: t(2 * 16, dim, 47),
norm: ones(16),
ape: vec![0.01; 4 * 2 * 16],
ratio: 4,
overlap: true,
},
})
} else {
None
},
hc_attn_fn: w((2 + hc) * hc * hc * dim, 15 + li),
hc_attn_base: w((2 + hc) * hc, 17 + li),
hc_attn_scale: [1.0, 1.0, 1.0],
hc_ffn_fn: w((2 + hc) * hc * hc * dim, 19 + li),
hc_ffn_base: w((2 + hc) * hc, 21 + li),
hc_ffn_scale: [1.0, 1.0, 1.0],
gate: t(cfg.n_routed_experts, dim, 23 + li),
gate_bias: if li == 1 {
Some(vec![0.0; cfg.n_routed_experts])
} else {
None
},
tid2eid: if li == 0 {
Some(
(0..cfg.vocab * cfg.top_k)
.map(|i| (i % cfg.n_routed_experts) as f32)
.collect(),
)
} else {
None
},
experts,
mask: None,
shared: Dsv4Expert {
w1: t(cfg.moe_inter, dim, 25 + li),
w2: t(dim, cfg.moe_inter, 27 + li),
w3: t(cfg.moe_inter, dim, 29 + li),
},
});
}
let inv = |base: f32| -> Vec<f32> {
(0..cfg.rope_head_dim / 2)
.map(|i| 1.0 / base.powf(2.0 * i as f32 / cfg.rope_head_dim as f32))
.collect()
};
let g = Dsv4Globals {
inv_freq_compress: inv(160000.0),
inv_freq_window: inv(10000.0),
embed: t(cfg.vocab, dim, 31),
norm: ones(dim),
head: t(cfg.vocab, dim, 33),
hc_head_fn: w(hc * hc * dim, 35),
hc_head_base: w(hc, 37),
hc_head_scale: 1.0,
};
(g, layers, cfg)
}
#[test]
fn forward_token_decodes_a_sequence_without_falling_over() {
let (g, layers, cfg) = toy();
let mut st = Dsv4State::new(layers.len());
let inv_freq: Vec<f32> = (0..cfg.rope_head_dim / 2)
.map(|i| 1.0 / 10000f32.powf(2.0 * i as f32 / cfg.rope_head_dim as f32))
.collect();
let mut logits = Vec::new();
let mut first: Option<Vec<f32>> = None;
for (step, tok) in [3u32, 7, 1, 9, 4, 2, 8, 5, 6, 0].into_iter().enumerate() {
forward_token(
&g,
&layers,
&cfg,
&mut st,
tok,
&inv_freq,
None,
&mut logits,
);
assert_eq!(logits.len(), cfg.vocab, "step {step}: logit count");
assert!(
logits.iter().all(|v| v.is_finite()),
"step {step}: non-finite logit — {logits:?}"
);
let spread = logits.iter().cloned().fold(f32::MIN, f32::max)
- logits.iter().cloned().fold(f32::MAX, f32::min);
assert!(spread > 1e-6, "step {step}: logits are flat ({spread})");
if step == 0 {
first = Some(logits.clone());
}
assert_eq!(st.pos, step + 1, "position bookkeeping");
}
assert!(!st.window[0].is_empty(), "sliding window never filled");
for (li, w) in st.window.iter().enumerate() {
assert!(
w.len() / cfg.head_dim <= cfg.window,
"layer {li}: window holds {} positions, cap is {}",
w.len() / cfg.head_dim,
cfg.window
);
}
assert!(
!st.compressed[1].is_empty(),
"compressor layer produced no compressed KV in 10 tokens"
);
assert_eq!(
st.compressed[1].len() / cfg.head_dim,
2,
"expected two folds in ten tokens at ratio 4"
);
assert!(
!st.prev_kv[1].is_empty(),
"the overlapping compressor never kept a previous window"
);
for (li, l) in layers.iter().enumerate() {
if l.indexer.is_some() {
assert!(
!st.index_kv[li].is_empty(),
"layer {li} has an indexer but its cache stayed empty"
);
}
}
let mut fresh = Dsv4State::new(layers.len());
let mut relogits = Vec::new();
forward_token(
&g,
&layers,
&cfg,
&mut fresh,
3,
&inv_freq,
None,
&mut relogits,
);
assert_eq!(
relogits,
first.unwrap(),
"the same token from a fresh state must reproduce exactly"
);
}
#[test]
fn swiglu_limit_clamps_up_both_ways_and_gate_only_from_above() {
let inter = 4;
let gate_src = [-50.0f32, 50.0, 1.0, -1.0];
let up_src = [50.0f32, -50.0, 1.0, -1.0];
let limit = 10.0f32;
let mut got = vec![0.0f32; inter];
expert_swiglu(
&[0.0],
&|_, d| d.copy_from_slice(&gate_src),
&|_, d| d.copy_from_slice(&up_src),
&|src, d| d.copy_from_slice(src),
inter,
1.0,
limit,
&mut got,
);
let silu = |g: f32| g / (1.0 + (-g).exp());
let want = [
silu(-50.0) * limit,
silu(limit) * -limit,
silu(1.0) * 1.0,
silu(-1.0) * -1.0,
];
for (i, w) in want.iter().enumerate() {
assert!(
(got[i] - w).abs() < 1e-5,
"lane {i}: got {} want {w}",
got[i]
);
}
let mut raw = vec![0.0f32; inter];
expert_swiglu(
&[0.0],
&|_, d| d.copy_from_slice(&gate_src),
&|_, d| d.copy_from_slice(&up_src),
&|src, d| d.copy_from_slice(src),
inter,
1.0,
0.0,
&mut raw,
);
assert!(
(raw[1] - silu(50.0) * -50.0).abs() < 1e-3,
"limit 0 must not clamp"
);
}
#[test]
fn grouped_projection_is_identical_with_and_without_a_pool() {
let (groups, lora, per_group, dim) = (4usize, 128usize, 64usize, 32usize);
let attn: Vec<f32> = (0..groups * per_group)
.map(|i| ((i * 13) as f32 * 0.021).sin())
.collect();
let wo_a: Vec<f32> = (0..groups * lora * per_group)
.map(|i| ((i * 7) as f32 * 0.011).cos())
.collect();
let wo_b: Vec<f32> = (0..dim * groups * lora)
.map(|i| ((i * 5) as f32 * 0.009).sin())
.collect();
let row = |r: usize, x: &[f32], _sc: &mut [f32]| -> f32 {
wo_a[r * per_group..(r + 1) * per_group]
.iter()
.zip(x)
.map(|(a, b)| a * b)
.sum()
};
let project = |mid: &[f32], dst: &mut [f32]| {
for (d, o) in dst.iter_mut().enumerate() {
*o = wo_b[d * mid.len()..(d + 1) * mid.len()]
.iter()
.zip(mid)
.map(|(a, b)| a * b)
.sum();
}
};
let mut serial = vec![0.0f32; dim];
o_project(
&attn,
&row,
per_group,
&project,
groups,
lora,
None,
&mut serial,
);
let pool = crate::pool::Pool::new(4);
let mut pooled = vec![0.0f32; dim];
o_project(
&attn,
&row,
per_group,
&project,
groups,
lora,
Some(&pool),
&mut pooled,
);
assert_eq!(serial, pooled, "the pooled projection diverged");
assert!(
serial.iter().any(|v| v.abs() > 1e-6),
"test data is degenerate"
);
}
#[test]
fn overlapping_compressor_folds_both_windows() {
let (ratio, d) = (2usize, 3usize);
let cur_kv: Vec<f32> = vec![
1.0, 1.0, 1.0, 10.0, 20.0, 30.0, 2.0, 2.0, 2.0, 40.0, 50.0, 60.0, ];
let cur_sc: Vec<f32> = vec![
0.0, 0.0, 0.0, 0.0, 0.0, 100.0, 0.0, 0.0, 0.0, 100.0, 100.0, 0.0,
];
let prev_kv: Vec<f32> = vec![
7.0, 8.0, 9.0, 0.0, 0.0, 0.0, 5.0, 6.0, 7.0, 0.0, 0.0, 0.0,
];
let prev_sc = vec![0.0f32; ratio * 2 * d];
let mut out = vec![0.0f32; d];
compress_window_overlap(&prev_kv, &prev_sc, &cur_kv, &cur_sc, ratio, d, &mut out);
assert!((out[0] - 40.0).abs() < 1e-3, "dim0 = {}", out[0]);
assert!((out[1] - 50.0).abs() < 1e-3, "dim1 = {}", out[1]);
assert!((out[2] - 30.0).abs() < 1e-3, "dim2 = {}", out[2]);
let mut first = vec![0.0f32; d];
compress_window_overlap(&[], &[], &cur_kv, &cur_sc, ratio, d, &mut first);
assert!(
first.iter().all(|v| v.is_finite()),
"first window: {first:?}"
);
assert!((first[0] - 40.0).abs() < 1e-3, "first dim0 = {}", first[0]);
let mut both = vec![0.0f32; d];
let strong_prev = vec![100.0f32; ratio * 2 * d];
compress_window_overlap(
&prev_kv,
&strong_prev,
&cur_kv,
&cur_sc,
ratio,
d,
&mut both,
);
assert!(
(both[0] - 40.0).abs() > 1.0,
"a scored previous window must move the fold, got {}",
both[0]
);
}
#[test]
fn sinkhorn_matches_the_reference_numbers() {
let hc = 4;
let mixes: Vec<f32> = (0..24).map(|i| (i as f32 * 0.37).sin() * 3.0).collect();
let base: Vec<f32> = (0..24).map(|i| (i as f32 * 0.11).cos()).collect();
let (mut pre, mut post, mut comb) = (vec![0.0; hc], vec![0.0; hc], vec![0.0; hc * hc]);
hc_split_sinkhorn(
&mixes,
&[1.0, 1.0, 1.0],
&base,
hc,
20,
1e-6,
&mut pre,
&mut post,
&mut comb,
);
let want_pre = [0.7310596, 0.8888268, 0.9525191, 0.97424865];
let want_post = [1.9600224, 1.9534285, 1.9201256, 1.8160983];
let want_comb = [
0.5996052,
0.28253591,
0.09218107,
0.025676856,
0.17564717,
0.22228767,
0.27174541,
0.33031881,
0.029528176,
0.12206022,
0.32619134,
0.5222193,
0.19521846,
0.37311527,
0.30988118,
0.12178412,
];
for (i, w) in want_pre.iter().enumerate() {
assert!((pre[i] - w).abs() < 1e-5, "pre[{i}]: {} vs {w}", pre[i]);
}
for (i, w) in want_post.iter().enumerate() {
assert!((post[i] - w).abs() < 1e-5, "post[{i}]: {} vs {w}", post[i]);
}
for (i, w) in want_comb.iter().enumerate() {
assert!((comb[i] - w).abs() < 1e-4, "comb[{i}]: {} vs {w}", comb[i]);
}
}
#[test]
fn sinkhorn_leaves_the_mixing_matrix_doubly_stochastic() {
let hc = 4;
let mix_hc = (2 + hc) * hc;
let mixes: Vec<f32> = (0..mix_hc).map(|i| (i as f32 * 0.37).sin() * 3.0).collect();
let base: Vec<f32> = (0..mix_hc).map(|i| (i as f32 * 0.11).cos()).collect();
let (mut pre, mut post, mut comb) = (vec![0.0; hc], vec![0.0; hc], vec![0.0; hc * hc]);
hc_split_sinkhorn(
&mixes,
&[1.0, 1.0, 1.0],
&base,
hc,
20,
1e-6,
&mut pre,
&mut post,
&mut comb,
);
for j in 0..hc {
let r: f32 = comb[j * hc..(j + 1) * hc].iter().sum();
assert!((r - 1.0).abs() < 2e-3, "row {j} sums to {r}");
let c: f32 = (0..hc).map(|k| comb[k * hc + j]).sum();
assert!((c - 1.0).abs() < 2e-3, "col {j} sums to {c}");
}
assert!(pre.iter().all(|&v| v > 0.0 && v < 1.001));
assert!(post.iter().all(|&v| v >= 0.0 && v <= 2.0));
}
#[test]
fn expand_of_identical_copies_is_a_fixed_point() {
let (hc, dim) = (4usize, 3usize);
let residual: Vec<f32> = std::iter::repeat([1.5f32, -2.0, 0.25])
.take(hc)
.flatten()
.collect();
let comb = {
vec![0.25f32; hc * hc]
};
let post = vec![0.0f32; hc];
let mut out = vec![0.0f32; hc * dim];
hc_expand(&[0.0; 3], &residual, &post, &comb, hc, dim, &mut out);
for (o, r) in out.iter().zip(&residual) {
assert!((o - r).abs() < 1e-6, "{o} vs {r}");
}
}
#[test]
fn selection_bias_steers_the_choice_but_not_the_weights() {
let scores = [3.0f32, 0.1, 2.0, 0.05];
let bias = [0.0f32, 10.0, 0.0, 0.0];
let (mut idx, mut w) = (Vec::new(), Vec::new());
route(&scores, Some(&bias), 2, 1.5, None, None, &mut idx, &mut w);
assert_eq!(idx[0], 1, "the biased expert must win selection");
assert_eq!(idx[1], 0);
assert!(w[0] < w[1], "biased expert kept its own (small) weight");
let sum: f32 = w.iter().sum();
assert!((sum - 1.5).abs() < 1e-5, "weights renormalize then scale");
}
#[test]
fn attention_sink_drains_weight_without_contributing_output() {
let hd = 2;
let q = [1.0f32, 0.0];
let kv = [1.0f32, 0.0, 0.0, 1.0];
let mut out = vec![0.0f32; hd];
sparse_attend(&q, &kv, &[0, 1], f32::NEG_INFINITY, 1.0, hd, &mut out);
let plain = out.clone();
assert!(plain[0] > plain[1], "the aligned key must dominate");
sparse_attend(&q, &kv, &[0, 1], 20.0, 1.0, hd, &mut out);
assert!(
out[0] < plain[0] * 0.01 && out[1] < plain[1] * 0.01,
"a large sink must drain nearly all the mass: {out:?}"
);
}
#[test]
fn masked_positions_leave_the_denominator_alone() {
let hd = 2;
let q = [1.0f32, 0.0];
let kv = [1.0f32, 0.0, 0.0, 1.0];
let (mut a, mut b) = (vec![0.0f32; hd], vec![0.0f32; hd]);
sparse_attend(&q, &kv, &[0], f32::NEG_INFINITY, 1.0, hd, &mut a);
sparse_attend(
&q,
&kv,
&[0, usize::MAX],
f32::NEG_INFINITY,
1.0,
hd,
&mut b,
);
for (x, y) in a.iter().zip(&b) {
assert!((x - y).abs() < 1e-6, "{x} vs {y}");
}
}
#[test]
fn rope_tail_inverts_itself() {
let inv_freq = [1.0f32, 0.5];
let orig = [9.0f32, 8.0, 1.0, 2.0, 3.0, 4.0];
let mut v = orig;
rope_tail(&mut v, &inv_freq, 7, 4, false);
assert!(v[..2] == orig[..2], "the non-rope head must not move");
assert!(v[2..] != orig[2..], "the tail must actually rotate");
rope_tail(&mut v, &inv_freq, 7, 4, true);
for (a, b) in v.iter().zip(&orig) {
assert!((a - b).abs() < 1e-5, "{a} vs {b}");
}
}
#[test]
fn compressor_pools_the_window_per_dimension() {
let (ratio, width) = (2usize, 2usize);
let kv = [1.0f32, 10.0, 3.0, 20.0];
let score = [0.0f32, 0.0, 0.0, 50.0];
let ape = vec![0.0f32; ratio * width];
let mut out = vec![0.0f32; width];
compress_window(&kv, &score, &ape, ratio, width, &mut out);
assert!(
(out[0] - 2.0).abs() < 1e-5,
"equal scores average: {}",
out[0]
);
assert!(
(out[1] - 20.0).abs() < 1e-3,
"a dominant score wins: {}",
out[1]
);
}
#[test]
fn index_scores_relu_before_weighting() {
let (nh, hd) = (2usize, 2usize);
let q = [1.0f32, 0.0, -1.0, 0.0];
let kv = [1.0f32, 0.0, 0.0, 1.0];
let w = [1.0f32, 1.0];
let mut sc = Vec::new();
index_scores(&q, &kv, &w, nh, hd, 2, 2, None, &mut sc);
assert!(sc[0] > 0.9, "abstention, not veto: {:?}", sc);
}
#[test]
fn index_scores_mask_the_future() {
let (nh, hd) = (1usize, 2usize);
let q = [1.0f32, 0.0];
let kv = [1.0f32, 0.0, 1.0, 0.0, 1.0, 0.0];
let w = [1.0f32];
let mut sc = Vec::new();
index_scores(&q, &kv, &w, nh, hd, 3, 2, None, &mut sc);
assert!(sc[0].is_finite() && sc[1].is_finite());
assert!(sc[2] == f32::NEG_INFINITY, "position 2 is in the future");
let mut idx = Vec::new();
top_k_positions(&sc, 3, &mut idx);
assert_eq!(idx, vec![0, 1], "a masked slot never wins a slot");
}
#[test]
fn top_k_is_deterministic_on_ties() {
let sc = [1.0f32, 1.0, 1.0, 0.0];
let mut idx = Vec::new();
top_k_positions(&sc, 2, &mut idx);
assert_eq!(idx, vec![0, 1], "ties resolve to the lower index");
}
#[test]
fn hc_block_preserves_the_copy_structure_and_applies_the_block() {
let cfg = Dsv4Cfg {
dim: 4,
n_heads: 1,
head_dim: 4,
rope_head_dim: 2,
q_lora_rank: 4,
o_lora_rank: 2,
o_groups: 1,
hc_mult: 4,
hc_sinkhorn_iters: 20,
hc_eps: 1e-6,
norm_eps: 1e-6,
n_routed_experts: 2,
top_k: 1,
moe_inter: 4,
route_scale: 1.0,
swiglu_limit: 10.0,
window: 128,
index_topk: 4,
vocab: 8,
};
let (hc, dim) = (cfg.hc_mult, cfg.dim);
let mix_hc = (2 + hc) * hc;
let hc_fn: Vec<f32> = (0..mix_hc * hc * dim)
.map(|i| ((i % 13) as f32 - 6.0) * 0.05)
.collect();
let hc_base: Vec<f32> = (0..mix_hc).map(|i| (i as f32 * 0.2).sin()).collect();
let norm_w = vec![1.0f32; dim];
let mut state: Vec<f32> = (0..hc * dim).map(|i| (i as f32 * 0.3).cos()).collect();
let before = state.clone();
let mut scratch = HcScratch::new(&cfg);
hc_block(
&mut state,
&hc_fn,
&[1.0, 1.0, 1.0],
&hc_base,
&norm_w,
&cfg,
&mut scratch,
None,
|_folded, out: &mut [f32]| out.iter_mut().for_each(|o| *o = 1.0),
);
assert_eq!(state.len(), before.len(), "copy structure must survive");
assert!(state.iter().all(|v| v.is_finite()), "{state:?}");
assert!(
state.iter().zip(&before).any(|(a, b)| (a - b).abs() > 1e-4),
"the block's output has to reach the state"
);
}
#[test]
fn hash_route_reads_the_table_row() {
let table = [7.0f32, 9.0, 1.0, 2.0, 5.0, 6.0];
assert_eq!(hash_route(&table, 3, 2, 0), vec![7, 9]);
assert_eq!(hash_route(&table, 3, 2, 2), vec![5, 6]);
assert_eq!(hash_route(&table, 3, 2, 99), vec![5, 6]);
}
#[test]
fn a_task_mask_restricts_selection_and_renormalizes() {
let scores = [0.1f32, 4.0, 1.0, 9.0];
let (mut idx, mut w) = (Vec::new(), Vec::new());
route(&scores, None, 2, 1.0, None, None, &mut idx, &mut w);
assert_eq!(idx, vec![3, 1], "unmasked: the two best win");
let sum: f32 = w.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "weights must sum to route_scale");
let mask = [true, false, true, true];
let (mut i2, mut w2) = (Vec::new(), Vec::new());
route(&scores, None, 2, 1.0, None, Some(&mask), &mut i2, &mut w2);
assert_eq!(i2, vec![3, 2], "masked expert must not be selected");
let sum2: f32 = w2.iter().sum();
assert!((sum2 - 1.0).abs() < 1e-5, "masked weights must renormalize");
let tight = [false, false, false, true];
let (mut i3, mut w3) = (Vec::new(), Vec::new());
route(&scores, None, 2, 1.0, None, Some(&tight), &mut i3, &mut w3);
assert_eq!(i3, vec![3]);
assert_eq!(w3.len(), 1);
}
#[test]
fn hash_layers_weight_the_experts_the_table_names() {
let scores = [0.1f32, 0.4, 0.2, 5.0];
let table = vec![0.0f32, 1.0];
let idx_forced = hash_route(&table, 1, 2, 0);
assert_eq!(idx_forced, vec![0, 1]);
let (mut idx, mut w) = (Vec::new(), Vec::new());
route(
&scores,
None,
2,
1.0,
Some(&idx_forced),
None,
&mut idx,
&mut w,
);
assert_eq!(idx, vec![0, 1], "the table must decide the experts");
let sp = |x: f32| (1.0 + x.exp()).ln().sqrt();
let (s0, s1) = (sp(scores[0]), sp(scores[1]));
let tot = s0 + s1;
assert!(
(w[0] - s0 / tot).abs() < 1e-6,
"w[0]={} want {}",
w[0],
s0 / tot
);
assert!(
(w[1] - s1 / tot).abs() < 1e-6,
"w[1]={} want {}",
w[1],
s1 / tot
);
let (mut idx2, mut w2) = (Vec::new(), Vec::new());
route(&scores, None, 2, 1.0, None, None, &mut idx2, &mut w2);
assert_eq!(idx2[0], 3, "without a table the highest score still wins");
}
}