spark_model/layers/ops/
prefill_attn_a.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Auto-extracted from `ops.rs` during refactor wave 4a.
4
5#![allow(unused_imports)]
6
7use anyhow::Result;
8use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
9use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
10
11use crate::layers::moe;
12use crate::weight_map::{DenseWeight, Fp8DenseWeight, Fp8Weight, QuantizedWeight};
13
14use super::*;
15
16/// Batched Q rope extract: [N, nq, hd] → [N, nq, rope] at offset nope per head.
17/// 1 kernel replaces N*nq D2D copies per layer.
18#[allow(clippy::too_many_arguments)]
19pub fn mla_q_rope_extract_batched(
20    gpu: &dyn GpuBackend,
21    kernel: KernelHandle,
22    q_full: DevicePtr,
23    q_rope_out: DevicePtr,
24    num_tokens: u32,
25    nq: u32,
26    hd: u32,
27    nope: u32,
28    rope: u32,
29    q_dim: u32,
30    stream: u64,
31) -> Result<()> {
32    let total = num_tokens * nq * rope;
33    KernelLaunch::new(gpu, kernel)
34        .grid([div_ceil(total, 256), 1, 1])
35        .block([256, 1, 1])
36        .arg_ptr(q_full)
37        .arg_ptr(q_rope_out)
38        .arg_u32(num_tokens)
39        .arg_u32(nq)
40        .arg_u32(hd)
41        .arg_u32(nope)
42        .arg_u32(rope)
43        .arg_u32(q_dim)
44        .launch(stream)
45}
46
47/// Batched Q rope writeback: [N, nq, rope] → [N, nq, hd] at offset nope per head.
48#[allow(clippy::too_many_arguments)]
49pub fn mla_q_rope_writeback_batched(
50    gpu: &dyn GpuBackend,
51    kernel: KernelHandle,
52    q_rope_in: DevicePtr,
53    q_full: DevicePtr,
54    num_tokens: u32,
55    nq: u32,
56    hd: u32,
57    nope: u32,
58    rope: u32,
59    q_dim: u32,
60    stream: u64,
61) -> Result<()> {
62    let total = num_tokens * nq * rope;
63    KernelLaunch::new(gpu, kernel)
64        .grid([div_ceil(total, 256), 1, 1])
65        .block([256, 1, 1])
66        .arg_ptr(q_rope_in)
67        .arg_ptr(q_full)
68        .arg_u32(num_tokens)
69        .arg_u32(nq)
70        .arg_u32(hd)
71        .arg_u32(nope)
72        .arg_u32(rope)
73        .arg_u32(q_dim)
74        .launch(stream)
75}
76
77/// Batched K/V assembly from kv_expanded + k_rope for N tokens.
78/// 1 kernel replaces N*nkv*3 D2D copies per layer.
79#[allow(clippy::too_many_arguments)]
80pub fn mla_kv_assemble_batched(
81    gpu: &dyn GpuBackend,
82    kernel: KernelHandle,
83    kv_expanded: DevicePtr,
84    k_rope_buf: DevicePtr,
85    k_out: DevicePtr,
86    v_out: DevicePtr,
87    num_tokens: u32,
88    nkv: u32,
89    nope: u32,
90    v_dim: u32,
91    rope: u32,
92    hd: u32,
93    kv_expanded_stride: u32,
94    stream: u64,
95) -> Result<()> {
96    KernelLaunch::new(gpu, kernel)
97        .grid([num_tokens, 2, 1])
98        .block([256, 1, 1])
99        .arg_ptr(kv_expanded)
100        .arg_ptr(k_rope_buf)
101        .arg_ptr(k_out)
102        .arg_ptr(v_out)
103        .arg_u32(nkv)
104        .arg_u32(nope)
105        .arg_u32(v_dim)
106        .arg_u32(rope)
107        .arg_u32(hd)
108        .arg_u32(kv_expanded_stride)
109        .launch(stream)
110}
111
112/// Batched MLA cache assembly for N tokens: K=[latent|rope], V=[latent|zeros].
113/// 1 kernel replaces N*4 D2D copies+memsets per layer.
114#[allow(clippy::too_many_arguments)]
115pub fn mla_cache_assemble_batched(
116    gpu: &dyn GpuBackend,
117    kernel: KernelHandle,
118    kv_latent: DevicePtr,
119    k_rope: DevicePtr,
120    k_cache: DevicePtr,
121    v_cache: DevicePtr,
122    num_tokens: u32,
123    kv_lora: u32,
124    rope: u32,
125    mla_cache_dim: u32,
126    stream: u64,
127) -> Result<()> {
128    KernelLaunch::new(gpu, kernel)
129        .grid([num_tokens, 1, 1])
130        .block([mla_cache_dim.max(256), 1, 1])
131        .arg_ptr(kv_latent)
132        .arg_ptr(k_rope)
133        .arg_ptr(k_cache)
134        .arg_ptr(v_cache)
135        .arg_u32(kv_lora)
136        .arg_u32(rope)
137        .arg_u32(mla_cache_dim)
138        .launch(stream)
139}
140
141/// Fused MLA prefill: Q_absorption + attention + V_extraction in one kernel.
142/// Grid: (num_heads, seq_len, 1)  Block: (256, 1, 1)
143#[allow(clippy::too_many_arguments)]
144pub fn mla_fused_prefill(
145    gpu: &dyn GpuBackend,
146    kernel: KernelHandle,
147    q_full: DevicePtr,
148    q_rope: DevicePtr,
149    kv_latent: DevicePtr,
150    k_rope: DevicePtr,
151    w_uk: DevicePtr,
152    w_uv: DevicePtr,
153    v_out: DevicePtr,
154    k_cache_out: DevicePtr,
155    v_cache_out: DevicePtr,
156    seq_len: u32,
157    nq: u32,
158    nope: u32,
159    rope: u32,
160    kv_lora: u32,
161    v_dim: u32,
162    hd: u32,
163    num_kv_heads: u32,
164    inv_sqrt_d: f32,
165    stream: u64,
166) -> Result<()> {
167    KernelLaunch::new(gpu, kernel)
168        .grid([nq, seq_len, 1])
169        .block([256, 1, 1])
170        .arg_ptr(q_full)
171        .arg_ptr(q_rope)
172        .arg_ptr(kv_latent)
173        .arg_ptr(k_rope)
174        .arg_ptr(w_uk)
175        .arg_ptr(w_uv)
176        .arg_ptr(v_out)
177        .arg_ptr(k_cache_out)
178        .arg_ptr(v_cache_out)
179        .arg_u32(seq_len)
180        .arg_u32(nq)
181        .arg_u32(nope)
182        .arg_u32(rope)
183        .arg_u32(kv_lora)
184        .arg_u32(v_dim)
185        .arg_u32(hd)
186        .arg_u32(num_kv_heads)
187        .arg_f32(inv_sqrt_d)
188        .launch(stream)
189}
190
191/// Assemble Q_final from Q_absorbed + Q_rope: [absorbed|rope] per head per token.
192#[allow(clippy::too_many_arguments)]
193pub fn mla_q_final_assemble_batched(
194    gpu: &dyn GpuBackend,
195    kernel: KernelHandle,
196    q_absorbed: DevicePtr,
197    q_rope: DevicePtr,
198    q_final: DevicePtr,
199    num_tokens: u32,
200    nq: u32,
201    kv_lora: u32,
202    rope: u32,
203    mla_cache_dim: u32,
204    stream: u64,
205) -> Result<()> {
206    let total = num_tokens * nq * mla_cache_dim;
207    KernelLaunch::new(gpu, kernel)
208        .grid([div_ceil(total, 256), 1, 1])
209        .block([256, 1, 1])
210        .arg_ptr(q_absorbed)
211        .arg_ptr(q_rope)
212        .arg_ptr(q_final)
213        .arg_u32(num_tokens)
214        .arg_u32(nq)
215        .arg_u32(kv_lora)
216        .arg_u32(rope)
217        .arg_u32(mla_cache_dim)
218        .launch(stream)
219}
220
221/// Grouped GEMM for MLA: G independent `[M,K_g]@[N_g,K_g]^T→[M,N_g]` in one launch.
222/// Grid: (M*G, ceil(N_g/4), 1)  Block: (256, 1, 1)
223#[allow(clippy::too_many_arguments)]
224pub fn grouped_gemm_mla(
225    gpu: &dyn GpuBackend,
226    kernel: KernelHandle,
227    a: DevicePtr,
228    b: DevicePtr,
229    c: DevicePtr,
230    m: u32,
231    g: u32,
232    k_g: u32,
233    n_g: u32,
234    a_stride: u32,
235    c_stride: u32,
236    stream: u64,
237) -> Result<()> {
238    KernelLaunch::new(gpu, kernel)
239        .grid([m * g, div_ceil(n_g, 4), 1])
240        .block([256, 1, 1])
241        .arg_ptr(a)
242        .arg_ptr(b)
243        .arg_ptr(c)
244        .arg_u32(m)
245        .arg_u32(g)
246        .arg_u32(k_g)
247        .arg_u32(n_g)
248        .arg_u32(a_stride)
249        .arg_u32(c_stride)
250        .launch(stream)
251}
252
253/// MLA absorbed prefill attention (HDIM=320, simple scalar kernel).
254/// Grid: (num_q_heads, ceil(seq_len/16), batch)  Block: (256, 1, 1)
255#[allow(clippy::too_many_arguments)]
256pub fn mla_prefill_attention_320(
257    gpu: &dyn GpuBackend,
258    kernel: KernelHandle,
259    q: DevicePtr,
260    k: DevicePtr,
261    v: DevicePtr,
262    output: DevicePtr,
263    seq_len: u32,
264    batch: u32,
265    num_q_heads: u32,
266    num_kv_heads: u32,
267    head_dim: u32,
268    inv_sqrt_d: f32,
269    causal: bool,
270    stream: u64,
271) -> Result<()> {
272    let br = 16u32; // MLA_BR in the kernel
273    KernelLaunch::new(gpu, kernel)
274        .grid([num_q_heads, div_ceil(seq_len, br), batch])
275        .block([256, 1, 1])
276        .arg_ptr(q)
277        .arg_ptr(k)
278        .arg_ptr(v)
279        .arg_ptr(output)
280        .arg_u32(seq_len)
281        .arg_u32(num_q_heads)
282        .arg_u32(num_kv_heads)
283        .arg_u32(head_dim)
284        .arg_f32(inv_sqrt_d)
285        .arg_u32(if causal { 1 } else { 0 })
286        .launch(stream)
287}
288
289pub fn paged_decode_attn_bf16(
290    gpu: &dyn GpuBackend,
291    kernel: KernelHandle,
292    q: DevicePtr,
293    k_cache: DevicePtr,
294    v_cache: DevicePtr,
295    output: DevicePtr,
296    block_tables: DevicePtr,
297    seq_lens: DevicePtr,
298    max_blocks_per_seq: u32,
299    num_seqs: u32,
300    num_q_heads: u32,
301    num_kv_heads: u32,
302    head_dim: u32,
303    block_size: u32,
304    inv_sqrt_d: f32,
305    q_stride: u32,
306    sliding_window: u32, // 0 = full attention; >0 = window size (Gemma-4 sliding layers)
307    stream: u64,
308) -> Result<()> {
309    KernelLaunch::new(gpu, kernel)
310        .grid([num_q_heads, num_seqs, 1])
311        .block([256, 1, 1])
312        .arg_ptr(q)
313        .arg_ptr(k_cache)
314        .arg_ptr(v_cache)
315        .arg_ptr(output)
316        .arg_ptr(block_tables)
317        .arg_ptr(seq_lens)
318        .arg_u32(max_blocks_per_seq)
319        .arg_u32(num_q_heads)
320        .arg_u32(num_kv_heads)
321        .arg_u32(head_dim)
322        .arg_u32(block_size)
323        .arg_f32(inv_sqrt_d)
324        .arg_u32(q_stride)
325        .arg_u32(sliding_window)
326        .launch(stream)
327}
328
329pub fn paged_decode_attn_fp8(
330    gpu: &dyn GpuBackend,
331    kernel: KernelHandle,
332    q: DevicePtr,
333    k_cache: DevicePtr,
334    v_cache: DevicePtr,
335    output: DevicePtr,
336    block_tables: DevicePtr,
337    seq_lens: DevicePtr,
338    max_blocks_per_seq: u32,
339    num_seqs: u32,
340    num_q_heads: u32,
341    num_kv_heads: u32,
342    head_dim: u32,
343    block_size: u32,
344    inv_sqrt_d: f32,
345    k_scale: f32,
346    v_scale: f32,
347    q_stride: u32,
348    cache_stride: u64,
349    sliding_window: u32,
350    stream: u64,
351) -> Result<()> {
352    KernelLaunch::new(gpu, kernel)
353        .grid([num_q_heads, num_seqs, 1])
354        .block([256, 1, 1])
355        .arg_ptr(q)
356        .arg_ptr(k_cache)
357        .arg_ptr(v_cache)
358        .arg_ptr(output)
359        .arg_ptr(block_tables)
360        .arg_ptr(seq_lens)
361        .arg_u32(max_blocks_per_seq)
362        .arg_u32(num_q_heads)
363        .arg_u32(num_kv_heads)
364        .arg_u32(head_dim)
365        .arg_u32(block_size)
366        .arg_f32(inv_sqrt_d)
367        .arg_f32(k_scale)
368        .arg_f32(v_scale)
369        .arg_u32(q_stride)
370        .arg_u64(cache_stride)
371        .arg_u32(sliding_window)
372        .launch(stream)
373}