spark_storage/
tiled_attention.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// Host wrapper for the tiled decode-attention kernels. Owns the (m, l, o)
4// running-state device buffers and exposes the per-step lifecycle:
5//
6//   1. `begin_step()`   — reset state to (-INF, 0, 0).
7//   2. `step_tile(...)` — one launch per tile of blocks.
8//   3. `finalize(...)`  — divide o by l and store as BF16 output.
9//
10// The state buffers are sized for the maximum (num_seqs, num_q_heads,
11// head_dim) the engine will see, so they can be reused across steps without
12// reallocation.
13
14use anyhow::{Context, Result, bail};
15use std::ffi::c_void;
16
17use crate::cuda_min::{CudaCtx, CudaModule, DeviceBuffer, launch_kernel};
18
19include!(concat!(env!("OUT_DIR"), "/storage_ptx.rs"));
20
21unsafe extern "C" {
22    fn cuMemsetD32Async(dst: u64, value: u32, count: usize, stream: u64) -> i32;
23}
24
25#[derive(Clone, Copy, Debug)]
26pub struct TiledAttentionDims {
27    pub max_seqs: usize,
28    pub num_q_heads: usize,
29    pub num_kv_heads: usize,
30    pub head_dim: usize,
31    pub block_size: usize,
32    pub tile_capacity: usize,
33}
34
35impl TiledAttentionDims {
36    pub fn validate(&self) -> Result<()> {
37        if !self.num_q_heads.is_multiple_of(self.num_kv_heads) {
38            bail!(
39                "num_q_heads ({}) must divide num_kv_heads ({})",
40                self.num_q_heads,
41                self.num_kv_heads
42            );
43        }
44        if self.head_dim > 256 {
45            bail!("head_dim {} exceeds MAX_HEAD_DIM=256", self.head_dim);
46        }
47        Ok(())
48    }
49    pub fn gqa_ratio(&self) -> i32 {
50        (self.num_q_heads / self.num_kv_heads) as i32
51    }
52    fn n_q_slots(&self) -> usize {
53        self.max_seqs * self.num_q_heads
54    }
55    pub fn m_bytes(&self) -> usize {
56        self.n_q_slots() * 4
57    }
58    pub fn l_bytes(&self) -> usize {
59        self.n_q_slots() * 4
60    }
61    pub fn o_bytes(&self) -> usize {
62        self.n_q_slots() * self.head_dim * 4
63    }
64    pub fn output_bytes(&self) -> usize {
65        self.n_q_slots() * self.head_dim * 2
66    }
67}
68
69pub struct TiledAttention {
70    dims: TiledAttentionDims,
71    _modules: Vec<CudaModule>,
72    f_step: u64,
73    f_finalize: u64,
74    pub m_state: DeviceBuffer,
75    pub l_state: DeviceBuffer,
76    pub o_state: DeviceBuffer,
77}
78
79const NEG_INF_F32_BITS: u32 = 0xFF800000;
80
81impl TiledAttention {
82    pub fn new(dims: TiledAttentionDims) -> Result<Self> {
83        dims.validate()?;
84        let mut modules = Vec::new();
85        let mut f_step = 0u64;
86        let mut f_finalize = 0u64;
87        for entry in STORAGE_PTX.iter() {
88            match entry.name {
89                "paged_decode_attn_tiled" => {
90                    let m = CudaModule::from_ptx(entry.ptx)
91                        .with_context(|| format!("load {}", entry.name))?;
92                    f_step = m.function("paged_decode_attn_tiled")?;
93                    modules.push(m);
94                }
95                "attention_finalize" => {
96                    let m = CudaModule::from_ptx(entry.ptx)
97                        .with_context(|| format!("load {}", entry.name))?;
98                    f_finalize = m.function("attention_finalize")?;
99                    modules.push(m);
100                }
101                _ => {}
102            }
103        }
104        if f_step == 0 || f_finalize == 0 {
105            bail!("tiled-attention PTX modules missing");
106        }
107        let m_state = DeviceBuffer::new(dims.m_bytes())?;
108        let l_state = DeviceBuffer::new(dims.l_bytes())?;
109        let o_state = DeviceBuffer::new(dims.o_bytes())?;
110        Ok(Self {
111            dims,
112            _modules: modules,
113            f_step,
114            f_finalize,
115            m_state,
116            l_state,
117            o_state,
118        })
119    }
120
121    /// Reset (m, l, o) for `num_seqs` sequences. Call once at the start of
122    /// each decode step before the first `step_tile`.
123    pub fn begin_step(&self, ctx: &CudaCtx, num_seqs: usize) -> Result<()> {
124        self.begin_step_on_stream(ctx.stream, num_seqs)
125    }
126
127    pub fn begin_step_on_stream(&self, stream: u64, num_seqs: usize) -> Result<()> {
128        if num_seqs > self.dims.max_seqs {
129            bail!(
130                "begin_step num_seqs {} > max_seqs {}",
131                num_seqs,
132                self.dims.max_seqs
133            );
134        }
135        let n_q = num_seqs * self.dims.num_q_heads;
136        let n_o = n_q * self.dims.head_dim;
137        unsafe {
138            let s = cuMemsetD32Async(self.m_state.ptr, NEG_INF_F32_BITS, n_q, stream);
139            if s != 0 {
140                bail!("cuMemsetD32Async m_state failed: {s}");
141            }
142            let s = cuMemsetD32Async(self.l_state.ptr, 0, n_q, stream);
143            if s != 0 {
144                bail!("cuMemsetD32Async l_state failed: {s}");
145            }
146            let s = cuMemsetD32Async(self.o_state.ptr, 0, n_o, stream);
147            if s != 0 {
148                bail!("cuMemsetD32Async o_state failed: {s}");
149            }
150        }
151        Ok(())
152    }
153
154    /// Stride triple for the kernel-native paged K/V layout:
155    /// `[num_blocks, block_size, num_kv_heads, head_dim]`. All in BF16
156    /// elements.
157    pub fn paged_strides(&self) -> (i64, i64, i64) {
158        let blk = (self.dims.block_size * self.dims.num_kv_heads * self.dims.head_dim) as i64;
159        let tok = (self.dims.num_kv_heads * self.dims.head_dim) as i64;
160        let kvh = self.dims.head_dim as i64;
161        (blk, tok, kvh)
162    }
163
164    /// Stride triple for the scratch-pool slot layout
165    /// `[slot, K|V, kv_head, block_size, head_dim]`. Reads of contiguous-
166    /// per-(kv_head) groups stay efficient on disk; the kernel pays a single
167    /// extra multiply per address calculation.
168    pub fn scratch_pool_strides(&self) -> (i64, i64, i64) {
169        let kv_stripe = (self.dims.num_kv_heads * self.dims.block_size * self.dims.head_dim) as i64;
170        let blk = 2 * kv_stripe; // K stripes + V stripes per slot
171        let tok = self.dims.head_dim as i64;
172        let kvh = (self.dims.block_size * self.dims.head_dim) as i64;
173        (blk, tok, kvh)
174    }
175
176    /// One tile of blocks across `num_seqs` sequences. Stride triple
177    /// (`blk_stride`, `tok_stride`, `kvh_stride`) selects the K/V layout —
178    /// see [`paged_strides`](Self::paged_strides) and
179    /// [`scratch_pool_strides`](Self::scratch_pool_strides).
180    #[allow(clippy::too_many_arguments)]
181    pub fn step_tile(
182        &self,
183        ctx: &CudaCtx,
184        q: u64,
185        k_pool: u64,
186        v_pool: u64,
187        tile_blocks: u64,
188        tile_block_counts: u64,
189        num_seqs: usize,
190        blk_stride: i64,
191        tok_stride: i64,
192        kvh_stride: i64,
193        last_block_valid_slots: i32,
194    ) -> Result<()> {
195        self.step_tile_on_stream(
196            ctx.stream,
197            q,
198            k_pool,
199            v_pool,
200            tile_blocks,
201            tile_block_counts,
202            num_seqs,
203            blk_stride,
204            tok_stride,
205            kvh_stride,
206            last_block_valid_slots,
207        )
208    }
209
210    #[allow(clippy::too_many_arguments)]
211    pub fn step_tile_on_stream(
212        &self,
213        stream: u64,
214        q: u64,
215        k_pool: u64,
216        v_pool: u64,
217        tile_blocks: u64,
218        tile_block_counts: u64,
219        num_seqs: usize,
220        blk_stride: i64,
221        tok_stride: i64,
222        kvh_stride: i64,
223        last_block_valid_slots: i32,
224    ) -> Result<()> {
225        let mut q_v = q;
226        let mut k_v = k_pool;
227        let mut v_v = v_pool;
228        let mut tb = tile_blocks;
229        let mut tc = tile_block_counts;
230        let mut m_v = self.m_state.ptr;
231        let mut l_v = self.l_state.ptr;
232        let mut o_v = self.o_state.ptr;
233        let mut nq = self.dims.num_q_heads as i32;
234        let mut nk = self.dims.num_kv_heads as i32;
235        let mut hd = self.dims.head_dim as i32;
236        let mut bs = self.dims.block_size as i32;
237        let mut tcap = self.dims.tile_capacity as i32;
238        let mut gqa = self.dims.gqa_ratio();
239        let mut blk_s = blk_stride;
240        let mut tok_s = tok_stride;
241        let mut kvh_s = kvh_stride;
242        let mut lbvs = last_block_valid_slots;
243        let mut params = [
244            &mut q_v as *mut _ as *mut c_void,
245            &mut k_v as *mut _ as *mut c_void,
246            &mut v_v as *mut _ as *mut c_void,
247            &mut tb as *mut _ as *mut c_void,
248            &mut tc as *mut _ as *mut c_void,
249            &mut m_v as *mut _ as *mut c_void,
250            &mut l_v as *mut _ as *mut c_void,
251            &mut o_v as *mut _ as *mut c_void,
252            &mut nq as *mut _ as *mut c_void,
253            &mut nk as *mut _ as *mut c_void,
254            &mut hd as *mut _ as *mut c_void,
255            &mut bs as *mut _ as *mut c_void,
256            &mut tcap as *mut _ as *mut c_void,
257            &mut gqa as *mut _ as *mut c_void,
258            &mut blk_s as *mut _ as *mut c_void,
259            &mut tok_s as *mut _ as *mut c_void,
260            &mut kvh_s as *mut _ as *mut c_void,
261            &mut lbvs as *mut _ as *mut c_void,
262        ];
263        launch_kernel(
264            self.f_step,
265            (num_seqs as u32, self.dims.num_q_heads as u32, 1),
266            (self.dims.head_dim as u32, 1, 1),
267            0,
268            stream,
269            &mut params,
270        )
271    }
272
273    /// Divide o_state by l_state and store as BF16 in `output`.
274    pub fn finalize(&self, ctx: &CudaCtx, output: u64, num_seqs: usize) -> Result<()> {
275        self.finalize_on_stream(ctx.stream, output, num_seqs)
276    }
277
278    pub fn finalize_on_stream(&self, stream: u64, output: u64, num_seqs: usize) -> Result<()> {
279        let mut l_v = self.l_state.ptr;
280        let mut o_v = self.o_state.ptr;
281        let mut out_v = output;
282        let mut nq = self.dims.num_q_heads as i32;
283        let mut hd = self.dims.head_dim as i32;
284        let mut params = [
285            &mut l_v as *mut _ as *mut c_void,
286            &mut o_v as *mut _ as *mut c_void,
287            &mut out_v as *mut _ as *mut c_void,
288            &mut nq as *mut _ as *mut c_void,
289            &mut hd as *mut _ as *mut c_void,
290        ];
291        launch_kernel(
292            self.f_finalize,
293            (num_seqs as u32, self.dims.num_q_heads as u32, 1),
294            (self.dims.head_dim as u32, 1, 1),
295            0,
296            stream,
297            &mut params,
298        )
299    }
300
301    pub fn dims(&self) -> TiledAttentionDims {
302        self.dims
303    }
304}