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
16pub fn silu_mul(
21 gpu: &dyn GpuBackend,
22 kernel: KernelHandle,
23 gate: DevicePtr,
24 up: DevicePtr,
25 output: DevicePtr,
26 num_elements: u32,
27 stream: u64,
28) -> Result<()> {
29 KernelLaunch::new(gpu, kernel)
30 .grid([div_ceil(num_elements, 256), 1, 1])
31 .block([256, 1, 1])
32 .arg_ptr(gate)
33 .arg_ptr(up)
34 .arg_ptr(output)
35 .arg_u32(num_elements)
36 .launch(stream)
37}
38
39#[allow(clippy::too_many_arguments)]
53pub fn silu_mul_strided(
54 gpu: &dyn GpuBackend,
55 kernel: KernelHandle,
56 gate: DevicePtr,
57 up: DevicePtr,
58 output: DevicePtr,
59 rows: u32,
60 cols: u32,
61 in_stride: u32,
62 out_stride: u32,
63 stream: u64,
64) -> Result<()> {
65 KernelLaunch::new(gpu, kernel)
66 .grid([div_ceil(cols, 256), rows, 1])
67 .block([256, 1, 1])
68 .arg_ptr(gate)
69 .arg_ptr(up)
70 .arg_ptr(output)
71 .arg_u32(rows)
72 .arg_u32(cols)
73 .arg_u32(in_stride)
74 .arg_u32(out_stride)
75 .launch(stream)
76}
77
78#[allow(clippy::too_many_arguments)]
92pub fn silu_mul_quant_fp8(
93 gpu: &dyn GpuBackend,
94 kernel: KernelHandle,
95 gate: DevicePtr,
96 up: DevicePtr,
97 out_fp8: DevicePtr,
98 a_scale: DevicePtr,
99 out_bf16: DevicePtr,
100 m: u32,
101 k: u32,
102 stream: u64,
103) -> Result<()> {
104 KernelLaunch::new(gpu, kernel)
105 .grid([m, 1, 1])
106 .block([128, 1, 1])
107 .arg_ptr(gate)
108 .arg_ptr(up)
109 .arg_ptr(out_fp8)
110 .arg_ptr(a_scale)
111 .arg_ptr(out_bf16)
112 .arg_u32(m)
113 .arg_u32(k)
114 .launch(stream)
115}
116
117pub fn l2_norm(
125 gpu: &dyn GpuBackend,
126 kernel: KernelHandle,
127 data: DevicePtr,
128 num_heads: u32,
129 head_dim: u32,
130 eps: f32,
131 num_tokens: u32,
132 stride: u32,
133 stream: u64,
134) -> Result<()> {
135 KernelLaunch::new(gpu, kernel)
136 .grid([num_heads, num_tokens, 1])
137 .block([head_dim.min(1024), 1, 1])
138 .arg_ptr(data)
139 .arg_u32(head_dim)
140 .arg_f32(eps)
141 .arg_u32(stride)
142 .launch(stream)
143}
144
145pub fn sigmoid_gate_mul(
152 gpu: &dyn GpuBackend,
153 kernel: KernelHandle,
154 input: DevicePtr,
155 gate: DevicePtr,
156 output: DevicePtr,
157 num_elements: u32,
158 stream: u64,
159) -> Result<()> {
160 KernelLaunch::new(gpu, kernel)
161 .grid([div_ceil(num_elements, 256), 1, 1])
162 .block([256, 1, 1])
163 .arg_ptr(input)
164 .arg_ptr(gate)
165 .arg_ptr(output)
166 .arg_u32(num_elements)
167 .launch(stream)
168}
169
170pub fn sigmoid_gate_mul_head_broadcast(
179 gpu: &dyn GpuBackend,
180 kernel: KernelHandle,
181 input: DevicePtr,
182 gate: DevicePtr,
183 output: DevicePtr,
184 nq: u32,
185 hd: u32,
186 num_tokens: u32,
187 stream: u64,
188) -> Result<()> {
189 let total = num_tokens * nq * hd;
190 KernelLaunch::new(gpu, kernel)
191 .grid([div_ceil(total, 256), 1, 1])
192 .block([256, 1, 1])
193 .arg_ptr(input)
194 .arg_ptr(gate)
195 .arg_ptr(output)
196 .arg_u32(nq)
197 .arg_u32(hd)
198 .arg_u32(total)
199 .launch(stream)
200}
201
202#[allow(clippy::too_many_arguments)]
204pub fn softplus_gate_mul_head_broadcast(
205 gpu: &dyn GpuBackend,
206 kernel: KernelHandle,
207 input: DevicePtr,
208 gate: DevicePtr,
209 output: DevicePtr,
210 nq: u32,
211 hd: u32,
212 num_tokens: u32,
213 stream: u64,
214) -> Result<()> {
215 let total = num_tokens * nq * hd;
216 KernelLaunch::new(gpu, kernel)
217 .grid([div_ceil(total, 256), 1, 1])
218 .block([256, 1, 1])
219 .arg_ptr(input)
220 .arg_ptr(gate)
221 .arg_ptr(output)
222 .arg_u32(nq)
223 .arg_u32(hd)
224 .arg_u32(total)
225 .launch(stream)
226}
227
228pub fn residual_add(
233 gpu: &dyn GpuBackend,
234 kernel: KernelHandle,
235 residual: DevicePtr,
236 src: DevicePtr,
237 num_elements: u32,
238 stream: u64,
239) -> Result<()> {
240 KernelLaunch::new(gpu, kernel)
241 .grid([div_ceil(num_elements, 256), 1, 1])
242 .block([256, 1, 1])
243 .arg_ptr(residual)
244 .arg_ptr(src)
245 .arg_u32(num_elements)
246 .launch(stream)
247}
248
249pub fn scaled_add(
254 gpu: &dyn GpuBackend,
255 kernel: KernelHandle,
256 output: DevicePtr,
257 src: DevicePtr,
258 scale: f32,
259 num_elements: u32,
260 stream: u64,
261) -> Result<()> {
262 KernelLaunch::new(gpu, kernel)
263 .grid([div_ceil(num_elements, 256), 1, 1])
264 .block([256, 1, 1])
265 .arg_ptr(output)
266 .arg_ptr(src)
267 .arg_f32(scale)
268 .arg_u32(num_elements)
269 .launch(stream)
270}
271
272pub fn sigmoid_blend(
277 gpu: &dyn GpuBackend,
278 kernel: KernelHandle,
279 output: DevicePtr,
280 src: DevicePtr,
281 sigmoid_gate: f32,
282 num_elements: u32,
283 stream: u64,
284) -> Result<()> {
285 KernelLaunch::new(gpu, kernel)
286 .grid([div_ceil(num_elements, 256), 1, 1])
287 .block([256, 1, 1])
288 .arg_ptr(output)
289 .arg_ptr(src)
290 .arg_f32(sigmoid_gate)
291 .arg_u32(num_elements)
292 .launch(stream)
293}
294
295