1#![allow(unused_imports)]
8
9use anyhow::Result;
10use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
11use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
12
13use crate::layers::moe;
14use crate::weight_map::{DenseWeight, Fp8DenseWeight, Fp8Weight, QuantizedWeight};
15
16use super::*;
17
18#[allow(clippy::too_many_arguments)]
29pub fn gdn_prefill(
30 gpu: &dyn GpuBackend,
31 kernel: KernelHandle,
32 h_state: DevicePtr,
33 query: DevicePtr,
34 key: DevicePtr,
35 value: DevicePtr,
36 gate: DevicePtr,
37 beta: DevicePtr,
38 output: DevicePtr,
39 batch_size: u32,
40 seq_len: u32,
41 num_k_heads: u32,
42 num_v_heads: u32,
43 k_dim: u32,
44 v_dim: u32,
45 qk_stride: u32,
46 v_stride: u32,
47 gb_stride: u32,
48 stream: u64,
49) -> Result<()> {
50 KernelLaunch::new(gpu, kernel)
51 .grid([num_v_heads, batch_size, 1])
52 .block([128, 1, 1])
53 .shared_mem(4 * k_dim * 4) .arg_ptr(h_state)
55 .arg_ptr(query)
56 .arg_ptr(key)
57 .arg_ptr(value)
58 .arg_ptr(gate)
59 .arg_ptr(beta)
60 .arg_ptr(output)
61 .arg_u32(batch_size)
62 .arg_u32(seq_len)
63 .arg_u32(num_k_heads)
64 .arg_u32(num_v_heads)
65 .arg_u32(k_dim)
66 .arg_u32(v_dim)
67 .arg_u32(qk_stride)
68 .arg_u32(v_stride)
69 .arg_u32(gb_stride)
70 .launch(stream)
71}
72
73#[allow(clippy::too_many_arguments)]
80pub fn gdn_prefill_split(
81 gpu: &dyn GpuBackend,
82 kernel: KernelHandle,
83 h_state: DevicePtr,
84 query: DevicePtr,
85 key: DevicePtr,
86 value: DevicePtr,
87 gate: DevicePtr,
88 beta: DevicePtr,
89 output: DevicePtr,
90 batch_size: u32,
91 seq_len: u32,
92 num_k_heads: u32,
93 num_v_heads: u32,
94 k_dim: u32,
95 v_dim: u32,
96 qk_stride: u32,
97 v_stride: u32,
98 gb_stride: u32,
99 stream: u64,
100) -> Result<()> {
101 KernelLaunch::new(gpu, kernel)
102 .grid([num_v_heads * 2, batch_size, 1])
103 .block([64, 1, 1])
104 .shared_mem(4 * k_dim * 4) .arg_ptr(h_state)
106 .arg_ptr(query)
107 .arg_ptr(key)
108 .arg_ptr(value)
109 .arg_ptr(gate)
110 .arg_ptr(beta)
111 .arg_ptr(output)
112 .arg_u32(batch_size)
113 .arg_u32(seq_len)
114 .arg_u32(num_k_heads)
115 .arg_u32(num_v_heads)
116 .arg_u32(k_dim)
117 .arg_u32(v_dim)
118 .arg_u32(qk_stride)
119 .arg_u32(v_stride)
120 .arg_u32(gb_stride)
121 .launch(stream)
122}
123
124#[allow(clippy::too_many_arguments)]
131pub fn gdn_prefill_split4(
132 gpu: &dyn GpuBackend,
133 kernel: KernelHandle,
134 h_state: DevicePtr,
135 query: DevicePtr,
136 key: DevicePtr,
137 value: DevicePtr,
138 gate: DevicePtr,
139 beta: DevicePtr,
140 output: DevicePtr,
141 batch_size: u32,
142 seq_len: u32,
143 num_k_heads: u32,
144 num_v_heads: u32,
145 k_dim: u32,
146 v_dim: u32,
147 qk_stride: u32,
148 v_stride: u32,
149 gb_stride: u32,
150 stream: u64,
151) -> Result<()> {
152 KernelLaunch::new(gpu, kernel)
153 .grid([num_v_heads * 4, batch_size, 1])
154 .block([32, 1, 1])
155 .shared_mem(4 * k_dim * 4) .arg_ptr(h_state)
157 .arg_ptr(query)
158 .arg_ptr(key)
159 .arg_ptr(value)
160 .arg_ptr(gate)
161 .arg_ptr(beta)
162 .arg_ptr(output)
163 .arg_u32(batch_size)
164 .arg_u32(seq_len)
165 .arg_u32(num_k_heads)
166 .arg_u32(num_v_heads)
167 .arg_u32(k_dim)
168 .arg_u32(v_dim)
169 .arg_u32(qk_stride)
170 .arg_u32(v_stride)
171 .arg_u32(gb_stride)
172 .launch(stream)
173}
174
175#[allow(clippy::too_many_arguments)]
185pub fn gdn_prefill_persistent(
186 gpu: &dyn GpuBackend,
187 kernel: KernelHandle,
188 h_state: DevicePtr,
189 query: DevicePtr,
190 key: DevicePtr,
191 value: DevicePtr,
192 gate: DevicePtr,
193 beta: DevicePtr,
194 output: DevicePtr,
195 batch_size: u32,
196 seq_len: u32,
197 num_k_heads: u32,
198 num_v_heads: u32,
199 k_dim: u32,
200 v_dim: u32,
201 qk_stride: u32,
202 v_stride: u32,
203 gb_stride: u32,
204 stream: u64,
205) -> Result<()> {
206 let smem = k_dim * v_dim * 4 + 4 * k_dim * 4; KernelLaunch::new(gpu, kernel)
208 .grid([num_v_heads, batch_size, 1])
209 .block([128, 1, 1])
210 .shared_mem(smem)
211 .arg_ptr(h_state)
212 .arg_ptr(query)
213 .arg_ptr(key)
214 .arg_ptr(value)
215 .arg_ptr(gate)
216 .arg_ptr(beta)
217 .arg_ptr(output)
218 .arg_u32(batch_size)
219 .arg_u32(seq_len)
220 .arg_u32(num_k_heads)
221 .arg_u32(num_v_heads)
222 .arg_u32(k_dim)
223 .arg_u32(v_dim)
224 .arg_u32(qk_stride)
225 .arg_u32(v_stride)
226 .arg_u32(gb_stride)
227 .launch(stream)
228}
229
230#[allow(clippy::too_many_arguments)]
233pub fn gdn_prefill_persistent_smem(
234 gpu: &dyn GpuBackend,
235 kernel: KernelHandle,
236 h_state: DevicePtr,
237 query: DevicePtr,
238 key: DevicePtr,
239 value: DevicePtr,
240 gate: DevicePtr,
241 beta: DevicePtr,
242 output: DevicePtr,
243 batch_size: u32,
244 seq_len: u32,
245 num_k_heads: u32,
246 num_v_heads: u32,
247 k_dim: u32,
248 v_dim: u32,
249 qk_stride: u32,
250 v_stride: u32,
251 gb_stride: u32,
252 smem: u32,
253 stream: u64,
254) -> Result<()> {
255 KernelLaunch::new(gpu, kernel)
256 .grid([num_v_heads, batch_size, 1])
257 .block([128, 1, 1])
258 .shared_mem(smem)
259 .arg_ptr(h_state)
260 .arg_ptr(query)
261 .arg_ptr(key)
262 .arg_ptr(value)
263 .arg_ptr(gate)
264 .arg_ptr(beta)
265 .arg_ptr(output)
266 .arg_u32(batch_size)
267 .arg_u32(seq_len)
268 .arg_u32(num_k_heads)
269 .arg_u32(num_v_heads)
270 .arg_u32(k_dim)
271 .arg_u32(v_dim)
272 .arg_u32(qk_stride)
273 .arg_u32(v_stride)
274 .arg_u32(gb_stride)
275 .launch(stream)
276}