spark_model/layers/mod.rs
1// SPDX-License-Identifier: AGPL-3.0-only
2
3pub mod deepseek_v4_mtp;
4pub mod dense_ffn;
5pub mod dflash_head;
6pub mod ep_dispatch;
7pub mod fp8_calibration;
8mod gemv_tier;
9/// GLM-5.3-Flash KDA integrated layer (Slice 6 -- one layer, no scheduler/cache wiring).
10pub mod glm5next_dsa;
11/// GLM-5.3-Flash DSA + kpool indexer CPU reference (Slice 8 design artifact).
12pub mod glm5next_dsa_ref;
13pub mod glm5next_kda;
14/// GLM-5.3-Flash KDA CPU reference (Slice 2 design artifact -- not a production forward path).
15pub mod glm5next_kda_ref;
16/// GLM-5.3-Flash composite decoder layer -- mixer (KDA|DSA) + MLP (dense|MoE) + mHC.
17pub mod glm5next_layer;
18/// GLM-5.3-Flash MLP production surface -- dense FFN + routed NVFP4 MoE (TP + EP sharded).
19pub mod glm5next_mlp;
20pub mod glm5next_mtp_head;
21/// GLM-5.3-Flash 45-layer text-model skeleton (Slice 9 -- topology, wiring, structural binding).
22pub mod glm5next_skeleton;
23pub mod moe;
24pub mod mtp_head;
25pub(crate) mod mtp_meta;
26pub mod mtp_multi;
27pub mod nemotron_mamba2;
28pub mod nemotron_moe;
29pub mod ngram_embed;
30pub mod ops;
31pub mod ple;
32pub mod qsa;
33pub mod qwen3_attention;
34pub mod qwen3_ssm;
35pub mod vision_encoder;
36pub mod w4a16_gemv_tiers;
37
38/// Minimum K at which the deep-K `w4a16_gemm_t_k64` (K_STEP_T=64) beats the
39/// K_STEP_T=32 `w4a16_gemm_t`.
40///
41/// ★ 6144, not 4096. Measured with `w4a16_m17_bench` on the REAL decode shapes at
42/// M=16 against the STREAM-measured 230 GB/s ceiling — `_k64` is the WORST tile
43/// variant at K=5120 and the best only at K>=6144:
44///
45/// ssm_qkvz N=16384 K=5120 _t 281.9us _k64 341.6us _m128 272.4us
46/// attn qkv N=14336 K=5120 _t 273.9us _k64 328.5us _m128 262.8us
47/// ssm_out_proj N=5120 K=6144 _t 237.7us _k64 163.3us _m128 240.7us
48///
49/// The original 4096 threshold (this session) was derived from the ffn/out_proj
50/// shapes and wrongly generalised to K=5120, sending 48 qkvz + 16 fused-qkv
51/// launches per step to the slowest variant. Both variants accumulate K
52/// sequentially, so moving between them is byte-identical.
53///
54/// `ATLAS_NO_W4A16_K64=1` restores the pre-session 8192 threshold.
55pub(crate) fn w4a16_k64_min_k() -> u32 {
56 static MIN_K: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
57 *MIN_K.get_or_init(|| {
58 // Explicit override so an A/B can pin a previous threshold exactly.
59 if let Some(n) = std::env::var("ATLAS_W4A16_K64_MIN_K")
60 .ok()
61 .and_then(|v| v.parse::<u32>().ok())
62 {
63 return n;
64 }
65 if std::env::var("ATLAS_NO_W4A16_K64").ok().as_deref() == Some("1") {
66 8192
67 } else {
68 6144
69 }
70 })
71}
72
73pub use deepseek_v4_mtp::{DeepseekV4MtpHead, DeepseekV4MtpProposerState};
74pub use dense_ffn::{DenseFfnLayer, DenseFfnWeights, FfnActivation};
75pub use dflash_head::{
76 BlockDiffusionDraftHead, DflashLayer, DflashProposerState, DflashQuantization, dflash_ctx_cap,
77};
78pub use glm5next_mtp_head::Glm5NextMtpHead;
79pub use moe::MoeLayer;
80pub use mtp_head::{MtpHead, MtpQuantization, mtp_drafter_prefill_enabled};
81pub use nemotron_mamba2::NemotronMamba2Layer;
82pub use nemotron_moe::NemotronMoeLayer;
83pub use qwen3_attention::Qwen3AttentionLayer;
84pub use qwen3_ssm::Qwen3SsmLayer;
85pub use vision_encoder::{MergerLayer, ViTBlock, VisionEncoder};
86
87use crate::layer::ForwardContext;
88use anyhow::Result;
89use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
90
91/// Try to load an optional kernel, logging at debug level if it's not found.
92/// Returns `KernelHandle(0)` (null) on failure — callers must check before use.
93///
94/// Debug (not warn) because misses are expected when a model doesn't use a
95/// given feature: e.g. Qwen3-Coder-Next (GDN+attention) never calls MLA
96/// kernels, but the layer builder still probes them. Warning on expected
97/// misses drowned out genuine problems in startup logs.
98/// Resolve the `w4a16_gemm_t_m128_v2` handle honoring `ATLAS_W4A16_VARIANT`.
99///
100/// One resolver for the THREE sites that dispatch on this handle (attention
101/// projections, dense-FFN prefill, SSM batched decode), so variant policy and
102/// rollback live in exactly one place. Default (unset) resolves to a ZERO
103/// handle — v1 everywhere — because the 27B port measured SLOWER than v1
104/// (see body). `ATLAS_W4A16_VARIANT=v2` opts in on all three sites at once;
105/// requesting it on a target without the kernel is a HARD startup error
106/// (fail fast, not a silent fallback discovered in a perf regression).
107#[track_caller]
108pub fn w4a16_v2_kernel(gpu: &dyn GpuBackend) -> KernelHandle {
109 let variant = std::env::var("ATLAS_W4A16_VARIANT").ok();
110 // DEFAULT OFF on the qwen3 layer stack: the 27B port of the 8-warp v2
111 // crush kernel is bit-identical to v1 (microtest 100% on 8 shapes) but
112 // MEASURED SLOWER on the 27B FFN shapes — 0.78-0.82x of v1 standalone
113 // (w4a16_bf16_v2_bench, 2026-07-30; v1 58-74 TFLOP/s). The kernel stays
114 // in the PTX set for A/B and for shape regimes where the extra warps
115 // might pay; nothing auto-activates it. `ATLAS_W4A16_VARIANT=v2` opts in
116 // (hard error if the target lacks the kernel).
117 if !matches!(variant.as_deref(), Some("v2") | Some("v3")) {
118 if variant.as_deref() == Some("v1") {
119 tracing::debug!("ATLAS_W4A16_VARIANT=v1: w4a16 m128 v2 suppressed (explicit)");
120 }
121 return KernelHandle(0);
122 }
123 let h = try_kernel(gpu, "w4a16_v2", "w4a16_gemm_t_m128_v2");
124 if h.0 == 0 {
125 panic!(
126 "ATLAS_W4A16_VARIANT={} requested but w4a16_v2::w4a16_gemm_t_m128_v2 is not in this \
127 target's kernel set — refusing to start with a silently-degraded config",
128 variant.unwrap()
129 );
130 }
131 tracing::debug!(
132 handle = h.0,
133 "w4a16_gemm_t_m128_v2 resolution (explicit opt-in)"
134 );
135 h
136}
137
138/// Resolve the W4A16 m128 **v3** GEMM. Opt-in ONLY, same contract as
139/// [`w4a16_v2_kernel`]: `ATLAS_W4A16_VARIANT=v3` selects it, anything else
140/// resolves to a ZERO handle WITHOUT issuing a lookup.
141///
142/// Not issuing the lookup is the point. `prefill_weights` dispatches on
143/// `v == 3 && handle != 0`, so on the default (`v1`) the probe could never be
144/// used — it only ever added a permanently-failing row to the boot audit on
145/// every target that does not ship `w4a16_v3`. Requesting the variant on such a
146/// target is a HARD error, not a silent fallback discovered in a perf report.
147#[track_caller]
148pub fn w4a16_v3_kernel(gpu: &dyn GpuBackend) -> KernelHandle {
149 if std::env::var("ATLAS_W4A16_VARIANT").as_deref() != Ok("v3") {
150 return KernelHandle(0);
151 }
152 let h = try_kernel(gpu, "w4a16_v3", "w4a16_gemm_t_m128_v3");
153 if h.0 == 0 {
154 panic!(
155 "ATLAS_W4A16_VARIANT=v3 requested but w4a16_v3::w4a16_gemm_t_m128_v3 is not in this \
156 target's kernel set — refusing to start with a silently-degraded config"
157 );
158 }
159 h
160}
161
162/// Resolve the N128/M64 tile GEMM, preferring the 3-deep weight-pipeline variant.
163/// **ON by default**; `ATLAS_NO_TGEMM_PIPELINE3` (presence — `=0` is NOT "off")
164/// falls back to the 2-stage parent. Falls back automatically on any target that
165/// does not ship `_p3`.
166///
167/// Same mechanism as [`k64_kernel`]: the parent drains its cp.async group before
168/// the dequant phase, which only a co-resident CTA can cover. This kernel's live
169/// shapes — ssm_qkvz (128 CTAs) and the fused QKV (112) — sit in the exposed
170/// band of the grid.x-vs-efficiency curve. Bit-identical.
171#[track_caller]
172pub fn tgemm_kernel(gpu: &dyn GpuBackend) -> KernelHandle {
173 if std::env::var("ATLAS_NO_TGEMM_PIPELINE3").is_err() {
174 let h = try_kernel(gpu, "w4a16", "w4a16_gemm_t_p3");
175 if h.0 != 0 {
176 return h;
177 }
178 }
179 try_kernel(gpu, "w4a16", "w4a16_gemm_t")
180}
181
182/// Resolve the k64 deep-K tile GEMM, preferring the 3-deep weight-pipeline
183/// variant. **ON by default**; `ATLAS_NO_K64_PIPELINE3` (presence — `=0` is NOT
184/// "off") falls back to the 2-stage parent.
185///
186/// The parent issues one cp.async group then `wait_all`s it before the dequant
187/// phase, so with a small grid there are ZERO outstanding loads across that
188/// phase. The out_proj/o_proj shapes (N=5120, K=6144) launch 40 CTAs on 48 SMs —
189/// exactly 1 CTA/SM — so nothing covers the drain, and they measure ~38% of
190/// achievable while lm_head (1938 CTAs) reaches 83% on the identical loop.
191/// `_p3` keeps step i+2's loads in flight across dequant(i+1). Bit-identical.
192#[track_caller]
193pub fn k64_kernel(gpu: &dyn GpuBackend) -> Result<KernelHandle> {
194 let want_p3 = std::env::var("ATLAS_NO_K64_PIPELINE3").is_err();
195 if want_p3 {
196 let h = try_kernel(gpu, "w4a16", "w4a16_gemm_t_k64_p3");
197 if h.0 != 0 {
198 return Ok(h);
199 }
200 }
201 gpu.kernel("w4a16", "w4a16_gemm_t_k64")
202}
203
204/// Resolve the NARROW-N (N_TILE=64) deep-K twin. `KernelHandle(0)` when the
205/// kernel is absent or the presence kill switch `ATLAS_NO_K64_N64` is set
206/// (`=0` is NOT "off"). Callers must store the handle — `kernel()` is an
207/// init-time lookup, not a per-launch one.
208#[track_caller]
209pub fn k64_n64_kernel(gpu: &dyn GpuBackend) -> KernelHandle {
210 if std::env::var("ATLAS_NO_K64_N64").is_ok() {
211 return KernelHandle(0);
212 }
213 try_kernel(gpu, "w4a16", "w4a16_gemm_t_k64_n64_p3")
214}
215
216/// Wide-tile CTA count below which the N_TILE=64 deep-K twin wins.
217///
218/// `w4a16_gemm_t_k64_p3` owns a 128-wide N tile and a 64-row M tile, so a
219/// launch is `ceil(n/128) * ceil(m/64)` CTAs of 128 threads. At the out_proj /
220/// o_proj shape (N=5120, K=6144 — 64 launches per decode step at n=64) that is
221/// **40 CTAs on a 48-SM device**: 8 SMs idle by construction and the other 40
222/// hold ONE CTA each, i.e. 4 warps against a 48-warp SM budget.
223///
224/// That is NOT a bandwidth problem. A full-working-set replay (rotating over
225/// enough distinct weight tensors that every launch sees a cold 24 MB L2)
226/// decomposes the 167.8 us launch as: memory pipe alone 87.2 us = **95.6% of
227/// the 229.6 GB/s row-strided ceiling** — already saturated — with the
228/// remaining 80.6 us being barrier-serialized dequant that has no co-resident
229/// warp to hide it. Halving the N tile doubles the grid to 80 CTAs, filling all
230/// 48 SMs with 1-2 co-resident CTAs, and moves the SAME bytes: each CTA still
231/// owns a disjoint N slice and the extra A re-reads are L2 hits (A is 786 KB).
232///
233/// Measured (replay, cold L2, bit-identical output at every point):
234/// wide CTAs 20 -> 1.74x 40 -> 1.42x 48 -> 1.33x 64 -> 1.15x
235/// wide CTAs 76 -> 0.77x 80 -> 0.80x 96 -> 0.79x 244 -> 0.95x
236/// Above ~64 CTAs the wide tile already fills the machine and the narrow tile
237/// only pays more epilogue and A traffic, so the gate is a hard `<= 64`.
238const K64_N64_MAX_WIDE_CTAS: u32 = 64;
239
240/// Should the narrow-N deep-K twin serve this shape? See
241/// `K64_N64_MAX_WIDE_CTAS` for the derivation and the measured curve.
242pub fn k64_n64_wins(m: u32, n: u32) -> bool {
243 n.div_ceil(128) * m.div_ceil(64) <= K64_N64_MAX_WIDE_CTAS
244}
245
246/// Optional kernel lookup: `KernelHandle(0)` instead of an error.
247///
248/// `#[track_caller]` so the audit names the DISPATCH SITE — this helper stands
249/// between ~500 call sites and `GpuBackend::kernel`, and without it every
250/// optional lookup in the binary would be reported against this one line.
251///
252/// A zero handle is a SILENT slower path, so a lookup that lands here for a
253/// model that genuinely needs the kernel is a bug. Either gate the call on the
254/// model's config so it is never issued, or declare it in the target's
255/// MODEL.toml `[expected_absent]` with a reason; the boot gate
256/// (`kernel_audit::classify_failures`) fails closed on anything else.
257/// Minimum rows in flight for the grouped-GEMM MoE decode arm. SSOT for the
258/// SSM stack (`qwen3_ssm::trait_decode_multi_seq`) and the attention layers
259/// (`qwen3_attention::…::multi_seq::ffn`), which must agree — they are the
260/// same trade on the same weights.
261///
262/// The arm reads each routed expert ONCE instead of once per token, so it
263/// wins when there are enough tokens to amortise the expert sort/permute
264/// launch overhead, and loses when there are not. Both ends are measured:
265///
266/// | n | verdict | measurement |
267/// |---|---|---|
268/// | 4 | LOSS | 31 vs 56 tok/s on Holo — the fixed per-layer sort/permute dominates at small N |
269/// | >=16 | WIN | SSM-side alone C=32 172.7 -> 216.2 tok/s (+25%); #415's attention-side extension +7.9% at C=32 / +9.7% at C=64 on Qwen3.6-35B-A3B-NVFP4, paired gsm8k n=200 strict 0.960 vs 0.900 baseline, zero regressions |
270///
271/// 16 is the smallest width measured on the winning side. n=5..15 is
272/// UNMEASURED, not a known win — it sits on the losing side of this gate on
273/// purpose, because the one thing we know about the gap is that the loss at
274/// n=4 is large (-45%) and the win at n=16 is smaller (+25%).
275pub fn moe_grouped_decode_min_rows() -> usize {
276 16
277}
278
279/// Kill switch for the grouped-GEMM MoE decode arm. PRESENCE check per the
280/// house convention (`ATLAS_NO_MOE_GROUPED_DECODE=0` is NOT off), read once
281/// per process — this predicate sits in the decode path, and the `env::var`
282/// it replaces ran on every dispatch for MoE models.
283pub fn moe_grouped_decode_enabled() -> bool {
284 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
285 *ON.get_or_init(|| std::env::var_os("ATLAS_NO_MOE_GROUPED_DECODE").is_none())
286}
287
288/// Force the grouped arm BELOW `moe_grouped_decode_min_rows()`. Diagnostic
289/// only — it exists so the n=5..15 gap can be measured without a rebuild, and
290/// it is the same var #415's measurements used, kept working on purpose.
291/// Never a production setting: if forcing wins at a width, move the THRESHOLD.
292pub fn moe_grouped_decode_forced() -> bool {
293 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
294 *ON.get_or_init(|| std::env::var("ATLAS_MOE_GROUPED_DECODE").as_deref() == Ok("1"))
295}
296
297/// Whether the grouped-GEMM MoE decode arm should run for `n` rows —
298/// PURE, so both polarities are testable without touching process env or the
299/// `OnceLock`s below (which latch, and would make the tests order-dependent).
300pub fn moe_grouped_decode_decide(n: usize, enabled: bool, forced: bool) -> bool {
301 enabled && (n >= moe_grouped_decode_min_rows() || forced)
302}
303
304/// Whether the grouped-GEMM MoE decode arm should run for `n` rows.
305pub fn moe_grouped_decode_for(n: usize) -> bool {
306 moe_grouped_decode_decide(n, moe_grouped_decode_enabled(), moe_grouped_decode_forced())
307}
308
309#[cfg(test)]
310#[path = "moe_grouped_decode_tests.rs"]
311mod moe_grouped_decode_tests;
312
313#[track_caller]
314pub fn try_kernel(gpu: &dyn GpuBackend, module: &str, func: &str) -> KernelHandle {
315 match gpu.kernel(module, func) {
316 Ok(h) => h,
317 Err(_) => {
318 tracing::debug!("Optional kernel '{module}::{func}' not loaded");
319 KernelHandle(0)
320 }
321 }
322}
323
324/// FFN component: MoE (expert routing), dense SwiGLU, or None (standalone attention).
325#[allow(clippy::large_enum_variant)]
326pub enum FfnComponent {
327 Moe(MoeLayer),
328 Dense(DenseFfnLayer),
329 /// No FFN — used by Nemotron-H standalone attention layers.
330 None,
331}
332
333impl FfnComponent {
334 pub fn is_none(&self) -> bool {
335 matches!(self, Self::None)
336 }
337
338 /// True for a plain dense (SwiGLU) FFN. Wide-batch verify paths gate their
339 /// `forward_prefill` fast path on this: batching reads dense weights once
340 /// (big win at N=17), but on a 256-expert MoE the grouped-GEMM is a net
341 /// loss at small batch (per-expert M~1 + sort/permute overhead), so MoE
342 /// keeps its per-token loop.
343 pub fn is_dense(&self) -> bool {
344 matches!(self, Self::Dense(_))
345 }
346
347 /// True when this MoE FFN can serve DECODE through the grouped read-once
348 /// GEMM (forward_prefill) instead of the pairwise per-slot loop. The
349 /// is_dense() comment above asserts grouped is "a net loss at small batch"
350 /// on a 256-expert MoE, but that was never measured for decode CONCURRENCY
351 /// (n=4) where the pairwise path re-reads ~14-20 distinct experts as 40
352 /// per-slot CTAs. Native-NVFP4-routed only (forward_prefill's unconditional
353 /// grouped path); dense/none are false.
354 pub fn moe_grouped_decode_ok(&self) -> bool {
355 match self {
356 Self::Moe(m) => m.grouped_decode_ok(),
357 _ => false,
358 }
359 }
360
361 /// ATLAS_FP32_ROUTING active for this FFN (MoE only; false otherwise).
362 pub fn fp32_routing_active(&self) -> bool {
363 match self {
364 Self::Moe(m) => m.fp32_routing_active(),
365 _ => false,
366 }
367 }
368
369 pub fn forward(
370 &self,
371 input: DevicePtr,
372 ctx: &ForwardContext,
373 stream: u64,
374 ) -> Result<DevicePtr> {
375 match self {
376 Self::Moe(m) => m.forward(input, ctx, stream),
377 Self::Dense(d) => d.forward(input, ctx, stream),
378 Self::None => Ok(input),
379 }
380 }
381
382 pub fn forward_k2(&self, input: DevicePtr, ctx: &ForwardContext, stream: u64) -> Result<()> {
383 match self {
384 Self::Moe(m) => m.forward_k2(input, ctx, stream),
385 Self::Dense(d) => d.forward_k2(input, ctx, stream),
386 Self::None => Ok(()),
387 }
388 }
389
390 pub fn forward_k3(&self, input: DevicePtr, ctx: &ForwardContext, stream: u64) -> Result<()> {
391 match self {
392 Self::Moe(m) => m.forward_k3(input, ctx, stream),
393 Self::Dense(d) => d.forward_k3(input, ctx, stream),
394 Self::None => Ok(()),
395 }
396 }
397
398 /// Whether the K=m (m<=8) batched-GEMV verify FFN is available (dense
399 /// only — MoE / missing batch4/batch8 kernel / non-NVFP4 weights →
400 /// false). Lets callers gate branch entry BEFORE computing the pre-FFN
401 /// norm, so there is no half-done fallthrough to `forward_prefill`.
402 pub fn can_forward_km(&self, m: u32) -> bool {
403 matches!(self, Self::Dense(d) if d.can_forward_km(m))
404 }
405
406 /// K=m (m=4..8) verify FFN via batched GEMV (dense only). Returns
407 /// `false` when the path is unavailable (MoE / missing batchm kernel /
408 /// non-NVFP4 weights) so the caller can fall back to `forward_prefill`.
409 pub fn try_forward_km(
410 &self,
411 input: DevicePtr,
412 m: u32,
413 ctx: &ForwardContext,
414 stream: u64,
415 ) -> Result<bool> {
416 match self {
417 Self::Dense(d) if d.can_forward_km(m) => {
418 d.forward_km(input, m, ctx, stream)?;
419 Ok(true)
420 }
421 _ => Ok(false),
422 }
423 }
424
425 pub fn forward_prefill(
426 &self,
427 input: DevicePtr,
428 num_tokens: usize,
429 ctx: &ForwardContext,
430 stream: u64,
431 ) -> Result<()> {
432 match self {
433 Self::Moe(m) => m.forward_prefill(input, num_tokens, ctx, stream),
434 Self::Dense(d) => d.forward_prefill(input, num_tokens, ctx, stream),
435 Self::None => {
436 let _ = (input, num_tokens);
437 Ok(())
438 }
439 }
440 }
441
442 pub fn forward_batched(
443 &self,
444 input: DevicePtr,
445 num_tokens: usize,
446 ctx: &ForwardContext,
447 stream: u64,
448 ) -> Result<()> {
449 match self {
450 Self::Moe(m) => m.forward_batched(input, num_tokens, ctx, stream),
451 Self::Dense(d) => d.forward_batched(input, num_tokens, ctx, stream),
452 Self::None => {
453 let _ = (input, num_tokens);
454 Ok(())
455 }
456 }
457 }
458
459 pub fn forward_token_major_decode(
460 &self,
461 input: DevicePtr,
462 num_tokens: usize,
463 ctx: &ForwardContext,
464 stream: u64,
465 ) -> Result<()> {
466 match self {
467 Self::Moe(m) => m.forward_token_major_decode(input, num_tokens, ctx, stream),
468 Self::Dense(d) => d.forward_batched(input, num_tokens, ctx, stream),
469 Self::None => {
470 let _ = (input, num_tokens);
471 Ok(())
472 }
473 }
474 }
475
476 pub fn forward_atomic_c4_decode(
477 &self,
478 input: DevicePtr,
479 num_tokens: usize,
480 ctx: &ForwardContext,
481 stream: u64,
482 ) -> Result<()> {
483 match self {
484 Self::Moe(m) => m.forward_atomic_c4_decode(input, num_tokens, ctx, stream),
485 Self::Dense(d) => d.forward_batched(input, num_tokens, ctx, stream),
486 Self::None => {
487 let _ = (input, num_tokens);
488 Ok(())
489 }
490 }
491 }
492}
493
494pub(crate) use gemv_tier::batch8_kernel;