spark_storage/
attention_ref.rs1use half::bf16;
9
10#[inline]
11fn b2f(x: bf16) -> f32 {
12 x.to_f32()
13}
14
15#[derive(Clone)]
16pub struct AttnState {
17 pub m: Vec<f32>, pub l: Vec<f32>, pub o: Vec<f32>, }
21
22impl AttnState {
23 pub fn new(num_seqs: usize, num_q_heads: usize, head_dim: usize) -> Self {
24 let n_q = num_seqs * num_q_heads;
25 Self {
26 m: vec![f32::NEG_INFINITY; n_q],
27 l: vec![0.0; n_q],
28 o: vec![0.0; n_q * head_dim],
29 }
30 }
31}
32
33#[allow(clippy::too_many_arguments)]
36pub fn step_tile_ref(
37 state: &mut AttnState,
38 q: &[bf16],
39 k_pool: &[bf16],
40 v_pool: &[bf16],
41 tile_blocks: &[i32], tile_block_counts: &[i32], num_seqs: usize,
44 num_q_heads: usize,
45 num_kv_heads: usize,
46 head_dim: usize,
47 block_size: usize,
48 tile_capacity: usize,
49 gqa_ratio: usize,
50) {
51 let inv_sqrt_d = 1.0_f32 / (head_dim as f32).sqrt();
52 let kv_token_stride = num_kv_heads * head_dim;
53 for seq in 0..num_seqs {
54 let n_blocks = tile_block_counts[seq] as usize;
55 for qh in 0..num_q_heads {
56 let kh = qh / gqa_ratio;
57 let q_off = (seq * num_q_heads + qh) * head_dim;
58 let m_idx = seq * num_q_heads + qh;
59 let o_off = m_idx * head_dim;
60 let mut m_run = state.m[m_idx];
61 let mut l_run = state.l[m_idx];
62 let mut o_run: Vec<f32> = state.o[o_off..o_off + head_dim].to_vec();
63 for b in 0..n_blocks {
64 let blk_id = tile_blocks[seq * tile_capacity + b] as usize;
65 let blk_base = blk_id * block_size * kv_token_stride;
66 for t in 0..block_size {
67 let kv_base = blk_base + t * kv_token_stride + kh * head_dim;
68 let mut dot = 0.0_f32;
70 for i in 0..head_dim {
71 dot += b2f(q[q_off + i]) * b2f(k_pool[kv_base + i]);
72 }
73 let logit = dot * inv_sqrt_d;
74 let m_new = m_run.max(logit);
75 let scale_old = (m_run - m_new).exp();
76 let scale_new = (logit - m_new).exp();
77 let l_new = l_run * scale_old + scale_new;
78 for i in 0..head_dim {
79 let v = b2f(v_pool[kv_base + i]);
80 o_run[i] = o_run[i] * scale_old + v * scale_new;
81 }
82 m_run = m_new;
83 l_run = l_new;
84 }
85 }
86 state.m[m_idx] = m_run;
87 state.l[m_idx] = l_run;
88 state.o[o_off..o_off + head_dim].copy_from_slice(&o_run);
89 }
90 }
91}
92
93pub fn finalize_ref(
94 state: &AttnState,
95 num_seqs: usize,
96 num_q_heads: usize,
97 head_dim: usize,
98) -> Vec<bf16> {
99 let mut out = vec![bf16::from_f32(0.0); num_seqs * num_q_heads * head_dim];
100 for seq in 0..num_seqs {
101 for qh in 0..num_q_heads {
102 let idx = seq * num_q_heads + qh;
103 let l = state.l[idx];
104 let inv_l = if l > 0.0 { 1.0 / l } else { 0.0 };
105 for i in 0..head_dim {
106 let val = state.o[idx * head_dim + i] * inv_l;
107 out[idx * head_dim + i] = bf16::from_f32(val);
108 }
109 }
110 }
111 out
112}