1use 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 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 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 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; 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 #[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 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}