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)]
33pub fn moe_weighted_sum_blend(
34 gpu: &dyn GpuBackend,
35 kernel: KernelHandle,
36 output: DevicePtr,
37 expert_out: DevicePtr,
38 expert_weights: DevicePtr,
39 shared_out: DevicePtr,
40 input: DevicePtr,
41 gate_weight: DevicePtr,
42 hidden: u32,
43 top_k: u32,
44 k: u32,
45 stream: u64,
46) -> Result<()> {
47 KernelLaunch::new(gpu, kernel)
48 .grid([div_ceil(hidden, 256), 1, 1])
49 .block([256, 1, 1])
50 .arg_ptr(output)
51 .arg_ptr(expert_out)
52 .arg_ptr(expert_weights)
53 .arg_ptr(shared_out)
54 .arg_ptr(input)
55 .arg_ptr(gate_weight)
56 .arg_u32(hidden)
57 .arg_u32(top_k)
58 .arg_u32(k)
59 .launch(stream)
60}
61
62#[allow(clippy::too_many_arguments)]
74pub fn moe_expert_gate_up_shared_batch2(
75 gpu: &dyn GpuBackend,
76 kernel: KernelHandle,
77 input: DevicePtr, gate_packed_ptrs: DevicePtr,
79 gate_scale_ptrs: DevicePtr,
80 gate_scale2_vals: DevicePtr,
81 gate_out: DevicePtr, up_packed_ptrs: DevicePtr,
83 up_scale_ptrs: DevicePtr,
84 up_scale2_vals: DevicePtr,
85 up_out: DevicePtr, expert_indices: DevicePtr, sh_gate: &QuantizedWeight,
88 sh_gate_out: DevicePtr, sh_up: &QuantizedWeight,
90 sh_up_out: DevicePtr, n: u32,
92 k: u32,
93 top_k: u32,
94 block_size: u32,
95 stream: u64,
96) -> Result<()> {
97 KernelLaunch::new(gpu, kernel)
98 .grid([div_ceil(n, 8), 2 * (top_k + 1), 2])
99 .block([block_size, 1, 1])
100 .arg_ptr(input)
101 .arg_ptr(gate_packed_ptrs)
102 .arg_ptr(gate_scale_ptrs)
103 .arg_ptr(gate_scale2_vals)
104 .arg_ptr(gate_out)
105 .arg_ptr(up_packed_ptrs)
106 .arg_ptr(up_scale_ptrs)
107 .arg_ptr(up_scale2_vals)
108 .arg_ptr(up_out)
109 .arg_ptr(expert_indices)
110 .arg_ptr(sh_gate.weight)
111 .arg_ptr(sh_gate.weight_scale)
112 .arg_f32(sh_gate.weight_scale_2)
113 .arg_ptr(sh_gate_out)
114 .arg_ptr(sh_up.weight)
115 .arg_ptr(sh_up.weight_scale)
116 .arg_f32(sh_up.weight_scale_2)
117 .arg_ptr(sh_up_out)
118 .arg_u32(n)
119 .arg_u32(k)
120 .arg_u32(top_k)
121 .launch(stream)
122}
123
124#[allow(clippy::too_many_arguments)]
128pub fn moe_expert_silu_down_shared_batch2(
129 gpu: &dyn GpuBackend,
130 kernel: KernelHandle,
131 gate_out: DevicePtr, up_out: DevicePtr, packed_ptrs: DevicePtr,
134 scale_ptrs: DevicePtr,
135 scale2_vals: DevicePtr,
136 output: DevicePtr, expert_indices: DevicePtr, sh_gate_in: DevicePtr, sh_up_in: DevicePtr, sh_down: &QuantizedWeight,
141 sh_down_out: DevicePtr, n: u32,
143 k: u32,
144 top_k: u32,
145 block_size: u32,
146 stream: u64,
147) -> Result<()> {
148 let smem_bytes = (k as usize * std::mem::size_of::<f32>()) as u32;
151 KernelLaunch::new(gpu, kernel)
152 .grid([div_ceil(n, 8), 2 * (top_k + 1), 1])
153 .block([block_size, 1, 1])
154 .shared_mem(smem_bytes)
155 .arg_ptr(gate_out)
156 .arg_ptr(up_out)
157 .arg_ptr(packed_ptrs)
158 .arg_ptr(scale_ptrs)
159 .arg_ptr(scale2_vals)
160 .arg_ptr(output)
161 .arg_ptr(expert_indices)
162 .arg_ptr(sh_gate_in)
163 .arg_ptr(sh_up_in)
164 .arg_ptr(sh_down.weight)
165 .arg_ptr(sh_down.weight_scale)
166 .arg_f32(sh_down.weight_scale_2)
167 .arg_ptr(sh_down_out)
168 .arg_u32(n)
169 .arg_u32(k)
170 .arg_u32(top_k)
171 .launch(stream)
172}
173
174#[allow(clippy::too_many_arguments)]
181pub fn moe_weighted_sum_blend_batch2(
182 gpu: &dyn GpuBackend,
183 kernel: KernelHandle,
184 output: DevicePtr, expert_out: DevicePtr, expert_weights: DevicePtr, shared_out: DevicePtr, input: DevicePtr, gate_weight: DevicePtr, hidden: u32,
191 top_k: u32,
192 k: u32,
193 stream: u64,
194) -> Result<()> {
195 KernelLaunch::new(gpu, kernel)
196 .grid([div_ceil(hidden, 256), 2, 1])
197 .block([256, 1, 1])
198 .arg_ptr(output)
199 .arg_ptr(expert_out)
200 .arg_ptr(expert_weights)
201 .arg_ptr(shared_out)
202 .arg_ptr(input)
203 .arg_ptr(gate_weight)
204 .arg_u32(hidden)
205 .arg_u32(top_k)
206 .arg_u32(k)
207 .launch(stream)
208}
209
210#[allow(clippy::too_many_arguments)]
214pub fn moe_expert_gate_up_shared_batch3(
215 gpu: &dyn GpuBackend,
216 kernel: KernelHandle,
217 input: DevicePtr, gate_packed_ptrs: DevicePtr,
219 gate_scale_ptrs: DevicePtr,
220 gate_scale2_vals: DevicePtr,
221 gate_out: DevicePtr, up_packed_ptrs: DevicePtr,
223 up_scale_ptrs: DevicePtr,
224 up_scale2_vals: DevicePtr,
225 up_out: DevicePtr, expert_indices: DevicePtr, sh_gate: &QuantizedWeight,
228 sh_gate_out: DevicePtr, sh_up: &QuantizedWeight,
230 sh_up_out: DevicePtr, n: u32,
232 k: u32,
233 top_k: u32,
234 stream: u64,
235) -> Result<()> {
236 KernelLaunch::new(gpu, kernel)
237 .grid([div_ceil(n, 8), 3 * (top_k + 1), 2])
238 .block([128, 1, 1])
239 .arg_ptr(input)
240 .arg_ptr(gate_packed_ptrs)
241 .arg_ptr(gate_scale_ptrs)
242 .arg_ptr(gate_scale2_vals)
243 .arg_ptr(gate_out)
244 .arg_ptr(up_packed_ptrs)
245 .arg_ptr(up_scale_ptrs)
246 .arg_ptr(up_scale2_vals)
247 .arg_ptr(up_out)
248 .arg_ptr(expert_indices)
249 .arg_ptr(sh_gate.weight)
250 .arg_ptr(sh_gate.weight_scale)
251 .arg_f32(sh_gate.weight_scale_2)
252 .arg_ptr(sh_gate_out)
253 .arg_ptr(sh_up.weight)
254 .arg_ptr(sh_up.weight_scale)
255 .arg_f32(sh_up.weight_scale_2)
256 .arg_ptr(sh_up_out)
257 .arg_u32(n)
258 .arg_u32(k)
259 .arg_u32(top_k)
260 .launch(stream)
261}
262
263#[allow(clippy::too_many_arguments)]
267pub fn moe_expert_silu_down_shared_batch3(
268 gpu: &dyn GpuBackend,
269 kernel: KernelHandle,
270 gate_out: DevicePtr, up_out: DevicePtr, packed_ptrs: DevicePtr,
273 scale_ptrs: DevicePtr,
274 scale2_vals: DevicePtr,
275 output: DevicePtr, expert_indices: DevicePtr, sh_gate_in: DevicePtr, sh_up_in: DevicePtr, sh_down: &QuantizedWeight,
280 sh_down_out: DevicePtr, n: u32,
282 k: u32,
283 top_k: u32,
284 stream: u64,
285) -> Result<()> {
286 let smem_bytes = (k as usize * std::mem::size_of::<f32>()) as u32;
289 KernelLaunch::new(gpu, kernel)
290 .grid([div_ceil(n, 8), 3 * (top_k + 1), 1])
291 .block([128, 1, 1])
292 .shared_mem(smem_bytes)
293 .arg_ptr(gate_out)
294 .arg_ptr(up_out)
295 .arg_ptr(packed_ptrs)
296 .arg_ptr(scale_ptrs)
297 .arg_ptr(scale2_vals)
298 .arg_ptr(output)
299 .arg_ptr(expert_indices)
300 .arg_ptr(sh_gate_in)
301 .arg_ptr(sh_up_in)
302 .arg_ptr(sh_down.weight)
303 .arg_ptr(sh_down.weight_scale)
304 .arg_f32(sh_down.weight_scale_2)
305 .arg_ptr(sh_down_out)
306 .arg_u32(n)
307 .arg_u32(k)
308 .arg_u32(top_k)
309 .launch(stream)
310}
311
312#[allow(clippy::too_many_arguments)]
316pub fn moe_weighted_sum_blend_batch3(
317 gpu: &dyn GpuBackend,
318 kernel: KernelHandle,
319 output: DevicePtr, expert_out: DevicePtr, expert_weights: DevicePtr, shared_out: DevicePtr, input: DevicePtr, gate_weight: DevicePtr, hidden: u32,
326 top_k: u32,
327 k: u32,
328 stream: u64,
329) -> Result<()> {
330 KernelLaunch::new(gpu, kernel)
331 .grid([div_ceil(hidden, 256), 3, 1])
332 .block([256, 1, 1])
333 .arg_ptr(output)
334 .arg_ptr(expert_out)
335 .arg_ptr(expert_weights)
336 .arg_ptr(shared_out)
337 .arg_ptr(input)
338 .arg_ptr(gate_weight)
339 .arg_u32(hidden)
340 .arg_u32(top_k)
341 .arg_u32(k)
342 .launch(stream)
343}
344
345