spark_model/layers/qwen3_attention/
init.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! `Qwen3AttentionLayer` constructors: `new`, `new_ungated`, and the
4//! private `new_with_gating` (kernel-loading core).
5
6use anyhow::Result;
7use spark_runtime::gpu::{GpuBackend, KernelHandle};
8use spark_runtime::kv_cache::KvCacheDtype;
9
10// `gate` must be called through a real path, not through a `let`-bound
11// function pointer: coercing a `#[track_caller]` fn to a pointer inserts a shim
12// and the audit would name the shim instead of the dispatch site below.
13use super::init_arch_gates::{ArchProbes, gated as gate};
14use super::types::{HeadGateActivation, Qwen3AttentionLayer};
15use crate::layers::FfnComponent;
16use crate::layers::fp8_calibration::Fp8KvCalibration;
17use crate::weight_map::{AttentionWeights, DenseWeight, QuantWeight, QuantizedWeight};
18
19impl Qwen3AttentionLayer {
20    pub fn new(
21        input_norm: DenseWeight,
22        attn: AttentionWeights,
23        post_attn_norm: DenseWeight,
24        ffn: FfnComponent,
25        attn_layer_idx: usize,
26        q_nvfp4: Option<QuantizedWeight>,
27        k_nvfp4: Option<QuantizedWeight>,
28        v_nvfp4: Option<QuantizedWeight>,
29        gpu: &dyn GpuBackend,
30        kv_dtype: KvCacheDtype,
31        fp8_calibration_tokens: usize,
32        config: &atlas_core::config::ModelConfig,
33    ) -> Result<Self> {
34        Self::new_with_gating(
35            input_norm,
36            attn,
37            post_attn_norm,
38            ffn,
39            attn_layer_idx,
40            q_nvfp4,
41            k_nvfp4,
42            v_nvfp4,
43            true,
44            gpu,
45            kv_dtype,
46            fp8_calibration_tokens,
47            config,
48        )
49    }
50
51    pub fn new_ungated(
52        input_norm: DenseWeight,
53        attn: AttentionWeights,
54        post_attn_norm: DenseWeight,
55        ffn: FfnComponent,
56        attn_layer_idx: usize,
57        q_nvfp4: Option<QuantizedWeight>,
58        k_nvfp4: Option<QuantizedWeight>,
59        v_nvfp4: Option<QuantizedWeight>,
60        gpu: &dyn GpuBackend,
61        kv_dtype: KvCacheDtype,
62        fp8_calibration_tokens: usize,
63        config: &atlas_core::config::ModelConfig,
64    ) -> Result<Self> {
65        Self::new_with_gating(
66            input_norm,
67            attn,
68            post_attn_norm,
69            ffn,
70            attn_layer_idx,
71            q_nvfp4,
72            k_nvfp4,
73            v_nvfp4,
74            false,
75            gpu,
76            kv_dtype,
77            fp8_calibration_tokens,
78            config,
79        )
80    }
81
82    #[allow(clippy::too_many_arguments)]
83    fn new_with_gating(
84        input_norm: DenseWeight,
85        attn: AttentionWeights,
86        post_attn_norm: DenseWeight,
87        ffn: FfnComponent,
88        attn_layer_idx: usize,
89        q_nvfp4: Option<QuantizedWeight>,
90        k_nvfp4: Option<QuantizedWeight>,
91        v_nvfp4: Option<QuantizedWeight>,
92        gated: bool,
93        gpu: &dyn GpuBackend,
94        kv_dtype: KvCacheDtype,
95        fp8_calibration_tokens: usize,
96        config: &atlas_core::config::ModelConfig,
97    ) -> Result<Self> {
98        let (reshape_mod, reshape_fn, decode_mod, decode_fn) =
99            super::init_kernel_dispatch::kernel_modules_for_dtype(kv_dtype, config.head_dim);
100        // Which cross-architecture kernel families this config says exist. A
101        // family the model does not have is never LOOKED UP, so it leaves no
102        // failed row in the boot audit. See `init_arch_gates`.
103        let probes = ArchProbes::from_config(config);
104        let mrope_interleaved = config.mrope_interleaved;
105        Ok(Self {
106            input_norm,
107            attn,
108            post_attn_norm,
109            ffn,
110            attn_layer_idx,
111            lora: None,
112            gated,
113            mrope_interleaved,
114            kv_dtype,
115            head_dim_override: None,
116            num_q_heads_override: None,
117            num_kv_heads_override: None,
118            sliding_window: None,
119            rope_theta_override: None,
120            rotary_dim_override: None,
121            rope_proportional: false,
122            attn_scale_override: None,
123            k_eq_v: false,
124            v_norm_weight: None,
125            head_gate_weight: None,
126            head_gate_activation: HeadGateActivation::Sigmoid,
127            sigmoid_gate_head_broadcast_k: super::super::try_kernel(
128                gpu,
129                "residual_add",
130                "sigmoid_gate_mul_head_broadcast",
131            ),
132            softplus_gate_head_broadcast_k: super::super::try_kernel(
133                gpu,
134                "residual_add",
135                "softplus_gate_mul_head_broadcast",
136            ),
137            yarn_inv_freq: spark_runtime::gpu::DevicePtr::NULL,
138            yarn_attention_factor: 1.0,
139            post_attn_out_norm: None,
140            post_ffn_out_norm: None,
141            layer_scalar: None,
142            moe_ffn: None,
143            shortcut_carry_out: None,
144            shortcut_carry_in: None,
145            pre_moe_norm: None,
146            post_moe_out_norm: None,
147            post_dense_ffn_norm: None,
148            sparse_v_threshold: 0.0,
149            q_weight: q_nvfp4.map(QuantWeight::Nvfp4),
150            k_weight: k_nvfp4.map(QuantWeight::Nvfp4),
151            v_weight: v_nvfp4.map(QuantWeight::Nvfp4),
152            o_weight: None,
153            o_dense_bf16: None,
154            mla: None,
155            // ── DeepSeek-V4 Manifold-Constrained Hyper-Connections (mHC) ──
156            // `hc` stays None for non-V4 models; the V4 loader attaches real
157            // HcWeights after this constructor. Kernel handles are lazy (null
158            // when the hyper_connection module is absent), so non-V4 models
159            // still start cleanly.
160            hc: None,
161            qsa: None,
162            hc_pre_k: gate(probes.hyper_connection, gpu, "hyper_connection", "hc_pre"),
163            hc_post_k: gate(probes.hyper_connection, gpu, "hyper_connection", "hc_post"),
164            hc_expand_k: gate(
165                probes.hyper_connection,
166                gpu,
167                "hyper_connection",
168                "hc_expand",
169            ),
170            hc_head_k: gate(probes.hyper_connection, gpu, "hyper_connection", "hc_head"),
171            qkv_nvfp4_t: None,
172            q_nvfp4_t: None,
173            k_nvfp4_t: None,
174            v_nvfp4_t: None,
175            o_nvfp4_t: None,
176            q_fp8w_t: None,
177            k_fp8w_t: None,
178            v_fp8w_t: None,
179            o_fp8w_t: None,
180            w8a16_gemm_t_k: super::super::try_kernel(gpu, "w8a16_gemm_t", "w8a16_gemm_t"),
181            w8a16_gemm_t_pipelined_k: super::super::try_kernel(
182                gpu,
183                "w8a16_gemm_t",
184                "w8a16_gemm_t_pipelined",
185            ),
186            w8a16_gemm_t_m128_k: super::super::try_kernel(
187                gpu,
188                "w8a16_gemm_t_m128",
189                "w8a16_gemm_t_m128",
190            ),
191            per_token_group_quant_fp8_k: super::super::try_kernel(
192                gpu,
193                "per_token_group_quant_fp8",
194                "per_token_group_quant_fp8",
195            ),
196            fp8_gemm_t_blockscaled_k: super::super::try_kernel(
197                gpu,
198                "fp8_gemm_t_blockscaled",
199                "fp8_gemm_t_blockscaled",
200            ),
201            rms_norm_k: gpu.kernel("norm", "rms_norm")?,
202            rms_norm_w_k: if crate::ships_vanilla_norm_weights(config) {
203                gpu.kernel("rms_norm_vanilla", "rms_norm_vanilla")?
204            } else {
205                gpu.kernel("norm", "rms_norm")?
206            },
207            rms_norm_w_warp_row_k: if crate::ships_vanilla_norm_weights(config) {
208                gpu.kernel("rms_norm_vanilla", "rms_norm_vanilla_warp_row")
209                    .unwrap_or(KernelHandle(0))
210            } else {
211                KernelHandle(0)
212            },
213            norm_vanilla: crate::ships_vanilla_norm_weights(config),
214            rms_norm_residual_k: if crate::ships_vanilla_norm_weights(config) {
215                gpu.kernel("norm", "rms_norm_residual_vanilla")?
216            } else {
217                gpu.kernel("norm", "rms_norm_residual")?
218            },
219            dense_gemv_k: gpu.kernel("gemv", "dense_gemv_bf16")?,
220            dequant_q2_0_gn_k: super::super::try_kernel(
221                gpu,
222                "dequant_gguf_bf16",
223                "dequant_q2_0_gn_to_bf16",
224            ),
225            // Resolved by `set_packed_q2_weights`, never here: q2_0_mmq /
226            // the Q8_1 quantizer ship only in GGUF-serving targets, and an
227            // unconditional probe fails the boot audit everywhere else.
228            q2_0_mmq_nc_k: KernelHandle(0),
229            q2_0_mmq_wc_k: KernelHandle(0),
230            q4k_quant_act_k: KernelHandle(0),
231            q2_0_gemv_k: super::super::try_kernel(gpu, "q2_0_gemv_vec", "q2_0_gemv_vec"),
232            dense_gemv_batchm_k: gpu
233                .kernel("dense_gemv_bf16_batchm", "dense_gemv_bf16_batchm")
234                .unwrap_or(KernelHandle(0)),
235            w4a16_gemv_k: gpu.kernel("w4a16_gemv", "w4a16_gemv")?,
236            w4a16_gemv_sw_k: super::super::try_kernel(gpu, "w4a16_gemv", "w4a16_gemv_sw"),
237            w8a16_gemv_k: gpu.kernel("w8a16_gemv", "w8a16_gemv")?,
238            w8a16_gemm_k: super::super::try_kernel(gpu, "w8a16_gemm", "w8a16_gemm"),
239            w8a16_gemm_pipelined_k: super::super::try_kernel(
240                gpu,
241                "w8a16_gemm_pipelined",
242                "w8a16_gemm_pipelined",
243            ),
244            w4a16_gemv_dual_k: gpu.kernel("w4a16_gemv_fused", "w4a16_gemv_dual")?,
245            rope_k: gpu.kernel("rope", "rope_forward")?,
246            rope_strided_k: super::super::try_kernel(gpu, "rope", "rope_forward_strided"),
247            rms_norm_strided_k: super::super::try_kernel(gpu, "norm", "rms_norm_strided"),
248            rope_mrope_interleaved_k: super::super::try_kernel(
249                gpu,
250                "rope_mrope_interleaved",
251                "rope_forward_mrope_interleaved",
252            ),
253            rope_mrope_interleaved_k_only_k: super::super::try_kernel(
254                gpu,
255                "rope_mrope_interleaved",
256                "rope_forward_mrope_interleaved_k_only",
257            ),
258            rope_yarn_k: super::super::try_kernel(gpu, "rope", "rope_forward_yarn"),
259            rope_yarn_scaled_k: super::super::try_kernel(gpu, "rope", "rope_forward_yarn_scaled"),
260            // Interleaved (GPT-J / is_neox_style=False) YaRN RoPE — DeepSeek-V4 MLA.
261            rope_yarn_interleaved_k: super::super::try_kernel(
262                gpu,
263                "rope",
264                "rope_forward_yarn_interleaved",
265            ),
266            rope_yarn_interleaved_inv_k: super::super::try_kernel(
267                gpu,
268                "rope",
269                "rope_forward_yarn_interleaved_inv",
270            ),
271            rope_proportional_k: super::super::try_kernel(gpu, "rope", "rope_forward_proportional"),
272            reshape_cache_k: gpu.kernel(reshape_mod, reshape_fn)?,
273            fused_k_norm_rope_cache_write_bf16_k: super::super::try_kernel(
274                gpu,
275                "fused_k_norm_rope_cache",
276                "fused_k_norm_rope_cache_write_bf16",
277            ),
278            fused_k_norm_rope_mrope_cache_write_bf16_k: super::super::try_kernel(
279                gpu,
280                "fused_k_norm_rope_cache",
281                "fused_k_norm_rope_mrope_cache_write_bf16",
282            ),
283            reshape_and_cache_flash_v_only_k: super::super::try_kernel(
284                gpu,
285                "reshape_and_cache",
286                "reshape_and_cache_flash_v_only",
287            ),
288            wht_bf16_k: super::super::try_kernel(gpu, "wht_bf16", "wht_bf16_inplace"),
289            wht_bf16_k_inv: super::super::try_kernel(gpu, "wht_bf16", "wht_bf16_inplace_inv"),
290            innerq_apply_q_k: super::super::try_kernel(
291                gpu,
292                "tq_plus_innerq_apply",
293                "tq_plus_innerq_apply_q",
294            ),
295            innerq_apply_k_k: super::super::try_kernel(
296                gpu,
297                "tq_plus_innerq_apply",
298                "tq_plus_innerq_apply_k",
299            ),
300            paged_decode_k: gpu.kernel(decode_mod, decode_fn)?,
301            // HDIM>256 decode arm. Every dispatch site gates on
302            // `head_dim > 256 && paged_decode_512_k.0 != 0`, so on a head_dim
303            // 128 model this whole family was a per-dtype probe that could
304            // never be used. `probes.wide_head_dim` is derived from the same
305            // `config.head_dim` those sites read.
306            paged_decode_512_k: match kv_dtype {
307                KvCacheDtype::Bf16 => gate(
308                    probes.wide_head_dim,
309                    gpu,
310                    "paged_decode_attn_512",
311                    "paged_decode_attn",
312                ),
313                KvCacheDtype::Turbo4 => gate(
314                    probes.wide_head_dim,
315                    gpu,
316                    "paged_decode_turbo4_512",
317                    "paged_decode_attn_turbo4",
318                ),
319                KvCacheDtype::Turbo8 => gate(
320                    probes.wide_head_dim,
321                    gpu,
322                    "paged_decode_turbo8_512",
323                    "paged_decode_attn_turbo8",
324                ),
325                KvCacheDtype::Turbo3 | KvCacheDtype::Turbo2 => gate(
326                    probes.wide_head_dim,
327                    gpu,
328                    "paged_decode_turbo4_512",
329                    "paged_decode_attn_turbo4",
330                ),
331                _ => gate(
332                    probes.wide_head_dim,
333                    gpu,
334                    "paged_decode_attn_fp8_512",
335                    "paged_decode_attn_fp8",
336                ),
337            },
338            paged_decode_mla_k: gate(probes.mla, gpu, "paged_decode_mla", "paged_decode_attn"),
339            // DeepSeek-V4-Flash MLA paged decode (compressed 576-dim KV cache).
340            mla_paged_decode_k: gate(
341                probes.mla,
342                gpu,
343                "mla_paged_decode",
344                "mla_paged_decode_nvfp4",
345            ),
346            mla_paged_decode_fp8_k: gate(
347                probes.mla,
348                gpu,
349                "mla_paged_decode_fp8",
350                "mla_paged_decode_fp8",
351            ),
352            mla_batched_gemv_k: gate(probes.mla, gpu, "mla_absorbed", "mla_batched_gemv"),
353            mla_q_rope_scatter_k: gate(probes.mla, gpu, "mla_absorbed", "mla_q_rope_scatter"),
354            mla_q_rope_writeback_k: gate(probes.mla, gpu, "mla_absorbed", "mla_q_rope_writeback"),
355            mla_cache_assemble_k: gate(probes.mla, gpu, "mla_absorbed", "mla_cache_assemble"),
356            mla_q_rope_extract_batched_k: gate(
357                probes.mla,
358                gpu,
359                "mla_absorbed",
360                "mla_q_rope_extract_batched",
361            ),
362            mla_q_rope_writeback_batched_k: gate(
363                probes.mla,
364                gpu,
365                "mla_absorbed",
366                "mla_q_rope_writeback_batched",
367            ),
368            mla_kv_assemble_batched_k: gate(
369                probes.mla,
370                gpu,
371                "mla_absorbed",
372                "mla_kv_assemble_batched",
373            ),
374            mla_cache_assemble_batched_k: gate(
375                probes.mla,
376                gpu,
377                "mla_absorbed",
378                "mla_cache_assemble_batched",
379            ),
380            prefill_attn_mla320_k: gate(
381                probes.mla,
382                gpu,
383                "mla_prefill_attn",
384                "mla_prefill_attn_320",
385            ),
386            grouped_gemm_mla_k: gate(probes.mla, gpu, "grouped_gemm_mla", "grouped_gemm_mla"),
387            mla_q_final_assemble_k: gate(
388                probes.mla,
389                gpu,
390                "mla_absorbed",
391                "mla_q_final_assemble_batched",
392            ),
393            mla_fused_prefill_k: gate(probes.mla, gpu, "mla_fused_prefill", "mla_fused_prefill"),
394            gemm_splitk_partial_k: super::super::try_kernel(
395                gpu,
396                "gemm_splitk",
397                "dense_gemm_splitk_partial",
398            ),
399            gemm_splitk_reduce_k: super::super::try_kernel(
400                gpu,
401                "gemm_splitk",
402                "dense_gemm_splitk_reduce",
403            ),
404            dense_gemm_tc_k: super::super::try_kernel(gpu, "gemm_tc", "dense_gemm_tc"),
405            paged_decode_splitk_k: match kv_dtype {
406                KvCacheDtype::Nvfp4 => {
407                    Some(gpu.kernel("paged_decode_nvfp4", "paged_decode_attn_splitk_nvfp4")?)
408                }
409                KvCacheDtype::Turbo3
410                | KvCacheDtype::Turbo4
411                | KvCacheDtype::Turbo8
412                | KvCacheDtype::Bf16KTurbo3V
413                | KvCacheDtype::Bf16KTurbo4V
414                | KvCacheDtype::Bf16KTurbo2V
415                | KvCacheDtype::Fp8KTurbo3V
416                | KvCacheDtype::Fp8KTurbo4V
417                | KvCacheDtype::Fp8KTurbo2V
418                | KvCacheDtype::Turbo4KTurbo3V
419                | KvCacheDtype::Turbo4KTurbo8V
420                | KvCacheDtype::Turbo3KTurbo8V => None,
421                _ => Some(gpu.kernel("paged_decode_fp8", "paged_decode_attn_splitk_fp8")?),
422            },
423            paged_decode_reduce_k: match kv_dtype {
424                KvCacheDtype::Nvfp4 => {
425                    Some(gpu.kernel("paged_decode_nvfp4", "paged_decode_attn_reduce_nvfp4")?)
426                }
427                KvCacheDtype::Turbo3
428                | KvCacheDtype::Turbo4
429                | KvCacheDtype::Turbo8
430                | KvCacheDtype::Bf16KTurbo3V
431                | KvCacheDtype::Bf16KTurbo4V
432                | KvCacheDtype::Bf16KTurbo2V
433                | KvCacheDtype::Fp8KTurbo3V
434                | KvCacheDtype::Fp8KTurbo4V
435                | KvCacheDtype::Fp8KTurbo2V
436                | KvCacheDtype::Turbo4KTurbo3V
437                | KvCacheDtype::Turbo4KTurbo8V
438                | KvCacheDtype::Turbo3KTurbo8V => None,
439                _ => Some(gpu.kernel("paged_decode_fp8", "paged_decode_attn_reduce_fp8")?),
440            },
441            residual_add_k: gpu.kernel("residual_add", "bf16_residual_add")?,
442            // Gemma-4 rms-norm uses the absolute formula `out = x * rms * w`.
443            rms_norm_f32_in_k: KernelHandle(0),
444            sigmoid_gate_mul_k: gpu.kernel("residual_add", "sigmoid_gate_mul")?,
445            deinterleave_qg_k: gpu.kernel("ssm_preprocess", "deinterleave_qg")?,
446            w4a16_gemv_qg_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_qg")?,
447            residual_add_rms_norm_k: if crate::ships_vanilla_norm_weights(config) {
448                gpu.kernel("norm", "residual_add_rms_norm_vanilla")?
449            } else {
450                gpu.kernel("norm", "residual_add_rms_norm")?
451            },
452            residual_add_rms_norm_gatef32_k: crate::layers::try_kernel(
453                gpu,
454                "norm",
455                "residual_add_rms_norm_gatef32",
456            ),
457            w4a16_gemv_qg_batch2_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_qg_batch2")?,
458            w4a16_gemv_dual_batch2_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_dual_batch2")?,
459            w4a16_gemv_batch2_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_batch2")?,
460            w4a16_gemv_qg_batch3_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_qg_batch3")?,
461            w4a16_gemv_dual_batch3_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_dual_batch3")?,
462            w4a16_gemv_batch3_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_batch3")?,
463            w4a16_batchm: crate::layers::w4a16_gemv_tiers::W4a16BatchmTiers::resolve(gpu),
464            w4a16_gemm_k: gpu.kernel("w4a16", "w4a16_gemm")?,
465            w4a16_gemm_t_k: crate::layers::tgemm_kernel(gpu),
466            w4a16_gemm_t_k64_k: crate::layers::k64_kernel(gpu)?,
467            w4a16_gemm_t_k64_n64_k: crate::layers::k64_n64_kernel(gpu),
468            w4a16_gemm_t_m128_k: gpu.kernel("w4a16", "w4a16_gemm_t_m128")?,
469            w4a16_gemm_t_m128_bf16_k: super::super::try_kernel(
470                gpu,
471                "w4a16",
472                "w4a16_gemm_t_m128_bf16",
473            ),
474            w4a16_gemm_t_m128_v2_k: super::super::w4a16_v2_kernel(gpu),
475            w4a16_gemm_t_m128_v3_k: super::super::w4a16_v3_kernel(gpu),
476            dense_gemm_k: gpu.kernel("gemm", "dense_gemm_bf16")?,
477            dense_gemm_pipelined_k: super::super::try_kernel(
478                gpu,
479                "gemm",
480                "dense_gemm_bf16_pipelined",
481            ),
482            prefill_attn_k: gpu.kernel("inferspark_prefill", "inferspark_prefill")?,
483            // Name comes from the SSOT helper that also supplies the BR the
484            // launcher builds its grid from — see `ops::wide_prefill_kernel`.
485            // Module and entry share a name for both variants.
486            // Resolved WITH FALLBACK — see `ops::wide_prefill_kernel`. A target
487            // that ships only the scalar HDIM=512 kernel must still get it.
488            prefill_attn_512_k: if probes.wide_head_dim {
489                crate::layers::ops::wide_prefill_kernel(gpu).0
490            } else {
491                spark_runtime::gpu::KernelHandle(0)
492            },
493            // BR=32 is the tensor-core instantiation; BR=16 the scalar reference.
494            prefill_attn_512_is_tc: probes.wide_head_dim
495                && crate::layers::ops::wide_prefill_kernel(gpu).1 == 32,
496            // DeepSeek-V4 sparse-attention compressor + compressed-KV prefill.
497            csa_compress_k: gate(probes.compressed_attn, gpu, "csa_compress", "csa_compress"),
498            prefill_attn_compressed_k: gate(
499                probes.compressed_attn,
500                gpu,
501                "prefill_attn_compressed",
502                "prefill_attn_compressed",
503            ),
504            v4_comp_pool_filled: std::sync::atomic::AtomicU32::new(0),
505            v4_comp_prev_valid: std::sync::atomic::AtomicBool::new(false),
506            v4_decode_started: std::sync::atomic::AtomicBool::new(false),
507            v4_decode_first_pos: std::sync::atomic::AtomicU32::new(0),
508            prefill_attn_paged_512_k: gate(
509                probes.wide_head_dim,
510                gpu,
511                "inferspark_prefill_paged_512",
512                "inferspark_prefill_paged_512",
513            ),
514            prefill_attn_64_k: gpu.kernel("inferspark_prefill", "inferspark_prefill_64")?,
515            prefill_attn_paged_k: gpu.kernel("prefill_paged", "inferspark_prefill_paged")?,
516            prefill_attn_paged_fp8_k: gpu
517                .kernel("prefill_paged_fp8", "inferspark_prefill_paged_fp8")?,
518            prefill_attn_paged_nvfp4_k: gpu
519                .kernel("prefill_paged_nvfp4", "inferspark_prefill_paged_nvfp4")?,
520            prefill_attn_paged_turbo4_k: super::super::try_kernel(
521                gpu,
522                "prefill_paged_turbo4",
523                "inferspark_prefill_paged_turbo4",
524            ),
525            prefill_attn_paged_64_k: gpu.kernel("prefill_paged", "inferspark_prefill_paged_64")?,
526            prefill_attn_paged_fp8_64_k: gpu
527                .kernel("prefill_paged_fp8", "inferspark_prefill_paged_fp8_64")?,
528            prefill_attn_paged_nvfp4_64_k: gpu
529                .kernel("prefill_paged_nvfp4", "inferspark_prefill_paged_nvfp4_64")?,
530            prefill_attn_paged_turbo2_64_k: super::super::try_kernel(
531                gpu,
532                "prefill_paged_turbo2",
533                "inferspark_prefill_paged_turbo2",
534            ),
535            prefill_attn_paged_turbo3_64_k: super::super::try_kernel(
536                gpu,
537                "prefill_paged_turbo3",
538                "inferspark_prefill_paged_turbo3_64",
539            ),
540            prefill_attn_paged_turbo4_64_k: super::super::try_kernel(
541                gpu,
542                "prefill_paged_turbo4",
543                "inferspark_prefill_paged_turbo4_64",
544            ),
545            prefill_attn_paged_turbo8_64_k: super::super::try_kernel(
546                gpu,
547                "prefill_paged_turbo8",
548                "inferspark_prefill_paged_turbo8_64",
549            ),
550            // TurboQuant+ safer-asym Bf16K + Turbo3V BR=64 prefill kernel.
551            // Compiled from inferspark_prefill_paged_bf16k_turbo3v.cu which
552            // forks prefill_paged_compute_asym.cuh (LOAD_K_TILE = bf16,
553            // LOAD_V_TILE = turbo3 3-bit dequant).
554            prefill_attn_paged_bf16k_turbo3v_64_k: super::super::try_kernel(
555                gpu,
556                "prefill_paged_bf16k_turbo3v",
557                "inferspark_prefill_paged_bf16k_turbo3v_64",
558            ),
559            // Bf16K + Turbo4V BR=64 prefill (4-bit V dequant in LOAD_V_TILE).
560            prefill_attn_paged_bf16k_turbo4v_64_k: super::super::try_kernel(
561                gpu,
562                "prefill_paged_bf16k_turbo4v",
563                "inferspark_prefill_paged_bf16k_turbo4v_64",
564            ),
565            // Bf16K + Turbo2V BR=64 prefill (2-bit V dequant in LOAD_V_TILE).
566            prefill_attn_paged_bf16k_turbo2v_64_k: super::super::try_kernel(
567                gpu,
568                "prefill_paged_bf16k_turbo2v",
569                "inferspark_prefill_paged_bf16k_turbo2v_64",
570            ),
571            // Fp8K + TurboNV BR=64 prefill kernels — K loaded as FP8 (per-tensor
572            // `k_scale` dequant in LOAD_K_TILE), V as 3/4/2-bit Lloyd-Max packed.
573            prefill_attn_paged_fp8k_turbo3v_64_k: super::super::try_kernel(
574                gpu,
575                "prefill_paged_fp8k_turbo3v",
576                "inferspark_prefill_paged_fp8k_turbo3v_64",
577            ),
578            prefill_attn_paged_fp8k_turbo4v_64_k: super::super::try_kernel(
579                gpu,
580                "prefill_paged_fp8k_turbo4v",
581                "inferspark_prefill_paged_fp8k_turbo4v_64",
582            ),
583            prefill_attn_paged_fp8k_turbo2v_64_k: super::super::try_kernel(
584                gpu,
585                "prefill_paged_fp8k_turbo2v",
586                "inferspark_prefill_paged_fp8k_turbo2v_64",
587            ),
588            // Both-sides-quantized TurboQuant+ asym BR=64 prefill kernels.
589            // K loaded via turbo* dequant in LOAD_K_TILE, V via the corresponding
590            // turbo* dequant in LOAD_V_TILE — separate (block_stride, data_section)
591            // pairs per side.
592            prefill_attn_paged_turbo4k_turbo3v_64_k: super::super::try_kernel(
593                gpu,
594                "prefill_paged_turbo4k_turbo3v",
595                "inferspark_prefill_paged_turbo4k_turbo3v_64",
596            ),
597            prefill_attn_paged_turbo4k_turbo8v_64_k: super::super::try_kernel(
598                gpu,
599                "prefill_paged_turbo4k_turbo8v",
600                "inferspark_prefill_paged_turbo4k_turbo8v_64",
601            ),
602            prefill_attn_paged_turbo3k_turbo8v_64_k: super::super::try_kernel(
603                gpu,
604                "prefill_paged_turbo3k_turbo8v",
605                "inferspark_prefill_paged_turbo3k_turbo8v_64",
606            ),
607            // ── Q12 Phase 3: batched paged-prefill kernel handles ──
608            prefill_attn_paged_batched_k: super::super::try_kernel(
609                gpu,
610                "inferspark_prefill_paged_batched",
611                "inferspark_prefill_paged_batched",
612            ),
613            prefill_attn_paged_fp8_batched_k: super::super::try_kernel(
614                gpu,
615                "inferspark_prefill_paged_fp8_batched",
616                "inferspark_prefill_paged_fp8_batched",
617            ),
618            prefill_attn_paged_nvfp4_batched_k: super::super::try_kernel(
619                gpu,
620                "inferspark_prefill_paged_nvfp4_batched",
621                "inferspark_prefill_paged_nvfp4_batched",
622            ),
623            prefill_attn_paged_batched_64_k: super::super::try_kernel(
624                gpu,
625                "inferspark_prefill_paged_batched",
626                "inferspark_prefill_paged_batched_64",
627            ),
628            prefill_attn_paged_fp8_batched_64_k: super::super::try_kernel(
629                gpu,
630                "inferspark_prefill_paged_fp8_batched",
631                "inferspark_prefill_paged_fp8_batched_64",
632            ),
633            prefill_attn_paged_nvfp4_batched_64_k: super::super::try_kernel(
634                gpu,
635                "inferspark_prefill_paged_nvfp4_batched",
636                "inferspark_prefill_paged_nvfp4_batched_64",
637            ),
638            deinterleave_qg_split_k: gpu.kernel("ssm_preprocess", "deinterleave_qg_split")?,
639            deinterleave_qg_split_qnorm_k: gpu
640                .kernel("ssm_preprocess", "deinterleave_qg_split_qnorm")?,
641            deinterleave_qg_split_qnorm_mrope_k: super::super::try_kernel(
642                gpu,
643                "ssm_preprocess",
644                "deinterleave_qg_split_qnorm_mrope",
645            ),
646            sigmoid_gate_mul_batched_k: gpu.kernel("residual_add", "sigmoid_gate_mul_batched")?,
647            q_fp8: None,
648            k_fp8: None,
649            v_fp8: None,
650            o_fp8: None,
651            fp8_gemm_k: gpu.kernel("w4a16", "fp8_gemm_t")?,
652            bf16_to_fp8_k: gpu.kernel("w4a16", "bf16_to_fp8")?,
653            fp8_fp8_gemm_k: gpu.kernel("w4a16", "fp8_fp8_gemm_t")?,
654            fp8_gemm_t_m128_k: gpu.kernel("w4a16", "fp8_gemm_t_m128")?,
655            fp8_fp8_gemm_t_m128_k: gpu.kernel("w4a16", "fp8_fp8_gemm_t_m128")?,
656            w4a4_gemm_k: crate::layers::try_kernel(gpu, "w4a4", "w4a4_gemm_mfast"),
657            quantize_nvfp4_k: crate::layers::try_kernel(
658                gpu,
659                "quantize_nvfp4",
660                "quantize_bf16_to_nvfp4",
661            ),
662            fp8_calibration: if fp8_calibration_tokens > 0
663                && crate::layers::fp8_calibration::dtype_runs_online_fp8_kv_calibration(kv_dtype)
664            {
665                Some(Fp8KvCalibration::new(
666                    fp8_calibration_tokens,
667                    config.fp8_kv_headroom,
668                    gpu,
669                )?)
670            } else {
671                None
672            },
673        })
674    }
675}