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, present};
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            // `Fp8ActQuant` probes the shared quantizer AND the Hopper
192            // twin, which only `kernels/hopper` ships, and carries both
193            // handles so a launcher can never pair one kernel's entry point
194            // with the other's grid. Every target still has the shared one.
195            // The shared name is the one `W8A8_PREFILL_KERNELS[0]` (#915)
196            // spells for preflight, which derives it from the same constants.
197            per_token_group_quant_fp8_k: crate::layers::ops::Fp8ActQuant::resolve(gpu),
198            fp8_gemm_t_blockscaled_k: super::super::try_kernel(
199                gpu,
200                super::types_weights::W8A8_PREFILL_KERNELS[1].0,
201                super::types_weights::W8A8_PREFILL_KERNELS[1].1,
202            ),
203            // Same optional adapter the SSM layer loads (`init.rs`): absent on
204            // a shadow that has no `fp8_scale_transpose` module, which makes
205            // the cuBLASLt W8A8 arms decline rather than hand the library the
206            // wrong scale order.
207            fp8_act_scale_kmajor_k: super::super::try_kernel(
208                gpu,
209                "fp8_scale_transpose",
210                "fp8_act_scale_to_kmajor",
211            ),
212            rms_norm_k: gpu.kernel("norm", "rms_norm")?,
213            rms_norm_w_k: if crate::ships_vanilla_norm_weights(config) {
214                gpu.kernel("rms_norm_vanilla", "rms_norm_vanilla")?
215            } else {
216                gpu.kernel("norm", "rms_norm")?
217            },
218            rms_norm_w_warp_row_k: if crate::ships_vanilla_norm_weights(config) {
219                gpu.kernel("rms_norm_vanilla", "rms_norm_vanilla_warp_row")
220                    .unwrap_or(KernelHandle(0))
221            } else {
222                KernelHandle(0)
223            },
224            norm_vanilla: crate::ships_vanilla_norm_weights(config),
225            rms_norm_residual_k: if crate::ships_vanilla_norm_weights(config) {
226                gpu.kernel("norm", "rms_norm_residual_vanilla")?
227            } else {
228                gpu.kernel("norm", "rms_norm_residual")?
229            },
230            dense_gemv_k: gpu.kernel("gemv", "dense_gemv_bf16")?,
231            dequant_q2_0_gn_k: super::super::try_kernel(
232                gpu,
233                "dequant_gguf_bf16",
234                "dequant_q2_0_gn_to_bf16",
235            ),
236            // Resolved by `set_packed_q2_weights`, never here: q2_0_mmq /
237            // the Q8_1 quantizer ship only in GGUF-serving targets, and an
238            // unconditional probe fails the boot audit everywhere else.
239            q2_0_mmq_nc_k: KernelHandle(0),
240            q2_0_mmq_wc_k: KernelHandle(0),
241            q4k_quant_act_k: KernelHandle(0),
242            q2_0_gemv_k: super::super::try_kernel(gpu, "q2_0_gemv_vec", "q2_0_gemv_vec"),
243            dense_gemv_batchm_k: gpu
244                .kernel("dense_gemv_bf16_batchm", "dense_gemv_bf16_batchm")
245                .unwrap_or(KernelHandle(0)),
246            w4a16_gemv_k: gpu.kernel("w4a16_gemv", "w4a16_gemv")?,
247            w4a16_gemv_sw_k: super::super::try_kernel(gpu, "w4a16_gemv", "w4a16_gemv_sw"),
248            w8a16_gemv_k: gpu.kernel("w8a16_gemv", "w8a16_gemv")?,
249            w8a16_gemv_batch4_k: super::super::try_kernel(
250                gpu,
251                "w8a16_gemv_batch4",
252                "w8a16_gemv_batch4",
253            ),
254            w8a16_gemv_batch16_k: super::super::try_kernel(
255                gpu,
256                "w8a16_gemv_batch4",
257                "w8a16_gemv_batch16",
258            ),
259            w8a16_gemv_batch4_strided_k: super::super::try_kernel(
260                gpu,
261                "w8a16_gemv_batch4",
262                "w8a16_gemv_batch4_strided",
263            ),
264            w8a16_gemv_batch16_strided_k: super::super::try_kernel(
265                gpu,
266                "w8a16_gemv_batch4",
267                "w8a16_gemv_batch16_strided",
268            ),
269            w8a16_gemm_m16_k: super::super::try_target_kernel(
270                gpu,
271                "w8a16_gemm_m16",
272                "w8a16_gemm_m16",
273            ),
274            w8a16_gemm_m16_strided_k: super::super::try_target_kernel(
275                gpu,
276                "w8a16_gemm_m16",
277                "w8a16_gemm_m16_strided",
278            ),
279            m16_tc: crate::layers::dense_ffn::m16_tc::m16_tc_levers().attn,
280            w8a16_gemv_ncol2_k: super::super::try_target_kernel(
281                gpu,
282                "w8a16_gemv_ncol",
283                "w8a16_gemv_batch16_ncol2",
284            ),
285            w8a16_gemv_ncol4_k: super::super::try_target_kernel(
286                gpu,
287                "w8a16_gemv_ncol",
288                "w8a16_gemv_batch16_ncol4",
289            ),
290            w8a16_gemv_ncol2_strided_k: super::super::try_target_kernel(
291                gpu,
292                "w8a16_gemv_ncol",
293                "w8a16_gemv_batch16_ncol2_strided",
294            ),
295            w8a16_gemv_ncol4_strided_k: super::super::try_target_kernel(
296                gpu,
297                "w8a16_gemv_ncol",
298                "w8a16_gemv_batch16_ncol4_strided",
299            ),
300            attn_ncol: super::attn_ncol_gemv::ncol_gemv_enabled()
301                .then(super::attn_ncol_gemv::ncol_gemv_width),
302            w8a16_gemm_k: super::super::try_kernel(gpu, "w8a16_gemm", "w8a16_gemm"),
303            w8a16_gemm_pipelined_k: super::super::try_kernel(
304                gpu,
305                "w8a16_gemm_pipelined",
306                "w8a16_gemm_pipelined",
307            ),
308            w4a16_gemv_dual_k: gpu.kernel("w4a16_gemv_fused", "w4a16_gemv_dual")?,
309            rope_k: gpu.kernel("rope", "rope_forward")?,
310            rope_strided_k: super::super::try_kernel(gpu, "rope", "rope_forward_strided"),
311            rms_norm_strided_k: super::super::try_kernel(gpu, "norm", "rms_norm_strided"),
312            rope_mrope_interleaved_k: super::super::try_kernel(
313                gpu,
314                "rope_mrope_interleaved",
315                "rope_forward_mrope_interleaved",
316            ),
317            rope_mrope_interleaved_k_only_k: super::super::try_kernel(
318                gpu,
319                "rope_mrope_interleaved",
320                "rope_forward_mrope_interleaved_k_only",
321            ),
322            rope_yarn_k: super::super::try_kernel(gpu, "rope", "rope_forward_yarn"),
323            rope_yarn_scaled_k: super::super::try_kernel(gpu, "rope", "rope_forward_yarn_scaled"),
324            // Interleaved (GPT-J / is_neox_style=False) YaRN RoPE — DeepSeek-V4 MLA.
325            rope_yarn_interleaved_k: super::super::try_kernel(
326                gpu,
327                "rope",
328                "rope_forward_yarn_interleaved",
329            ),
330            rope_yarn_interleaved_inv_k: super::super::try_kernel(
331                gpu,
332                "rope",
333                "rope_forward_yarn_interleaved_inv",
334            ),
335            rope_proportional_k: super::super::try_kernel(gpu, "rope", "rope_forward_proportional"),
336            reshape_cache_k: gpu.kernel(reshape_mod, reshape_fn)?,
337            fused_k_norm_rope_cache_write_bf16_k: super::super::try_kernel(
338                gpu,
339                "fused_k_norm_rope_cache",
340                "fused_k_norm_rope_cache_write_bf16",
341            ),
342            fused_k_norm_rope_mrope_cache_write_bf16_k: super::super::try_kernel(
343                gpu,
344                "fused_k_norm_rope_cache",
345                "fused_k_norm_rope_mrope_cache_write_bf16",
346            ),
347            reshape_and_cache_flash_v_only_k: super::super::try_kernel(
348                gpu,
349                "reshape_and_cache",
350                "reshape_and_cache_flash_v_only",
351            ),
352            wht_bf16_k: super::super::try_kernel(gpu, "wht_bf16", "wht_bf16_inplace"),
353            wht_bf16_k_inv: super::super::try_kernel(gpu, "wht_bf16", "wht_bf16_inplace_inv"),
354            innerq_apply_q_k: super::super::try_kernel(
355                gpu,
356                "tq_plus_innerq_apply",
357                "tq_plus_innerq_apply_q",
358            ),
359            innerq_apply_k_k: super::super::try_kernel(
360                gpu,
361                "tq_plus_innerq_apply",
362                "tq_plus_innerq_apply_k",
363            ),
364            paged_decode_k: gpu.kernel(decode_mod, decode_fn)?,
365            // HDIM>256 decode arm. Every dispatch site gates on
366            // `head_dim > 256 && paged_decode_512_k.0 != 0`, so on a head_dim
367            // 128 model this whole family was a per-dtype probe that could
368            // never be used. `probes.wide_head_dim` is derived from the same
369            // `config.head_dim` those sites read.
370            paged_decode_512_k: match kv_dtype {
371                KvCacheDtype::Bf16 => gate(
372                    probes.wide_head_dim,
373                    gpu,
374                    "paged_decode_attn_512",
375                    "paged_decode_attn",
376                ),
377                KvCacheDtype::Turbo4 => gate(
378                    probes.wide_head_dim,
379                    gpu,
380                    "paged_decode_turbo4_512",
381                    "paged_decode_attn_turbo4",
382                ),
383                KvCacheDtype::Turbo8 => gate(
384                    probes.wide_head_dim,
385                    gpu,
386                    "paged_decode_turbo8_512",
387                    "paged_decode_attn_turbo8",
388                ),
389                KvCacheDtype::Turbo3 | KvCacheDtype::Turbo2 => gate(
390                    probes.wide_head_dim,
391                    gpu,
392                    "paged_decode_turbo4_512",
393                    "paged_decode_attn_turbo4",
394                ),
395                _ => gate(
396                    probes.wide_head_dim,
397                    gpu,
398                    "paged_decode_attn_fp8_512",
399                    "paged_decode_attn_fp8",
400                ),
401            },
402            paged_decode_mla_k: gate(probes.mla, gpu, "paged_decode_mla", "paged_decode_attn"),
403            // DeepSeek-V4-Flash MLA paged decode (compressed 576-dim KV cache).
404            mla_paged_decode_k: gate(
405                probes.mla,
406                gpu,
407                "mla_paged_decode",
408                "mla_paged_decode_nvfp4",
409            ),
410            mla_paged_decode_fp8_k: gate(
411                probes.mla,
412                gpu,
413                "mla_paged_decode_fp8",
414                "mla_paged_decode_fp8",
415            ),
416            mla_batched_gemv_k: gate(probes.mla, gpu, "mla_absorbed", "mla_batched_gemv"),
417            mla_q_rope_scatter_k: gate(probes.mla, gpu, "mla_absorbed", "mla_q_rope_scatter"),
418            mla_q_rope_writeback_k: gate(probes.mla, gpu, "mla_absorbed", "mla_q_rope_writeback"),
419            mla_cache_assemble_k: gate(probes.mla, gpu, "mla_absorbed", "mla_cache_assemble"),
420            mla_q_rope_extract_batched_k: gate(
421                probes.mla,
422                gpu,
423                "mla_absorbed",
424                "mla_q_rope_extract_batched",
425            ),
426            mla_q_rope_writeback_batched_k: gate(
427                probes.mla,
428                gpu,
429                "mla_absorbed",
430                "mla_q_rope_writeback_batched",
431            ),
432            mla_kv_assemble_batched_k: gate(
433                probes.mla,
434                gpu,
435                "mla_absorbed",
436                "mla_kv_assemble_batched",
437            ),
438            mla_cache_assemble_batched_k: gate(
439                probes.mla,
440                gpu,
441                "mla_absorbed",
442                "mla_cache_assemble_batched",
443            ),
444            prefill_attn_mla320_k: gate(
445                probes.mla,
446                gpu,
447                "mla_prefill_attn",
448                "mla_prefill_attn_320",
449            ),
450            grouped_gemm_mla_k: gate(probes.mla, gpu, "grouped_gemm_mla", "grouped_gemm_mla"),
451            mla_q_final_assemble_k: gate(
452                probes.mla,
453                gpu,
454                "mla_absorbed",
455                "mla_q_final_assemble_batched",
456            ),
457            mla_fused_prefill_k: gate(probes.mla, gpu, "mla_fused_prefill", "mla_fused_prefill"),
458            gemm_splitk_partial_k: super::super::try_kernel(
459                gpu,
460                "gemm_splitk",
461                "dense_gemm_splitk_partial",
462            ),
463            gemm_splitk_reduce_k: super::super::try_kernel(
464                gpu,
465                "gemm_splitk",
466                "dense_gemm_splitk_reduce",
467            ),
468            dense_gemm_tc_k: super::super::try_kernel(gpu, "gemm_tc", "dense_gemm_tc"),
469            paged_decode_splitk_k: match kv_dtype {
470                KvCacheDtype::Nvfp4 => {
471                    Some(gpu.kernel("paged_decode_nvfp4", "paged_decode_attn_splitk_nvfp4")?)
472                }
473                KvCacheDtype::Turbo3
474                | KvCacheDtype::Turbo4
475                | KvCacheDtype::Turbo8
476                | KvCacheDtype::Bf16KTurbo3V
477                | KvCacheDtype::Bf16KTurbo4V
478                | KvCacheDtype::Bf16KTurbo2V
479                | KvCacheDtype::Fp8KTurbo3V
480                | KvCacheDtype::Fp8KTurbo4V
481                | KvCacheDtype::Fp8KTurbo2V
482                | KvCacheDtype::Turbo4KTurbo3V
483                | KvCacheDtype::Turbo4KTurbo8V
484                | KvCacheDtype::Turbo3KTurbo8V => None,
485                _ => Some(gpu.kernel("paged_decode_fp8", "paged_decode_attn_splitk_fp8")?),
486            },
487            paged_decode_reduce_k: match kv_dtype {
488                KvCacheDtype::Nvfp4 => {
489                    Some(gpu.kernel("paged_decode_nvfp4", "paged_decode_attn_reduce_nvfp4")?)
490                }
491                KvCacheDtype::Turbo3
492                | KvCacheDtype::Turbo4
493                | KvCacheDtype::Turbo8
494                | KvCacheDtype::Bf16KTurbo3V
495                | KvCacheDtype::Bf16KTurbo4V
496                | KvCacheDtype::Bf16KTurbo2V
497                | KvCacheDtype::Fp8KTurbo3V
498                | KvCacheDtype::Fp8KTurbo4V
499                | KvCacheDtype::Fp8KTurbo2V
500                | KvCacheDtype::Turbo4KTurbo3V
501                | KvCacheDtype::Turbo4KTurbo8V
502                | KvCacheDtype::Turbo3KTurbo8V => None,
503                _ => Some(gpu.kernel("paged_decode_fp8", "paged_decode_attn_reduce_fp8")?),
504            },
505            // The Hopper split-K twins (#928). `try_kernel`, not `kernel`: the
506            // sources live only in `kernels/hopper/common`, so on gb10, b200,
507            // strix and metal the lookup returns a zero handle and the dispatch
508            // keeps its existing arm. Resolved unconditionally rather than
509            // behind the `attn_decode_splitk` lever because the FP8 twin is a
510            // drop-in for the gb10 pair whenever split-K runs at all, and
511            // probing on a lever the operator can flip at boot would make the
512            // handle set depend on the environment — which a CUDA graph
513            // capture must not.
514            paged_decode_splitk_hopper_k: present(super::super::try_target_kernel(
515                gpu,
516                "paged_decode_fp8_splitk_hopper",
517                "paged_decode_attn_splitk_fp8_hopper",
518            )),
519            paged_decode_reduce_hopper_k: present(super::super::try_target_kernel(
520                gpu,
521                "paged_decode_fp8_splitk_hopper",
522                "paged_decode_attn_reduce_fp8_hopper",
523            )),
524            paged_decode_splitk_bf16_hopper_k: present(super::super::try_target_kernel(
525                gpu,
526                "paged_decode_bf16_splitk_hopper",
527                "paged_decode_attn_splitk_bf16_hopper",
528            )),
529            paged_decode_reduce_bf16_hopper_k: present(super::super::try_target_kernel(
530                gpu,
531                "paged_decode_bf16_splitk_hopper",
532                "paged_decode_attn_reduce_bf16_hopper",
533            )),
534            residual_add_k: gpu.kernel("residual_add", "bf16_residual_add")?,
535            // Gemma-4 rms-norm uses the absolute formula `out = x * rms * w`.
536            rms_norm_f32_in_k: KernelHandle(0),
537            sigmoid_gate_mul_k: gpu.kernel("residual_add", "sigmoid_gate_mul")?,
538            deinterleave_qg_k: gpu.kernel("ssm_preprocess", "deinterleave_qg")?,
539            w4a16_gemv_qg_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_qg")?,
540            residual_add_rms_norm_k: if crate::ships_vanilla_norm_weights(config) {
541                gpu.kernel("norm", "residual_add_rms_norm_vanilla")?
542            } else {
543                gpu.kernel("norm", "residual_add_rms_norm")?
544            },
545            residual_add_rms_norm_gatef32_k: crate::layers::try_kernel(
546                gpu,
547                "norm",
548                "residual_add_rms_norm_gatef32",
549            ),
550            w4a16_gemv_qg_batch2_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_qg_batch2")?,
551            w4a16_gemv_dual_batch2_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_dual_batch2")?,
552            w4a16_gemv_batch2_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_batch2")?,
553            w4a16_gemv_qg_batch3_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_qg_batch3")?,
554            w4a16_gemv_dual_batch3_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_dual_batch3")?,
555            w4a16_gemv_batch3_k: gpu.kernel("w4a16_gemv", "w4a16_gemv_batch3")?,
556            w4a16_batchm: crate::layers::w4a16_gemv_tiers::W4a16BatchmTiers::resolve(gpu),
557            w4a16_gemm_k: gpu.kernel("w4a16", "w4a16_gemm")?,
558            w4a16_gemm_t_k: crate::layers::tgemm_kernel(gpu),
559            w4a16_gemm_t_k64_k: crate::layers::k64_kernel(gpu)?,
560            w4a16_gemm_t_k64_n64_k: crate::layers::k64_n64_kernel(gpu),
561            w4a16_gemm_t_m128_k: gpu.kernel("w4a16", "w4a16_gemm_t_m128")?,
562            w4a16_gemm_t_m128_bf16_k: super::super::try_kernel(
563                gpu,
564                "w4a16",
565                "w4a16_gemm_t_m128_bf16",
566            ),
567            w4a16_gemm_t_m128_v2_k: super::super::w4a16_v2_kernel(gpu),
568            w4a16_gemm_t_m128_v3_k: super::super::w4a16_v3_kernel(gpu),
569            dense_gemm_k: gpu.kernel("gemm", "dense_gemm_bf16")?,
570            dense_gemm_pipelined_k: super::super::try_kernel(
571                gpu,
572                "gemm",
573                "dense_gemm_bf16_pipelined",
574            ),
575            prefill_attn_k: gpu.kernel("inferspark_prefill", "inferspark_prefill")?,
576            // Name comes from the SSOT helper that also supplies the BR the
577            // launcher builds its grid from — see `ops::wide_prefill_kernel`.
578            // Module and entry share a name for both variants.
579            // Resolved WITH FALLBACK — see `ops::wide_prefill_kernel`. A target
580            // that ships only the scalar HDIM=512 kernel must still get it.
581            prefill_attn_512_k: if probes.wide_head_dim {
582                crate::layers::ops::wide_prefill_kernel(gpu).0
583            } else {
584                spark_runtime::gpu::KernelHandle(0)
585            },
586            // BR=32 is the tensor-core instantiation; BR=16 the scalar reference.
587            prefill_attn_512_is_tc: probes.wide_head_dim
588                && crate::layers::ops::wide_prefill_kernel(gpu).1 == 32,
589            // DeepSeek-V4 sparse-attention compressor + compressed-KV prefill.
590            csa_compress_k: gate(probes.compressed_attn, gpu, "csa_compress", "csa_compress"),
591            prefill_attn_compressed_k: gate(
592                probes.compressed_attn,
593                gpu,
594                "prefill_attn_compressed",
595                "prefill_attn_compressed",
596            ),
597            v4_comp_pool_filled: std::sync::atomic::AtomicU32::new(0),
598            v4_comp_prev_valid: std::sync::atomic::AtomicBool::new(false),
599            v4_decode_started: std::sync::atomic::AtomicBool::new(false),
600            v4_decode_first_pos: std::sync::atomic::AtomicU32::new(0),
601            prefill_attn_paged_512_k: gate(
602                probes.wide_head_dim,
603                gpu,
604                "inferspark_prefill_paged_512",
605                "inferspark_prefill_paged_512",
606            ),
607            prefill_attn_64_k: gpu.kernel("inferspark_prefill", "inferspark_prefill_64")?,
608            prefill_attn_paged_k: gpu.kernel("prefill_paged", "inferspark_prefill_paged")?,
609            prefill_attn_paged_fp8_k: gpu
610                .kernel("prefill_paged_fp8", "inferspark_prefill_paged_fp8")?,
611            prefill_attn_paged_nvfp4_k: gpu
612                .kernel("prefill_paged_nvfp4", "inferspark_prefill_paged_nvfp4")?,
613            prefill_attn_paged_turbo4_k: super::super::try_kernel(
614                gpu,
615                "prefill_paged_turbo4",
616                "inferspark_prefill_paged_turbo4",
617            ),
618            prefill_attn_paged_64_k: gpu.kernel("prefill_paged", "inferspark_prefill_paged_64")?,
619            prefill_attn_paged_fp8_64_k: gpu
620                .kernel("prefill_paged_fp8", "inferspark_prefill_paged_fp8_64")?,
621            prefill_attn_paged_nvfp4_64_k: gpu
622                .kernel("prefill_paged_nvfp4", "inferspark_prefill_paged_nvfp4_64")?,
623            prefill_attn_paged_turbo2_64_k: super::super::try_kernel(
624                gpu,
625                "prefill_paged_turbo2",
626                "inferspark_prefill_paged_turbo2",
627            ),
628            prefill_attn_paged_turbo3_64_k: super::super::try_kernel(
629                gpu,
630                "prefill_paged_turbo3",
631                "inferspark_prefill_paged_turbo3_64",
632            ),
633            prefill_attn_paged_turbo4_64_k: super::super::try_kernel(
634                gpu,
635                "prefill_paged_turbo4",
636                "inferspark_prefill_paged_turbo4_64",
637            ),
638            prefill_attn_paged_turbo8_64_k: super::super::try_kernel(
639                gpu,
640                "prefill_paged_turbo8",
641                "inferspark_prefill_paged_turbo8_64",
642            ),
643            // TurboQuant+ safer-asym Bf16K + Turbo3V BR=64 prefill kernel.
644            // Compiled from inferspark_prefill_paged_bf16k_turbo3v.cu which
645            // forks prefill_paged_compute_asym.cuh (LOAD_K_TILE = bf16,
646            // LOAD_V_TILE = turbo3 3-bit dequant).
647            prefill_attn_paged_bf16k_turbo3v_64_k: super::super::try_kernel(
648                gpu,
649                "prefill_paged_bf16k_turbo3v",
650                "inferspark_prefill_paged_bf16k_turbo3v_64",
651            ),
652            // Bf16K + Turbo4V BR=64 prefill (4-bit V dequant in LOAD_V_TILE).
653            prefill_attn_paged_bf16k_turbo4v_64_k: super::super::try_kernel(
654                gpu,
655                "prefill_paged_bf16k_turbo4v",
656                "inferspark_prefill_paged_bf16k_turbo4v_64",
657            ),
658            // Bf16K + Turbo2V BR=64 prefill (2-bit V dequant in LOAD_V_TILE).
659            prefill_attn_paged_bf16k_turbo2v_64_k: super::super::try_kernel(
660                gpu,
661                "prefill_paged_bf16k_turbo2v",
662                "inferspark_prefill_paged_bf16k_turbo2v_64",
663            ),
664            // Fp8K + TurboNV BR=64 prefill kernels — K loaded as FP8 (per-tensor
665            // `k_scale` dequant in LOAD_K_TILE), V as 3/4/2-bit Lloyd-Max packed.
666            prefill_attn_paged_fp8k_turbo3v_64_k: super::super::try_kernel(
667                gpu,
668                "prefill_paged_fp8k_turbo3v",
669                "inferspark_prefill_paged_fp8k_turbo3v_64",
670            ),
671            prefill_attn_paged_fp8k_turbo4v_64_k: super::super::try_kernel(
672                gpu,
673                "prefill_paged_fp8k_turbo4v",
674                "inferspark_prefill_paged_fp8k_turbo4v_64",
675            ),
676            prefill_attn_paged_fp8k_turbo2v_64_k: super::super::try_kernel(
677                gpu,
678                "prefill_paged_fp8k_turbo2v",
679                "inferspark_prefill_paged_fp8k_turbo2v_64",
680            ),
681            // Both-sides-quantized TurboQuant+ asym BR=64 prefill kernels.
682            // K loaded via turbo* dequant in LOAD_K_TILE, V via the corresponding
683            // turbo* dequant in LOAD_V_TILE — separate (block_stride, data_section)
684            // pairs per side.
685            prefill_attn_paged_turbo4k_turbo3v_64_k: super::super::try_kernel(
686                gpu,
687                "prefill_paged_turbo4k_turbo3v",
688                "inferspark_prefill_paged_turbo4k_turbo3v_64",
689            ),
690            prefill_attn_paged_turbo4k_turbo8v_64_k: super::super::try_kernel(
691                gpu,
692                "prefill_paged_turbo4k_turbo8v",
693                "inferspark_prefill_paged_turbo4k_turbo8v_64",
694            ),
695            prefill_attn_paged_turbo3k_turbo8v_64_k: super::super::try_kernel(
696                gpu,
697                "prefill_paged_turbo3k_turbo8v",
698                "inferspark_prefill_paged_turbo3k_turbo8v_64",
699            ),
700            // ── Q12 Phase 3: batched paged-prefill kernel handles ──
701            prefill_attn_paged_batched_k: super::super::try_kernel(
702                gpu,
703                "inferspark_prefill_paged_batched",
704                "inferspark_prefill_paged_batched",
705            ),
706            prefill_attn_paged_fp8_batched_k: super::super::try_kernel(
707                gpu,
708                "inferspark_prefill_paged_fp8_batched",
709                "inferspark_prefill_paged_fp8_batched",
710            ),
711            prefill_attn_paged_nvfp4_batched_k: super::super::try_kernel(
712                gpu,
713                "inferspark_prefill_paged_nvfp4_batched",
714                "inferspark_prefill_paged_nvfp4_batched",
715            ),
716            prefill_attn_paged_batched_64_k: super::super::try_kernel(
717                gpu,
718                "inferspark_prefill_paged_batched",
719                "inferspark_prefill_paged_batched_64",
720            ),
721            prefill_attn_paged_fp8_batched_64_k: super::super::try_kernel(
722                gpu,
723                "inferspark_prefill_paged_fp8_batched",
724                "inferspark_prefill_paged_fp8_batched_64",
725            ),
726            prefill_attn_paged_nvfp4_batched_64_k: super::super::try_kernel(
727                gpu,
728                "inferspark_prefill_paged_nvfp4_batched",
729                "inferspark_prefill_paged_nvfp4_batched_64",
730            ),
731            deinterleave_qg_split_k: gpu.kernel("ssm_preprocess", "deinterleave_qg_split")?,
732            deinterleave_qg_split_qnorm_k: gpu
733                .kernel("ssm_preprocess", "deinterleave_qg_split_qnorm")?,
734            deinterleave_qg_split_qnorm_mrope_k: super::super::try_kernel(
735                gpu,
736                "ssm_preprocess",
737                "deinterleave_qg_split_qnorm_mrope",
738            ),
739            sigmoid_gate_mul_batched_k: gpu.kernel("residual_add", "sigmoid_gate_mul_batched")?,
740            q_fp8: None,
741            k_fp8: None,
742            v_fp8: None,
743            o_fp8: None,
744            fp8_gemm_k: gpu.kernel("w4a16", "fp8_gemm_t")?,
745            bf16_to_fp8_k: gpu.kernel("w4a16", "bf16_to_fp8")?,
746            fp8_fp8_gemm_k: gpu.kernel("w4a16", "fp8_fp8_gemm_t")?,
747            fp8_gemm_t_m128_k: gpu.kernel("w4a16", "fp8_gemm_t_m128")?,
748            fp8_fp8_gemm_t_m128_k: gpu.kernel("w4a16", "fp8_fp8_gemm_t_m128")?,
749            w4a4_gemm_k: crate::layers::try_kernel(gpu, "w4a4", "w4a4_gemm_mfast"),
750            quantize_nvfp4_k: crate::layers::try_kernel(
751                gpu,
752                "quantize_nvfp4",
753                "quantize_bf16_to_nvfp4",
754            ),
755            fp8_calibration: if fp8_calibration_tokens > 0
756                && crate::layers::fp8_calibration::dtype_runs_online_fp8_kv_calibration(kv_dtype)
757            {
758                Some(Fp8KvCalibration::new(
759                    attn_layer_idx,
760                    fp8_calibration_tokens,
761                    config.fp8_kv_headroom,
762                    gpu,
763                )?)
764            } else {
765                None
766            },
767        })
768    }
769}