1use anyhow::Result;
28use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
29
30use super::quant_weights::QuantWeights;
31
32mod full_attention;
33mod linear_attention;
34
35pub use full_attention::forward_full_attention;
36pub use linear_attention::forward_linear_attention;
37
38#[derive(Debug, Clone, Copy)]
41pub struct Qwen35ForwardConfig {
42 pub hidden: u32,
44 pub intermediate: u32,
45 pub num_layers: u32,
46 pub vocab: u32,
47 pub group_size: u32,
48 pub rms_eps: f32,
49
50 pub num_heads: u32,
52 pub num_kv_heads: u32,
53 pub head_dim: u32,
54 pub rope_theta: f32,
55 pub rotary_dim: u32,
59
60 pub num_k_heads_lin: u32,
62 pub num_v_heads_lin: u32,
63 pub k_head_dim_lin: u32,
64 pub v_head_dim_lin: u32,
65 pub conv_kernel_size: u32,
66}
67
68impl Qwen35ForwardConfig {
69 pub const fn qwen3_5_4b_mlx_int8() -> Self {
73 Self {
74 hidden: 2560,
75 intermediate: 9216,
76 num_layers: 32,
77 vocab: 248_320,
78 group_size: 64,
79 rms_eps: 1e-6,
80 num_heads: 16,
81 num_kv_heads: 4,
82 head_dim: 256,
83 rope_theta: 10_000_000.0,
84 rotary_dim: 64, num_k_heads_lin: 16,
86 num_v_heads_lin: 32,
87 k_head_dim_lin: 128,
88 v_head_dim_lin: 128,
89 conv_kernel_size: 4,
90 }
91 }
92
93 #[inline]
98 pub const fn q_total(&self) -> u32 {
99 self.num_heads * self.head_dim * 2
100 }
101 #[inline]
104 pub const fn q_only(&self) -> u32 {
105 self.num_heads * self.head_dim
106 }
107 #[inline]
109 pub const fn kv_dim(&self) -> u32 {
110 self.num_kv_heads * self.head_dim
111 }
112 #[inline]
114 pub const fn z_dim_lin(&self) -> u32 {
115 self.num_v_heads_lin * self.v_head_dim_lin
116 }
117 #[inline]
120 pub const fn qkv_total_lin(&self) -> u32 {
121 2 * self.num_k_heads_lin * self.k_head_dim_lin + self.num_v_heads_lin * self.v_head_dim_lin
122 }
123 #[inline]
126 pub const fn num_state_heads(&self) -> u32 {
127 self.num_v_heads_lin
128 }
129}
130
131pub struct Qwen35Kernels {
135 pub rms: KernelHandle,
136 pub rope: KernelHandle,
137 pub kvap: KernelHandle,
138 pub attn: KernelHandle,
139 pub sg: KernelHandle,
140 pub add_rms: KernelHandle,
141 pub qkv_split: KernelHandle,
142 pub conv1d: KernelHandle,
143 pub gdn_gate: KernelHandle,
144 pub sigmoid: KernelHandle,
145 pub gdn_dec: KernelHandle,
146 pub kvap_turbo8: KernelHandle,
151 pub attn_turbo8: KernelHandle,
152 pub kvap_turbo4: KernelHandle,
153 pub attn_turbo4: KernelHandle,
154 pub kvap_turbo3: KernelHandle,
155 pub attn_turbo3: KernelHandle,
156 pub kvap_turbo2: KernelHandle,
157 pub attn_turbo2: KernelHandle,
158 pub kvap_bf16k_turbo4v: KernelHandle,
159 pub attn_bf16k_turbo4v: KernelHandle,
160 pub kvap_bf16k_turbo3v: KernelHandle,
161 pub attn_bf16k_turbo3v: KernelHandle,
162 pub kvap_bf16k_turbo2v: KernelHandle,
163 pub attn_bf16k_turbo2v: KernelHandle,
164 pub wht: KernelHandle,
165 pub wht_inv: KernelHandle,
166}
167
168impl Qwen35Kernels {
169 pub fn resolve(gpu: &dyn GpuBackend) -> Result<Self> {
173 Ok(Self {
174 rms: gpu.kernel("rms_norm", "rms_norm")?,
175 rope: gpu.kernel("rope_apply", "rope_apply")?,
176 kvap: gpu.kernel("kv_cache_append", "kv_cache_append")?,
177 attn: gpu.kernel("attention_decode", "attention_decode")?,
178 sg: gpu.kernel("sigmoid_gate", "sigmoid_gate")?,
179 add_rms: gpu.kernel("add_rms_norm", "add_rms_norm")?,
180 qkv_split: gpu.kernel("qwen35_qkv_split", "qwen35_qkv_split")?,
181 conv1d: gpu.kernel("causal_conv1d_update_l2norm", "causal_conv1d_update_l2norm")?,
182 gdn_gate: gpu.kernel("gdn_helpers", "gdn_compute_gate")?,
183 sigmoid: gpu.kernel("gdn_helpers", "sigmoid_bf16_to_f32")?,
184 gdn_dec: gpu.kernel("gated_delta_rule_decode", "gated_delta_rule_decode")?,
185 kvap_turbo8: gpu.kernel("kv_cache_append_turbo8", "kv_cache_append_turbo8")?,
186 attn_turbo8: gpu.kernel("attention_decode_turbo8", "attention_decode_turbo8")?,
187 kvap_turbo4: gpu.kernel("kv_cache_append_turbo4", "kv_cache_append_turbo4")?,
188 attn_turbo4: gpu.kernel("attention_decode_turbo4", "attention_decode_turbo4")?,
189 kvap_turbo3: gpu.kernel("kv_cache_append_turbo3", "kv_cache_append_turbo3")?,
190 attn_turbo3: gpu.kernel("attention_decode_turbo3", "attention_decode_turbo3")?,
191 kvap_turbo2: gpu.kernel("kv_cache_append_turbo2", "kv_cache_append_turbo2")?,
192 attn_turbo2: gpu.kernel("attention_decode_turbo2", "attention_decode_turbo2")?,
193 kvap_bf16k_turbo4v: gpu.kernel(
194 "kv_cache_append_bf16k_turbov",
195 "kv_cache_append_bf16k_turbo4v",
196 )?,
197 attn_bf16k_turbo4v: gpu.kernel(
198 "attention_decode_bf16k_turbov",
199 "attention_decode_bf16k_turbo4v",
200 )?,
201 kvap_bf16k_turbo3v: gpu.kernel(
202 "kv_cache_append_bf16k_turbov",
203 "kv_cache_append_bf16k_turbo3v",
204 )?,
205 attn_bf16k_turbo3v: gpu.kernel(
206 "attention_decode_bf16k_turbov",
207 "attention_decode_bf16k_turbo3v",
208 )?,
209 kvap_bf16k_turbo2v: gpu.kernel(
210 "kv_cache_append_bf16k_turbov",
211 "kv_cache_append_bf16k_turbo2v",
212 )?,
213 attn_bf16k_turbo2v: gpu.kernel(
214 "attention_decode_bf16k_turbov",
215 "attention_decode_bf16k_turbo2v",
216 )?,
217 wht: gpu.kernel("wht_bf16", "wht_bf16_inplace")?,
218 wht_inv: gpu.kernel("wht_bf16", "wht_bf16_inplace_inv")?,
219 })
220 }
221}
222
223#[derive(Debug, Clone, Copy, PartialEq, Eq)]
225pub enum MetalKvDtype {
226 Bf16,
228 Turbo8,
231 Turbo4,
234 Turbo3,
237 Turbo2,
240 Bf16KTurbo4V,
243 Bf16KTurbo3V,
245 Bf16KTurbo2V,
247}
248
249impl MetalKvDtype {
250 pub fn k_is_rotated(self) -> bool {
253 matches!(
254 self,
255 Self::Turbo8 | Self::Turbo4 | Self::Turbo3 | Self::Turbo2
256 )
257 }
258 pub fn v_is_rotated(self) -> bool {
261 self != Self::Bf16
262 }
263}
264
265impl std::str::FromStr for MetalKvDtype {
266 type Err = anyhow::Error;
267 fn from_str(s: &str) -> Result<Self> {
268 match s {
269 "bf16" => Ok(Self::Bf16),
270 "turbo8" => Ok(Self::Turbo8),
271 "turbo4" => Ok(Self::Turbo4),
272 "turbo3" => Ok(Self::Turbo3),
273 "turbo2" => Ok(Self::Turbo2),
274 "bf16k_turbo4v" => Ok(Self::Bf16KTurbo4V),
275 "bf16k_turbo3v" => Ok(Self::Bf16KTurbo3V),
276 "bf16k_turbo2v" => Ok(Self::Bf16KTurbo2V),
277 other => {
278 anyhow::bail!(
279 "kv dtype {other:?} not supported on metal (bf16 | turbo8 | turbo4 | turbo3 | turbo2 | bf16k_turbo4v/3v/2v)"
280 )
281 }
282 }
283 }
284}
285
286pub struct LayerKvCache {
294 pub k: DevicePtr,
295 pub v: DevicePtr,
296 #[allow(dead_code)]
298 pub capacity: u32,
299 pub dtype: MetalKvDtype,
300 pub k_scales: Option<DevicePtr>,
304 pub v_scales: Option<DevicePtr>,
305}
306
307impl LayerKvCache {
308 pub fn alloc(
310 gpu: &dyn GpuBackend,
311 dtype: MetalKvDtype,
312 max_seq: u32,
313 kv_dim: u32,
314 ) -> Result<Self> {
315 assert!(
316 dtype == MetalKvDtype::Bf16 || kv_dim.is_multiple_of(16),
317 "turbo dtypes need KV_DIM divisible by 16"
318 );
319 let n = (max_seq * kv_dim) as usize;
320 let scale_bytes_e4m3 = (max_seq * kv_dim / 16) as usize;
321 let (kb, vb, ksb, vsb) = match dtype {
323 MetalKvDtype::Bf16 => (n * 2, n * 2, 0, 0),
324 MetalKvDtype::Turbo8 => (n, n, scale_bytes_e4m3 * 2, scale_bytes_e4m3 * 2),
326 MetalKvDtype::Turbo4 => (n / 2, n / 2, scale_bytes_e4m3, scale_bytes_e4m3),
328 MetalKvDtype::Turbo3 => (n * 3 / 8, n * 3 / 8, scale_bytes_e4m3, scale_bytes_e4m3),
330 MetalKvDtype::Turbo2 => (n / 4, n / 4, scale_bytes_e4m3, scale_bytes_e4m3),
332 MetalKvDtype::Bf16KTurbo4V => (n * 2, n / 2, 0, scale_bytes_e4m3),
334 MetalKvDtype::Bf16KTurbo3V => (n * 2, n * 3 / 8, 0, scale_bytes_e4m3),
335 MetalKvDtype::Bf16KTurbo2V => (n * 2, n / 4, 0, scale_bytes_e4m3),
336 };
337 let alloc_opt = |bytes: usize| -> Result<Option<DevicePtr>> {
338 Ok(if bytes > 0 {
339 Some(gpu.alloc(bytes)?)
340 } else {
341 None
342 })
343 };
344 Ok(Self {
345 k: gpu.alloc(kb)?,
346 v: gpu.alloc(vb)?,
347 capacity: max_seq,
348 dtype,
349 k_scales: alloc_opt(ksb)?,
350 v_scales: alloc_opt(vsb)?,
351 })
352 }
353}
354
355pub struct FullAttentionLayer<'a, Q: QuantWeights> {
358 pub input_ln: DevicePtr,
359 pub q_norm: DevicePtr,
360 pub k_norm: DevicePtr,
361 pub post_ln: DevicePtr,
362 pub q_proj: &'a Q,
363 pub k_proj: &'a Q,
364 pub v_proj: &'a Q,
365 pub o_proj: &'a Q,
366 pub gate_proj: &'a Q,
367 pub up_proj: &'a Q,
368 pub down_proj: &'a Q,
369}
370
371pub struct FullAttentionScratch {
373 pub x_norm: DevicePtr,
374 pub q_full: DevicePtr,
375 pub q_split: DevicePtr,
376 pub gate_split: DevicePtr,
377 pub k: DevicePtr,
378 pub v: DevicePtr,
379 pub q_norm_out: DevicePtr,
380 pub k_norm_out: DevicePtr,
381 pub attn_out: DevicePtr,
382 pub gated_attn: DevicePtr,
383 pub o: DevicePtr,
384 pub x_resid: DevicePtr,
385 pub x_norm2: DevicePtr,
386 pub gate_act: DevicePtr,
387 pub up_act: DevicePtr,
388 pub x_out: DevicePtr,
389}
390
391pub struct LinearAttentionLayer<'a, Q: QuantWeights> {
393 pub input_ln: DevicePtr,
394 pub a_log: DevicePtr,
396 pub dt_bias: DevicePtr,
398 pub conv1d_weight: DevicePtr,
400 pub in_proj_a: &'a Q,
401 pub in_proj_b: &'a Q,
402 pub in_proj_qkv: &'a Q,
403 pub in_proj_z: &'a Q,
404 pub norm_weight: DevicePtr,
406 pub out_proj: &'a Q,
407 pub post_ln: DevicePtr,
408 pub gate_proj: &'a Q,
409 pub up_proj: &'a Q,
410 pub down_proj: &'a Q,
411}
412
413pub struct LinearAttentionState {
416 pub conv1d_state: DevicePtr,
418 pub gdn_state: DevicePtr,
420}
421
422pub struct LinearAttentionScratch {
424 pub x_norm: DevicePtr,
425 pub dt_raw: DevicePtr,
426 pub b_raw: DevicePtr,
427 pub qkv: DevicePtr,
428 pub qkv_smooth: DevicePtr,
429 pub z: DevicePtr,
430 pub gate: DevicePtr,
432 pub beta: DevicePtr,
434 pub y: DevicePtr,
435 pub y_norm: DevicePtr,
436 pub out: DevicePtr,
437 pub x_resid: DevicePtr,
438 pub x_norm2: DevicePtr,
439 pub gate_act: DevicePtr,
440 pub up_act: DevicePtr,
441 pub x_final: DevicePtr,
442}