spark_storage/
attention_ref.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// Pure-Rust reference for FlashAttention-2 online-softmax decode attention,
4// matching the math the tiled CUDA kernel runs. Used to validate that
5// kernel-side tile fragmentation produces equivalent output to a single-shot
6// computation within float reordering tolerance.
7
8use 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>, // [num_seqs, num_q_heads]
18    pub l: Vec<f32>, // [num_seqs, num_q_heads]
19    pub o: Vec<f32>, // [num_seqs, num_q_heads, head_dim]
20}
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/// One tile update. Mirrors the CUDA kernel's per-token online-softmax
34/// recurrence exactly (same accumulation order, fp32 throughout).
35#[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],       // [num_seqs, tile_capacity]
42    tile_block_counts: &[i32], // [num_seqs]
43    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                    // Compute logit = Q ยท K_t / sqrt(d)
69                    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}