spark_model/layers/glm5next_dsa_ref/
mod.rs1#[derive(Clone, Copy, Debug)]
36pub struct DsaDims {
37 pub hidden: usize,
38 pub index_heads: usize,
40 pub index_head_dim: usize,
42 pub index_kpool: usize,
43 pub index_topk: usize,
44 pub always_select_tail: bool,
45 pub q_lora_rank: usize,
46 pub heads: usize,
48 pub kv_lora_rank: usize,
49 pub qk_nope_head_dim: usize,
50 pub qk_rope_head_dim: usize,
52 pub v_head_dim: usize,
53}
54
55impl DsaDims {
56 pub fn qk_head_dim(&self) -> usize {
57 self.qk_nope_head_dim + self.qk_rope_head_dim
58 }
59 pub fn select_k(&self, n_pools: usize) -> usize {
61 (self.index_topk / self.index_kpool).min(n_pools)
62 }
63 pub fn out_width(&self) -> usize {
65 self.index_topk
66 + if self.always_select_tail {
67 self.index_kpool - 1
68 } else {
69 0
70 }
71 }
72 pub fn is_nope(&self) -> bool {
73 self.qk_rope_head_dim == 0
74 }
75}
76
77pub const INVALID: i32 = -1;
79
80#[inline]
81fn sigmoid(x: f32) -> f32 {
82 1.0 / (1.0 + (-x).exp())
83}
84
85pub fn layer_norm(x: &[f32], w: &[f32], b: &[f32], d: usize, eps: f32) -> Vec<f32> {
90 let mut out = vec![0.0f32; x.len()];
91 for (row_in, row_out) in x.chunks_exact(d).zip(out.chunks_exact_mut(d)) {
92 let mean = row_in.iter().sum::<f32>() / d as f32;
93 let var = row_in.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / d as f32;
94 let inv = 1.0 / (var + eps).sqrt();
95 for i in 0..d {
96 row_out[i] = (row_in[i] - mean) * inv * w[i] + b[i];
97 }
98 }
99 out
100}
101
102pub fn rms_norm(x: &[f32], w: &[f32], d: usize, eps: f32) -> Vec<f32> {
104 let mut out = vec![0.0f32; x.len()];
105 for (row_in, row_out) in x.chunks_exact(d).zip(out.chunks_exact_mut(d)) {
106 let inv = 1.0 / (row_in.iter().map(|v| v * v).sum::<f32>() / d as f32 + eps).sqrt();
107 for i in 0..d {
108 row_out[i] = row_in[i] * inv * w[i];
109 }
110 }
111 out
112}
113
114pub fn linear(x: &[f32], m: usize, k: usize, w: &[f32], n: usize) -> Vec<f32> {
116 let mut out = vec![0.0f32; m * n];
117 for row in 0..m {
118 for col in 0..n {
119 let mut acc = 0.0f32;
120 for i in 0..k {
121 acc += x[row * k + i] * w[col * k + i];
122 }
123 out[row * n + col] = acc;
124 }
125 }
126 out
127}
128
129pub struct Pools {
131 pub keys: Vec<f32>,
133 pub indices: Vec<i32>,
135 pub valid: Vec<u8>,
137 pub n_pools: usize,
138}
139
140pub fn pool_states(
144 k: &[f32],
145 gate: &[f32],
146 valid: &[u8],
147 ape: &[f32],
148 dims: DsaDims,
149 seq: usize,
150) -> Pools {
151 let (d, kp) = (dims.index_head_dim, dims.index_kpool);
152 let n_pools = seq.div_ceil(kp);
153 let first_key = valid.iter().position(|v| *v != 0).unwrap_or(seq) as i64;
155
156 let mut keys = vec![0.0f32; n_pools * d];
157 let mut indices = vec![INVALID; n_pools * kp];
158 let mut pvalid = vec![0u8; n_pools];
159 let mut logits = vec![0.0f32; kp];
160
161 for p in 0..n_pools {
162 let mut slot_valid = [false; 64];
163 let mut slot_idx = [0usize; 64];
164 let mut all = true;
165 for s in 0..kp {
166 let raw = first_key + (p * kp + s) as i64;
167 let in_range = raw >= 0 && (raw as usize) < seq;
168 let ok = in_range && valid[raw as usize] != 0;
169 slot_valid[s] = ok;
170 slot_idx[s] = if in_range { raw as usize } else { 0 };
171 all &= ok;
172 indices[p * kp + s] = if ok { raw as i32 } else { INVALID };
173 }
174 pvalid[p] = all as u8;
175
176 for dd in 0..d {
178 let mut mx = f32::NEG_INFINITY;
179 for s in 0..kp {
180 logits[s] = if slot_valid[s] {
181 gate[slot_idx[s] * d + dd] + ape[s * d + dd]
182 } else {
183 f32::NEG_INFINITY
184 };
185 mx = mx.max(logits[s]);
186 }
187 let mut sum = 0.0f32;
188 for s in 0..kp {
189 let e = if logits[s] == f32::NEG_INFINITY {
190 0.0
191 } else {
192 (logits[s] - mx).exp()
193 };
194 logits[s] = e;
195 sum += e;
196 }
197 let inv = if sum > 0.0 { 1.0 / sum } else { 0.0 };
199 let mut acc = 0.0f32;
200 for s in 0..kp {
201 if slot_valid[s] {
202 acc += logits[s] * inv * k[slot_idx[s] * d + dd];
203 }
204 }
205 keys[p * d + dd] = acc;
206 }
207 }
208 let keep: Vec<usize> = (0..n_pools).filter(|p| pvalid[*p] != 0).collect();
213 if keep.len() == n_pools {
214 return Pools {
215 keys,
216 indices,
217 valid: pvalid,
218 n_pools,
219 };
220 }
221 let mut ck = vec![0.0f32; keep.len() * d];
222 let mut ci = vec![INVALID; keep.len() * kp];
223 let mut cv = vec![0u8; keep.len()];
224 for (j, p) in keep.iter().enumerate() {
225 ck[j * d..(j + 1) * d].copy_from_slice(&keys[p * d..(p + 1) * d]);
226 ci[j * kp..(j + 1) * kp].copy_from_slice(&indices[p * kp..(p + 1) * kp]);
227 cv[j] = pvalid[*p];
228 }
229 Pools {
230 keys: ck,
231 indices: ci,
232 valid: cv,
233 n_pools: keep.len(),
234 }
235}
236
237pub fn kept_pools(valid: &[u8], dims: DsaDims, seq: usize) -> Vec<i32> {
242 let kp = dims.index_kpool;
243 let n_pools = seq.div_ceil(kp);
244 let first_key = valid.iter().position(|v| *v != 0).unwrap_or(seq) as i64;
245 (0..n_pools)
246 .filter(|p| {
247 (0..kp).all(|s| {
248 let raw = first_key + (p * kp + s) as i64;
249 raw >= 0 && (raw as usize) < seq && valid[raw as usize] != 0
250 })
251 })
252 .map(|p| p as i32)
253 .collect()
254}
255
256pub fn index_scores(
261 q: &[f32],
262 weights: &[f32],
263 pools: &Pools,
264 dims: DsaDims,
265 q_rows: usize,
266) -> Vec<f32> {
267 let (h, d, p) = (dims.index_heads, dims.index_head_dim, pools.n_pools);
268 let scale = (d as f32).powf(-0.5);
269 let mut out = vec![0.0f32; q_rows * p];
270 for r in 0..q_rows {
271 for pp in 0..p {
272 let mut acc = 0.0f32;
273 for hh in 0..h {
274 let mut dot = 0.0f32;
275 for dd in 0..d {
276 dot += q[(r * h + hh) * d + dd] * pools.keys[pp * d + dd];
277 }
278 acc += weights[r * h + hh] * (scale * dot).max(0.0);
281 }
282 out[r * p + pp] = acc;
283 }
284 }
285 out
286}
287
288pub fn visible(valid_keys: &[u8], q_pos: usize, key_idx: usize) -> bool {
290 key_idx <= q_pos && valid_keys[key_idx] != 0
291}
292
293pub fn topk_pools(
301 scores: &[f32],
302 valid_candidates: &[u8],
303 n_pools: usize,
304 q_rows: usize,
305 select_k: usize,
306) -> Vec<i32> {
307 let mut out = vec![INVALID; q_rows * select_k];
308 let mut buf: Vec<(f32, usize)> = Vec::with_capacity(n_pools);
309 for r in 0..q_rows {
310 buf.clear();
311 for p in 0..n_pools {
312 let s = if valid_candidates[r * n_pools + p] != 0 {
313 scores[r * n_pools + p]
314 } else {
315 f32::MIN
316 };
317 buf.push((s, p));
318 }
319 buf.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap().then(a.1.cmp(&b.1)));
320 for (j, (_, p)) in buf.iter().take(select_k).enumerate() {
321 out[r * select_k + j] = *p as i32;
322 }
323 }
324 out
325}
326
327#[allow(clippy::too_many_arguments)]
334pub fn expand_selection(
335 selected: &[i32],
336 pools: &Pools,
337 valid_candidates: &[u8],
338 valid_keys: &[u8],
339 q_positions: &[usize],
340 q_mask: &[u8],
341 dims: DsaDims,
342 seq: usize,
343 select_k: usize,
344) -> Vec<i32> {
345 let (kp, width) = (dims.index_kpool, dims.out_width());
346 let q_rows = q_positions.len();
347 let mut out = vec![INVALID; q_rows * width];
348
349 let first_key = valid_keys.iter().position(|v| *v != 0).unwrap_or(seq) as i64;
350 for r in 0..q_rows {
351 let row = &mut out[r * width..(r + 1) * width];
352 if q_mask[r] == 0 {
353 continue; }
355 let mut w = 0usize;
356 for j in 0..select_k {
357 let p = selected[r * select_k + j];
358 let ok = p >= 0 && valid_candidates[r * pools.n_pools + p as usize] != 0;
359 for s in 0..kp {
360 row[w] = if ok {
361 pools.indices[p as usize * kp + s]
362 } else {
363 INVALID
364 };
365 w += 1;
366 }
367 }
368 if dims.always_select_tail {
369 let vis_count = (0..seq)
371 .filter(|k| visible(valid_keys, q_positions[r], *k))
372 .count();
373 let tail_count = vis_count % kp;
374 let tail_start = first_key + vis_count as i64 - tail_count as i64;
375 for t in 0..kp - 1 {
376 let idx = tail_start + t as i64;
377 let ok = t < tail_count
378 && idx >= 0
379 && (idx as usize) < seq
380 && visible(valid_keys, q_positions[r], idx as usize);
381 row[w] = if ok { idx as i32 } else { INVALID };
382 w += 1;
383 }
384 }
385 debug_assert_eq!(
386 w,
387 width.min(select_k * kp + if dims.always_select_tail { kp - 1 } else { 0 })
388 );
389 }
391 out
392}
393
394pub fn topk_to_mask(topk: &[i32], q_rows: usize, width: usize, kv_len: usize) -> Vec<u8> {
399 let mut mask = vec![0u8; q_rows * kv_len];
400 for r in 0..q_rows {
401 for j in 0..width {
402 let i = topk[r * width + j];
403 if i >= 0 && (i as usize) < kv_len {
404 mask[r * kv_len + i as usize] = 1;
405 }
406 }
407 }
408 mask
409}
410
411pub fn mla_masked_attention(
420 q: &[f32],
421 k: &[f32],
422 v: &[f32],
423 mask: &[u8],
424 dims: DsaDims,
425 q_rows: usize,
426 kv_len: usize,
427) -> Vec<f32> {
428 let (h, qd, vd) = (dims.heads, dims.qk_head_dim(), dims.v_head_dim);
429 let scale = (qd as f32).powf(-0.5);
430 let mut out = vec![0.0f32; q_rows * h * vd];
431 for r in 0..q_rows {
432 for hh in 0..h {
433 let mut m = f32::NEG_INFINITY;
435 let mut l = 0.0f32;
436 let mut acc = vec![0.0f32; vd];
437 for kk in 0..kv_len {
438 if mask[r * kv_len + kk] == 0 {
439 continue;
440 }
441 let mut dot = 0.0f32;
442 for dd in 0..qd {
443 dot += q[(r * h + hh) * qd + dd] * k[(kk * h + hh) * qd + dd];
444 }
445 let s = dot * scale;
446 let m_new = m.max(s);
447 let corr = if m == f32::NEG_INFINITY {
448 0.0
449 } else {
450 (m - m_new).exp()
451 };
452 let p = (s - m_new).exp();
453 l = l * corr + p;
454 for dd in 0..vd {
455 acc[dd] = acc[dd] * corr + p * v[(kk * h + hh) * vd + dd];
456 }
457 m = m_new;
458 }
459 let inv = if l > 0.0 { 1.0 / l } else { 0.0 };
460 for dd in 0..vd {
461 out[(r * h + hh) * vd + dd] = acc[dd] * inv;
462 }
463 }
464 }
465 out
466}
467
468pub fn expand_kv(kv_c: &[f32], w_kv_b: &[f32], dims: DsaDims, seq: usize) -> (Vec<f32>, Vec<f32>) {
473 let (h, nope, vd, r) = (
474 dims.heads,
475 dims.qk_nope_head_dim,
476 dims.v_head_dim,
477 dims.kv_lora_rank,
478 );
479 assert_eq!(dims.qk_rope_head_dim, 0, "expand_kv is the NoPE path");
480 let wide = linear(kv_c, seq, r, w_kv_b, h * (nope + vd));
481 let mut k = vec![0.0f32; seq * h * nope];
482 let mut v = vec![0.0f32; seq * h * vd];
483 for t in 0..seq {
484 for hh in 0..h {
485 let src = t * h * (nope + vd) + hh * (nope + vd);
486 k[(t * h + hh) * nope..(t * h + hh) * nope + nope]
487 .copy_from_slice(&wide[src..src + nope]);
488 v[(t * h + hh) * vd..(t * h + hh) * vd + vd]
489 .copy_from_slice(&wide[src + nope..src + nope + vd]);
490 }
491 }
492 (k, v)
493}
494
495pub fn sigmoid_f32(x: f32) -> f32 {
497 sigmoid(x)
498}
499
500#[cfg(test)]
501mod tests;