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