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_quant_fp8(
54 gpu: &dyn GpuBackend,
55 kernel: KernelHandle,
56 gate: DevicePtr,
57 up: DevicePtr,
58 out_fp8: DevicePtr,
59 a_scale: DevicePtr,
60 out_bf16: DevicePtr,
61 m: u32,
62 k: u32,
63 stream: u64,
64) -> Result<()> {
65 KernelLaunch::new(gpu, kernel)
66 .grid([m, 1, 1])
67 .block([128, 1, 1])
68 .arg_ptr(gate)
69 .arg_ptr(up)
70 .arg_ptr(out_fp8)
71 .arg_ptr(a_scale)
72 .arg_ptr(out_bf16)
73 .arg_u32(m)
74 .arg_u32(k)
75 .launch(stream)
76}
77
78pub fn l2_norm(
86 gpu: &dyn GpuBackend,
87 kernel: KernelHandle,
88 data: DevicePtr,
89 num_heads: u32,
90 head_dim: u32,
91 eps: f32,
92 num_tokens: u32,
93 stride: u32,
94 stream: u64,
95) -> Result<()> {
96 KernelLaunch::new(gpu, kernel)
97 .grid([num_heads, num_tokens, 1])
98 .block([head_dim.min(1024), 1, 1])
99 .arg_ptr(data)
100 .arg_u32(head_dim)
101 .arg_f32(eps)
102 .arg_u32(stride)
103 .launch(stream)
104}
105
106pub fn sigmoid_gate_mul(
113 gpu: &dyn GpuBackend,
114 kernel: KernelHandle,
115 input: DevicePtr,
116 gate: DevicePtr,
117 output: DevicePtr,
118 num_elements: u32,
119 stream: u64,
120) -> Result<()> {
121 KernelLaunch::new(gpu, kernel)
122 .grid([div_ceil(num_elements, 256), 1, 1])
123 .block([256, 1, 1])
124 .arg_ptr(input)
125 .arg_ptr(gate)
126 .arg_ptr(output)
127 .arg_u32(num_elements)
128 .launch(stream)
129}
130
131pub fn sigmoid_gate_mul_head_broadcast(
140 gpu: &dyn GpuBackend,
141 kernel: KernelHandle,
142 input: DevicePtr,
143 gate: DevicePtr,
144 output: DevicePtr,
145 nq: u32,
146 hd: u32,
147 num_tokens: u32,
148 stream: u64,
149) -> Result<()> {
150 let total = num_tokens * nq * hd;
151 KernelLaunch::new(gpu, kernel)
152 .grid([div_ceil(total, 256), 1, 1])
153 .block([256, 1, 1])
154 .arg_ptr(input)
155 .arg_ptr(gate)
156 .arg_ptr(output)
157 .arg_u32(nq)
158 .arg_u32(hd)
159 .arg_u32(total)
160 .launch(stream)
161}
162
163#[allow(clippy::too_many_arguments)]
165pub fn softplus_gate_mul_head_broadcast(
166 gpu: &dyn GpuBackend,
167 kernel: KernelHandle,
168 input: DevicePtr,
169 gate: DevicePtr,
170 output: DevicePtr,
171 nq: u32,
172 hd: u32,
173 num_tokens: u32,
174 stream: u64,
175) -> Result<()> {
176 let total = num_tokens * nq * hd;
177 KernelLaunch::new(gpu, kernel)
178 .grid([div_ceil(total, 256), 1, 1])
179 .block([256, 1, 1])
180 .arg_ptr(input)
181 .arg_ptr(gate)
182 .arg_ptr(output)
183 .arg_u32(nq)
184 .arg_u32(hd)
185 .arg_u32(total)
186 .launch(stream)
187}
188
189pub fn residual_add(
194 gpu: &dyn GpuBackend,
195 kernel: KernelHandle,
196 residual: DevicePtr,
197 src: DevicePtr,
198 num_elements: u32,
199 stream: u64,
200) -> Result<()> {
201 KernelLaunch::new(gpu, kernel)
202 .grid([div_ceil(num_elements, 256), 1, 1])
203 .block([256, 1, 1])
204 .arg_ptr(residual)
205 .arg_ptr(src)
206 .arg_u32(num_elements)
207 .launch(stream)
208}
209
210pub fn scaled_add(
215 gpu: &dyn GpuBackend,
216 kernel: KernelHandle,
217 output: DevicePtr,
218 src: DevicePtr,
219 scale: f32,
220 num_elements: u32,
221 stream: u64,
222) -> Result<()> {
223 KernelLaunch::new(gpu, kernel)
224 .grid([div_ceil(num_elements, 256), 1, 1])
225 .block([256, 1, 1])
226 .arg_ptr(output)
227 .arg_ptr(src)
228 .arg_f32(scale)
229 .arg_u32(num_elements)
230 .launch(stream)
231}
232
233pub fn sigmoid_blend(
238 gpu: &dyn GpuBackend,
239 kernel: KernelHandle,
240 output: DevicePtr,
241 src: DevicePtr,
242 sigmoid_gate: f32,
243 num_elements: u32,
244 stream: u64,
245) -> Result<()> {
246 KernelLaunch::new(gpu, kernel)
247 .grid([div_ceil(num_elements, 256), 1, 1])
248 .block([256, 1, 1])
249 .arg_ptr(output)
250 .arg_ptr(src)
251 .arg_f32(sigmoid_gate)
252 .arg_u32(num_elements)
253 .launch(stream)
254}
255
256