From 57f585f630478f56fdde8f6d7046542d227dd797 Mon Sep 17 00:00:00 2001 From: HUSRCF Date: Mon, 5 Oct 2026 19:30:12 +0800 Subject: [PATCH 1/2] perf(gfx1100): retain FA2 split-KV verifier Replace the long-context Q8 DFlash R4/R8 attention step with an opt-in FA2 split-KV S8 backend on exact gfx1100. Keep the live-context crossover fail-closed and retain the same fixed-grid route through HipGraph and Redline/PM4 with typed pointer effects, stable scratch, and product-identical lm-head handling. Current warpfront/beta@d5305333d W7900 fresh-process E2E: 42.6 to 48.3 tok/s (+13.38%) at unchanged tau=1.97 and 67 cycles. HipGraph off/on produced byte-identical transcripts. Redline four-arm parity passed 12/12 windows; PM4 was 6.00% faster than HipGraph in the five-window verify-body smoke. Validation: release product build, test_kernels 17/17, rdna-compute 438/438, hipfire-arch-qwen35 247/247, hipfire-generate 54/54, crate maps, env/lifecycle inventory, rustfmt, and diff checks. Local no-gpu-ci reached the unchanged beta Python select.py shadowing failure after its Rust/main checks; GitHub beta-runner baseline failures are disclosed in the PR. --- crates/hipfire-arch-qwen35/map.md | 14 +- crates/hipfire-arch-qwen35/src/dflash_spec.rs | 40 +- .../src/dflash_verify_pm4.rs | 36 ++ .../hipfire-arch-qwen35/src/qwen35/prefill.rs | 311 +++++++++-- crates/hipfire-arch-qwen35/src/speculative.rs | 118 +++- crates/hipfire-generate/map.md | 4 +- crates/hipfire-generate/src/redline.rs | 124 +++-- crates/rdna-compute/map.md | 10 +- crates/rdna-compute/src/attention.rs | 524 ++++++++++++++++-- crates/rdna-compute/src/mq_f16_producers.rs | 29 +- crates/rdna-compute/src/replay.rs | 413 ++++++++++++-- docs/env-vars.md | 2 + ...-gfx1100-fa2-splitkv-verifier-graph-pm4.md | 128 +++++ kernels/src/attention_q8_0_fa2_gqa.gfx11.hip | 153 ++++- ...ate_up_mq4g256v2_wmma_gfx1100_ldsstage.hip | 4 +- 15 files changed, 1633 insertions(+), 277 deletions(-) create mode 100644 docs/perf-checkpoints/2026-10-05-gfx1100-fa2-splitkv-verifier-graph-pm4.md diff --git a/crates/hipfire-arch-qwen35/map.md b/crates/hipfire-arch-qwen35/map.md index 3636f19683..22c858bc62 100644 --- a/crates/hipfire-arch-qwen35/map.md +++ b/crates/hipfire-arch-qwen35/map.md @@ -28,8 +28,8 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/carrier.rs`](src/carrier.rs) | 854 | 4 | 7 | | [`src/checkpoint.rs`](src/checkpoint.rs) | 108 | 3 | 0 | | [`src/dflash_slot.rs`](src/dflash_slot.rs) | 783 | 12 | 0 | -| [`src/dflash_spec.rs`](src/dflash_spec.rs) | 2,048 | 19 | 16 | -| [`src/dflash_verify_pm4.rs`](src/dflash_verify_pm4.rs) | 744 | 36 | 9 | +| [`src/dflash_spec.rs`](src/dflash_spec.rs) | 2,080 | 19 | 17 | +| [`src/dflash_verify_pm4.rs`](src/dflash_verify_pm4.rs) | 780 | 37 | 10 | | [`src/forward_slots.rs`](src/forward_slots.rs) | 4,346 | 24 | 6 | | [`src/grammar_config.rs`](src/grammar_config.rs) | 143 | 2 | 4 | | [`src/layer_driver.rs`](src/layer_driver.rs) | 646 | 0 | 2 | @@ -47,13 +47,13 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/qwen35/forward.rs`](src/qwen35/forward.rs) | 7,697 | 33 | 12 | | [`src/qwen35/load.rs`](src/qwen35/load.rs) | 7,635 | 16 | 14 | | [`src/qwen35/oracle.rs`](src/qwen35/oracle.rs) | 1,243 | 20 | 0 | -| [`src/qwen35/prefill.rs`](src/qwen35/prefill.rs) | 17,159 | 18 | 70 | +| [`src/qwen35/prefill.rs`](src/qwen35/prefill.rs) | 17,348 | 18 | 71 | | [`src/qwen35/weights.rs`](src/qwen35/weights.rs) | 2,991 | 43 | 11 | | [`src/qwen35.rs`](src/qwen35.rs) | 69 | 8 | 0 | | [`src/serve_engine.rs`](src/serve_engine.rs) | 7,434 | 12 | 33 | | [`src/spec_emit.rs`](src/spec_emit.rs) | 1,042 | 5 | 14 | | [`src/spec_impl.rs`](src/spec_impl.rs) | 643 | 1 | 0 | -| [`src/speculative.rs`](src/speculative.rs) | 9,932 | 87 | 14 | +| [`src/speculative.rs`](src/speculative.rs) | 10,002 | 88 | 14 | ### Public API surface @@ -63,7 +63,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside - [`src/checkpoint.rs`](src/checkpoint.rs): `hipfire_runtime`, `capture_checkpoint`, `restore_private` - [`src/dflash_slot.rs`](src/dflash_slot.rs): `DflashShared`, `hidden_capture`, `free_gpu`, `DflashSlotState`, `DflashSlotDraft`, `n_verify`, `load_dflash_shared`, `new_dflash_slot_state`, `reset_dflash_slot`, `scatter_staging_rows_to_interleaved`, `dflash_slot_draft_step`, `dflash_slot_verify_accept` - [`src/dflash_spec.rs`](src/dflash_spec.rs): `DdtreeState`, `DflashState`, `free_gpu`, `DEFAULT_DFLASH_CTX_CAP`, `load_dflash_state`, `DflashSpeculator`, `new`, `verify_pm4`, `verify_pm4_mut`, `build_dflash_speculator`, `admit_dflash_verify_pm4`, `DenseTpDflashRankState`, +7 more -- [`src/dflash_verify_pm4.rs`](src/dflash_verify_pm4.rs): `DFLASH_VERIFY_PM4_BLOCK`, `DflashVerifyPm4Phase`, `label`, `reason`, `is_live`, `DflashVerifyPm4Counters`, `DflashVerifyBinding`, `new`, `same_route`, `fingerprint_u64`, `DflashVerifyWindow`, `DflashVerifyRoute`, +24 more +- [`src/dflash_verify_pm4.rs`](src/dflash_verify_pm4.rs): `DFLASH_VERIFY_PM4_BLOCK`, `DflashVerifyPm4Phase`, `label`, `reason`, `is_live`, `DflashVerifyPm4Counters`, `DflashVerifyBinding`, `new`, `same_route`, `fingerprint_u64`, `DflashVerifyWindow`, `DflashVerifyRoute`, +25 more - [`src/forward_slots.rs`](src/forward_slots.rs): `DESC_BYTES`, `SlotKvTier`, `q8`, `is_q8`, `free_gpu`, `SlotDescStaging`, `new`, `SpecVerifyCapture`, `SpecHiddenCapture`, `is_verify_slot`, `forward_batch_slots`, `DECODE_GRAPH_CTX_BUCKET`, +12 more - [`src/grammar_config.rs`](src/grammar_config.rs): `resolve_qwen35_grammar_config`, `resolve_grammar_config` - [`src/layer_driver.rs`](src/layer_driver.rs): — @@ -87,7 +87,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside - [`src/serve_engine.rs`](src/serve_engine.rs): `EngineConfig`, `SlotEngine`, `submit`, `cancel_waiting`, `close`, `reset`, `stats`, `spawn`, `shutdown_engine`, `MTP_RETIRE_WINDOW`, `MTP_RETIRE_MIN_ADVANCE`, `MTP_RETIRE_WINDOWS` - [`src/spec_emit.rs`](src/spec_emit.rs): `Qwen35Emit`, `from_ctx`, `from_ctx_template_think_close`, `decoded_eot`, `visible_text` - [`src/spec_impl.rs`](src/spec_impl.rs): `Qwen35SpecScratch` -- [`src/speculative.rs`](src/speculative.rs): `SeedOracleStats`, `read_seed_oracle_stats`, `reset_seed_oracle_stats`, `record_ddtree_meta_nodes`, `DdtreeMetaStats`, `read_ddtree_meta_stats`, `reset_ddtree_meta_stats`, `KvMode`, `ModelSlotConfig`, `ModelSlot`, `from_bundle`, `into_bundle`, +75 more +- [`src/speculative.rs`](src/speculative.rs): `dflash_finish_retained_lm_head_argmax`, `SeedOracleStats`, `read_seed_oracle_stats`, `reset_seed_oracle_stats`, `record_ddtree_meta_nodes`, `DdtreeMetaStats`, `read_ddtree_meta_stats`, `reset_ddtree_meta_stats`, `KvMode`, `ModelSlotConfig`, `ModelSlot`, `from_bundle`, +76 more ### Dependencies (from `Cargo.toml`) @@ -102,6 +102,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 31 modules · 87,250 lines · 555 public items · 280 tests · 18 examples +- 31 modules · 87,577 lines · 557 public items · 283 tests · 18 examples diff --git a/crates/hipfire-arch-qwen35/src/dflash_spec.rs b/crates/hipfire-arch-qwen35/src/dflash_spec.rs index 6526dd6ed9..cdec47c388 100644 --- a/crates/hipfire-arch-qwen35/src/dflash_spec.rs +++ b/crates/hipfire-arch-qwen35/src/dflash_spec.rs @@ -466,12 +466,18 @@ pub fn load_dflash_state( .as_ref() .map(|p| p.max_batch) .unwrap_or(0); + let gfx1100_split_verify = gpu.arch == "gfx1100" + && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false) + && target_config.n_heads == 24 + && target_config.n_kv_heads == 4 + && target_config.head_dim == 256; // The draft's selector/dynamic-conv shape is deliberately NOT a gate: the // draft forward is outside the tape, so DFlash2 and legacy DFlash yield an // identical target verify body. let verify_pm4 = match admit_dflash_verify_pm4( env_opt_in, &gpu.arch, + gfx1100_split_verify, single_gpu, target_config.num_experts, kv_is_q8, @@ -496,7 +502,14 @@ pub fn load_dflash_state( " DFlash verify PM4: armed (B={}, exact {})", DFLASH_VERIFY_PM4_BLOCK, gpu.arch ); - DflashVerifyPm4::armed() + if gfx1100_split_verify { + // Keep the faster incumbent below the split verifier's 4K + // crossover, and do not capture a short-context tape that + // would silently omit the new partial + merge launches. + DflashVerifyPm4::armed_after_context(4096) + } else { + DflashVerifyPm4::armed() + } } Err(reason) => { eprintln!(" DFlash verify PM4: disabled ({reason})"); @@ -1323,6 +1336,7 @@ pub fn build_dflash_speculator( pub fn admit_dflash_verify_pm4( env_opt_in: bool, arch: &str, + gfx1100_split_verify: bool, single_gpu: bool, num_experts: usize, kv_is_q8: bool, @@ -1342,8 +1356,10 @@ pub fn admit_dflash_verify_pm4( if !env_opt_in { return Err("HIPFIRE_DFLASH_VERIFY_PM4 is not set to 1".into()); } - if arch != "gfx1201" { - return Err(format!("arch is {arch}, not exact gfx1201")); + if arch != "gfx1201" && !(arch == "gfx1100" && gfx1100_split_verify) { + return Err(format!( + "arch is {arch}, not exact gfx1201 or admitted gfx1100 split verifier" + )); } if !single_gpu { return Err("multi-GPU load is not admitted".into()); @@ -1795,6 +1811,7 @@ mod admit_dflash_verify_pm4_tests { struct Args { env_opt_in: bool, arch: &'static str, + gfx1100_split_verify: bool, single_gpu: bool, num_experts: usize, kv_is_q8: bool, @@ -1817,6 +1834,7 @@ mod admit_dflash_verify_pm4_tests { Self { env_opt_in: true, arch: "gfx1201", + gfx1100_split_verify: false, single_gpu: true, num_experts: 0, kv_is_q8: true, @@ -1840,6 +1858,7 @@ mod admit_dflash_verify_pm4_tests { admit_dflash_verify_pm4( a.env_opt_in, a.arch, + a.gfx1100_split_verify, a.single_gpu, a.num_experts, a.kv_is_q8, @@ -1880,7 +1899,20 @@ mod admit_dflash_verify_pm4_tests { ..Args::default() }) .unwrap_err(); - assert_eq!(err, "arch is gfx1100, not exact gfx1201"); + assert_eq!( + err, + "arch is gfx1100, not exact gfx1201 or admitted gfx1100 split verifier" + ); + } + + #[test] + fn admits_gfx1100_only_with_split_verifier() { + assert!(admit(Args { + arch: "gfx1100", + gfx1100_split_verify: true, + ..Args::default() + }) + .is_ok()); } #[test] diff --git a/crates/hipfire-arch-qwen35/src/dflash_verify_pm4.rs b/crates/hipfire-arch-qwen35/src/dflash_verify_pm4.rs index e4545810e0..bbe57458e1 100644 --- a/crates/hipfire-arch-qwen35/src/dflash_verify_pm4.rs +++ b/crates/hipfire-arch-qwen35/src/dflash_verify_pm4.rs @@ -224,6 +224,11 @@ pub enum DflashVerifyRoute { pub struct DflashVerifyPm4 { phase: DflashVerifyPm4Phase, controller: Option, + /// Do not build or submit the retained tape below this logical context. + /// The gfx1100 split-KV route wins only after 4K; deferring capture keeps + /// the short-context shipping kernel in place and guarantees that the + /// recorded body actually contains the split partial + merge pair. + min_context: usize, binding: Option, identity: Option, /// First calibration recording and the position it was taken at. @@ -239,6 +244,7 @@ impl DflashVerifyPm4 { reason: reason.into(), }, controller: None, + min_context: 0, binding: None, identity: None, calibration: None, @@ -249,9 +255,17 @@ impl DflashVerifyPm4 { /// Admitted route with a dedicated manual-PM4 controller. Allocates no GPU /// resource until the first capture. pub fn armed() -> Self { + Self::armed_after_context(0) + } + + /// Arm a route that remains on ordinary HIP until `start_pos + batch` + /// exceeds `min_context`. Once prepared, later short-context requests also + /// stay on HIP while the long-context tape remains ready for reuse. + pub fn armed_after_context(min_context: usize) -> Self { Self { phase: DflashVerifyPm4Phase::Armed, controller: Some(ReplayController::new_manual_pm4()), + min_context, binding: None, identity: None, calibration: None, @@ -301,6 +315,10 @@ impl DflashVerifyPm4 { self.note_hip_window(window); return DflashVerifyRoute::HipAuto; } + if window.position.saturating_add(window.batch) <= self.min_context { + self.note_hip_window(window); + return DflashVerifyRoute::HipAuto; + } if !window.eligible_shape() { self.note_hip_window(window); return DflashVerifyRoute::HipAuto; @@ -543,6 +561,7 @@ impl DflashVerifyPm4 { serde_json::json!({ "phase": self.phase.label(), "reason": self.phase.reason(), + "min_context": self.min_context, "binding": self.binding, "kv_mode": "q8", "dn_state_quant": "q8", @@ -660,6 +679,23 @@ mod tests { assert_eq!(route.counters().last_replay_position, Some(32)); } + #[test] + fn minimum_context_defers_capture_without_changing_phase() { + let bound = binding(1, 8192); + let mut route = DflashVerifyPm4::armed_after_context(4096); + assert_eq!( + route.plan_route(&window(&bound, 16, 4080)), + DflashVerifyRoute::HipAuto + ); + assert_eq!(*route.phase(), DflashVerifyPm4Phase::Armed); + assert_eq!(route.counters().full_hip, 1); + assert_eq!( + route.plan_route(&window(&bound, 16, 4081)), + DflashVerifyRoute::PrimeDirect + ); + assert_eq!(route.counters().prime_windows, 1); + } + #[test] fn position_beyond_prepared_max_rearms_instead_of_replaying() { let bound = binding(1, 64); diff --git a/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs b/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs index e7f89b4ee7..2307d47431 100644 --- a/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs +++ b/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs @@ -329,7 +329,9 @@ fn packed_mq4_down_admitted( down: (DType, usize, usize), hidden_dim: usize, ) -> bool { - packed_admitted && separate && residual + packed_admitted + && separate + && residual && down.0 == DType::MQ4G256 && (down.1, down.2) == (5120, 17408) && hidden_dim == down.2 @@ -350,18 +352,32 @@ fn try_packed_mq4_down( !pbs.lean && gpu.packed_mq4_admitted(down.m, down.k, n), h_source == FfnGateOutput::Separate, matches!(epilogue, BatchEpilogue::Residual), - (down.gpu_dtype, down.m, down.k), hidden_dim, + (down.gpu_dtype, down.m, down.k), + hidden_dim, ) { return Ok(false); } gpu.flush_residual_fold()?; fused_silu_mul_rotate_mq_batched_for( - gpu, down, &pbs.gate_ffn_batch, &pbs.up_batch, - &pbs.ffn_hidden_batch, hidden_dim, n, + gpu, + down, + &pbs.gate_ffn_batch, + &pbs.up_batch, + &pbs.ffn_hidden_batch, + hidden_dim, + n, + )?; + run_residual_gemm_key( + gpu, + hipfire_dispatch::types::KernelKey::GemmMq4PackedResidual, + &down.buf, + down.gpu_dtype, + &pbs.ffn_hidden_batch, + &pbs.x_batch, + down.m, + down.k, + n, )?; - run_residual_gemm_key(gpu, hipfire_dispatch::types::KernelKey::GemmMq4PackedResidual, - &down.buf, down.gpu_dtype, &pbs.ffn_hidden_batch, &pbs.x_batch, - down.m, down.k, n)?; Ok(true) } @@ -405,22 +421,34 @@ fn try_packed_mq4_ffn_gate_up( } gpu.flush_residual_fold()?; let separate_awq_inputs = gate.awq_scale.is_some() || up.awq_scale.is_some(); - for (index, (weight, output)) in [ - (gate, &pbs.gate_ffn_batch), - (up, &pbs.up_batch), - ].into_iter().enumerate() { + for (index, (weight, output)) in [(gate, &pbs.gate_ffn_batch), (up, &pbs.up_batch)] + .into_iter() + .enumerate() + { // Sidecars are per-weight and may differ (or exist on only one side). // Reuse the rotated input only when neither projection carries AWQ. if index == 0 || separate_awq_inputs { fused_rmsnorm_rotate_mq_batched_for( - gpu, &pbs.x_batch, norm, weight, &pbs.x_rot_batch, - weight.k, config.norm_eps, n, + gpu, + &pbs.x_batch, + norm, + weight, + &pbs.x_rot_batch, + weight.k, + config.norm_eps, + n, )?; } run_plain_gemm_key( - gpu, hipfire_dispatch::types::KernelKey::GemmMq4Packed, - &weight.buf, weight.gpu_dtype, &pbs.x_rot_batch, output, - weight.m, weight.k, n, + gpu, + hipfire_dispatch::types::KernelKey::GemmMq4Packed, + &weight.buf, + weight.gpu_dtype, + &pbs.x_rot_batch, + output, + weight.m, + weight.k, + n, )?; } Ok(true) @@ -1507,7 +1535,9 @@ fn dense_layers_have_projection_dtype(weights: &Qwen35Weights, dtype: DType) -> fn native_mq4_widened_requested() -> bool { hipfire_config::developer_var("HIPFIRE_GFX1100_MQ4_WIDE_PREFILL") - .ok().as_deref() == Some("1") + .ok() + .as_deref() + == Some("1") } /// Separately admitted native MQ4 route: no V2 producer or lean PBS contract. @@ -1520,7 +1550,9 @@ fn native_mq4_widened_route(gpu: &Gpu, weights: &Qwen35Weights) -> bool { } fn native_mq4_widened_flags_admitted( - arch: &str, flags: &rdna_compute::FeatureFlags, screen: bool, + arch: &str, + flags: &rdna_compute::FeatureFlags, + screen: bool, ) -> bool { arch == "gfx1100" && !flags.fp16_disabled @@ -1889,9 +1921,17 @@ pub fn minimum_prefill_reservation_bytes(config: &Qwen35Config, arch: &str) -> O // explicit experimental MQ4 opt-ins, conservatively reserve their prelude // too; the default V2/other-format reservation remains unchanged. let packed = hipfire_config::developer_var("HIPFIRE_GFX1100_PACKED_MQ4_PREFILL") - .ok().as_deref() == Some("1"); + .ok() + .as_deref() + == Some("1"); if arch == "gfx1100" && (native_mq4_widened_requested() || packed) { - base.checked_add(native_mq4_projection_deficit(config, WIDENED_COMMIT_ROWS, 0, 0, 0)?) + base.checked_add(native_mq4_projection_deficit( + config, + WIDENED_COMMIT_ROWS, + 0, + 0, + 0, + )?) } else { Some(base) } @@ -1907,11 +1947,18 @@ fn native_mq4_projection_deficit( live_f16: usize, live_partials: usize, ) -> Option { - let mmq = config.hidden_dim.checked_add(127)?.checked_div(128)? - .checked_mul(144)?.checked_mul(rows)?; + let mmq = config + .hidden_dim + .checked_add(127)? + .checked_div(128)? + .checked_mul(144)? + .checked_mul(rows)?; let tail = rows.min(127); let f16 = tail.checked_mul(config.hidden_dim)?.checked_mul(2)?; - let partials = tail.checked_mul(config.dim)?.checked_mul(4)?.checked_mul(4)?; + let partials = tail + .checked_mul(config.dim)? + .checked_mul(4)? + .checked_mul(4)?; mmq.saturating_sub(live_mmq) .checked_add(f16.saturating_sub(live_f16))? .checked_add(partials.saturating_sub(live_partials)) @@ -2082,11 +2129,13 @@ fn memory_admitted_rung( .unwrap_or(usize::MAX); let projection_deficit = if native_mq4 { native_mq4_projection_deficit( - config, rung, + config, + rung, gpu.scratch.q8_1_mmq_x_scratch_bytes, gpu.scratch.fp16_x_scratch_bytes, gpu.scratch.ksplit_det_partials_bytes, - ).unwrap_or(usize::MAX) + ) + .unwrap_or(usize::MAX) } else if fp8_projection { let x_need = rung.checked_mul(config.hidden_dim).unwrap_or(usize::MAX); let sums_need = rung @@ -2172,7 +2221,13 @@ fn ordinary_prefill_chunk_limit_with_cache( } let lean = lean_pbs_requested() && lean_pbs_route(gpu, weights, config, perf); let mut admitted = memory_admitted_rung( - gpu, config, kv_cache, perf, lean, native_mq4_widened_route(gpu, weights), cached, + gpu, + config, + kv_cache, + perf, + lean, + native_mq4_widened_route(gpu, weights), + cached, )?; if let Some(p) = pbs { admitted = admitted.min(p.max_batch); @@ -5479,19 +5534,23 @@ pub(crate) fn batch_chunk_upload_positions( /// S3-f16-projection-inputs: exact-route gate for the FP16 projection-input /// fast path (all four `batch_chunk_*` projection hooks below). /// -/// All predicates are cheap field reads — no env/lock/JIT in the cycle (the -/// kill switch resolves once at `FeatureFlags` init). Every failed predicate -/// runs the pre-change F32-producer + `convert_f32_to_f16` path byte-for-byte. +/// All ordinary-path predicates are cheap field reads. Capture/recording is +/// admitted only together with the exact-gfx1100 split-verifier experiment: +/// that route owns fixed, preallocated F16 sidecars and records every producer +/// and consumer through the blob funnel. Every failed predicate runs the +/// pre-change F32-producer + `convert_f32_to_f16` path byte-for-byte. #[inline] fn mq_f16_projection_fast_route(gpu: &Gpu, fusion: DflashFusionCtx, n: usize, dim: usize) -> bool { + let recording = gpu.graphs.capture_mode || gpu.replay.is_recording(); + let recording_supported = + !recording || hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false); matches!(fusion, DflashFusionCtx::ChainVerify) && gpu.arch_caps.is_gfx1100() && !gpu.flags.mq_f16_projection_off && n >= 1 && n <= 16 && dim % 256 == 0 - && !gpu.graphs.capture_mode - && !gpu.replay.is_recording() + && recording_supported } #[allow(clippy::too_many_arguments)] @@ -7232,7 +7291,13 @@ fn batch_chunk_delta_net_ffn_gate_up( ) -> HipResult { let _ = fusion; if try_packed_mq4_ffn_gate_up( - gpu, &layer.w_gate, &layer.w_up, &layer.ffn_norm, config, pbs, n, + gpu, + &layer.w_gate, + &layer.w_up, + &layer.ffn_norm, + config, + pbs, + n, )? { return Ok(FfnGateOutput::Separate); } @@ -8924,6 +8989,7 @@ fn q8_multirow_attn_admitted( is_independent: bool, capture_mode: bool, replay_recording: bool, + capture_safe_fixed_grid: bool, ) -> bool { matches!(arch, "gfx1100" | "gfx1151" | "gfx1201") && quant_q8 @@ -8932,8 +8998,7 @@ fn q8_multirow_attn_admitted( && min_ctx.is_some_and(|threshold| logical_ctx > threshold) && !is_tree && !is_independent - && !capture_mode - && !replay_recording + && ((!capture_mode && !replay_recording) || capture_safe_fixed_grid) } #[allow(clippy::too_many_arguments)] @@ -9203,7 +9268,7 @@ fn batch_chunk_fa_attend( } _ => {} } - if gpu.attention_flash_q8_0_rows_masked( + if gpu.attention_flash_q8_0_rows_masked_logical( &pbs.fa_q_batch, &kv_cache.k_gpu[layer_idx], &kv_cache.v_gpu[layer_idx], @@ -9213,6 +9278,7 @@ fn batch_chunk_fa_attend( config.n_kv_heads, config.head_dim, max_ctx_len, + start_pos + n, n, &s.flash_partials, )? { @@ -9562,7 +9628,13 @@ fn batch_chunk_full_attn_ffn_gate_up( ) -> HipResult { let _ = fusion; if try_packed_mq4_ffn_gate_up( - gpu, &layer.w_gate, &layer.w_up, &layer.ffn_norm, config, pbs, n, + gpu, + &layer.w_gate, + &layer.w_up, + &layer.ffn_norm, + config, + pbs, + n, )? { return Ok(FfnGateOutput::Separate); } @@ -12371,6 +12443,8 @@ fn forward_prefill_chunk_pair( fa_pertoken_min_ctx(gpu.arch_caps.arch()), gpu.graphs.capture_mode, gpu.replay.is_recording(), + gpu.arch_caps.is_gfx1100() + && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false), ); let multirow_admitted = |max_ctx: usize| { q8_multirow_attn_admitted( @@ -12384,6 +12458,7 @@ fn forward_prefill_chunk_pair( false, multirow_common.4, multirow_common.5, + multirow_common.6, ) }; let multirow_c = multirow_admitted(max_ctx_c); @@ -13054,8 +13129,10 @@ pub(crate) fn forward_batch_chunk_impl( // and re-scans the whole KV once per row, so a small verify block over a // long context pays the scan n times (202 vs 103 ms at 33k). The layer's // GEMMs stay batched either way — only the attend step switches to the - // multi-row tile. Its tile grid is sized from the live logical context on - // the host, so a captured replay would keep the first cycle's tile count. + // multi-row tile. The incumbent tile grid is sized from the live logical + // context and therefore stays out of capture. The gfx1100 split-KV route + // uses fixed S=8 geometry and is admitted only behind its explicit opt-in; + // the launcher independently rejects every other captured shape. let fa_attn_multirow = q8_multirow_attn_admitted( gpu.arch_caps.arch(), kv_cache.quant_q8, @@ -13067,6 +13144,8 @@ pub(crate) fn forward_batch_chunk_impl( batch_semantics.is_independent(), gpu.graphs.capture_mode, gpu.replay.is_recording(), + gpu.arch_caps.is_gfx1100() + && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false), ); let logical_max_ctx = match batch_semantics { BatchSemantics::Sequential => start_pos + n, @@ -14000,13 +14079,21 @@ mod tests { ] { assert!(!packed_mq4_ffn_gate_up_admitted(true, rejected, shape, 256)); assert!(!packed_mq4_ffn_gate_up_admitted(true, shape, rejected, 256)); - assert!(!packed_mq4_ffn_gate_up_admitted(true, rejected, rejected, 256)); + assert!(!packed_mq4_ffn_gate_up_admitted( + true, rejected, rejected, 256 + )); } } #[test] fn packed_mq4_down_refuses_prepared_partial_and_nonuniform_inputs() { - assert!(!packed_mq4_down_admitted(true, true, true, (DType::MQ4G256V2, 5120, 17408), 17408)); + assert!(!packed_mq4_down_admitted( + true, + true, + true, + (DType::MQ4G256V2, 5120, 17408), + 17408 + )); for dtype in [DType::MQ4G256] { let shape = (dtype, 5120, 17408); assert!(packed_mq4_down_admitted(true, true, true, shape, 17408)); @@ -14014,10 +14101,27 @@ mod tests { assert!(!packed_mq4_down_admitted(true, false, true, shape, 17408)); assert!(!packed_mq4_down_admitted(true, true, false, shape, 17408)); assert!(!packed_mq4_down_admitted(true, true, true, shape, 8704)); - assert!(!packed_mq4_down_admitted(true, true, true, (dtype, 2560, 17408), 17408)); + assert!(!packed_mq4_down_admitted( + true, + true, + true, + (dtype, 2560, 17408), + 17408 + )); } - for dtype in [DType::MQ4G256V2Lloyd, DType::MQ4G256Lloyd, DType::HFQ4G256, DType::Q8_0] { - assert!(!packed_mq4_down_admitted(true, true, true, (dtype, 5120, 17408), 17408)); + for dtype in [ + DType::MQ4G256V2Lloyd, + DType::MQ4G256Lloyd, + DType::HFQ4G256, + DType::Q8_0, + ] { + assert!(!packed_mq4_down_admitted( + true, + true, + true, + (dtype, 5120, 17408), + 17408 + )); } } @@ -14037,6 +14141,7 @@ mod tests { false, false, false, + false, )); } } @@ -14066,6 +14171,7 @@ mod tests { is_independent, capture_mode, replay_recording, + false, ) }; assert!(!admitted( @@ -14177,14 +14283,64 @@ mod tests { fn q8_multirow_attn_rejects_replay_recording_on_supported_arches() { for arch in ["gfx1100", "gfx1151", "gfx1201"] { assert!(q8_multirow_attn_admitted( - arch, true, 256, 8, 8192, Some(4096), false, false, false, false, + arch, + true, + 256, + 8, + 8192, + Some(4096), + false, + false, + false, + false, + false, )); assert!(!q8_multirow_attn_admitted( - arch, true, 256, 8, 8192, Some(4096), false, false, false, true, + arch, + true, + 256, + 8, + 8192, + Some(4096), + false, + false, + false, + true, + false, )); } } + #[test] + fn q8_multirow_attn_allows_only_explicit_fixed_grid_capture() { + assert!(q8_multirow_attn_admitted( + "gfx1100", + true, + 256, + 16, + 8192, + Some(4096), + false, + false, + true, + true, + true, + )); + assert!(!q8_multirow_attn_admitted( + "gfx1100", + true, + 256, + 16, + 8192, + Some(4096), + false, + false, + true, + true, + false, + )); + } + #[test] fn fa_pertoken_min_ctx_is_opt_in_on_gfx1151_only() { for arch in ["gfx1100", "gfx1201"] { @@ -15761,22 +15917,42 @@ mod tests { fn native_mq4_wide_rejects_unaudited_dispatch_overrides() { let make = || rdna_compute::FeatureFlags::for_test("gfx1100"); assert!(native_mq4_widened_flags_admitted("gfx1100", &make(), false)); - assert!(!native_mq4_widened_flags_admitted("gfx1201", &make(), false)); + assert!(!native_mq4_widened_flags_admitted( + "gfx1201", + &make(), + false + )); assert!(!native_mq4_widened_flags_admitted("gfx1100", &make(), true)); let mut variants = Vec::new(); - let mut f = make(); f.fp16_disabled = true; variants.push(f); - let mut f = make(); f.rocblas_all_archs = true; variants.push(f); - let mut f = make(); f.gemv_dp4a = Some(true); variants.push(f); - let mut f = make(); f.mmq_override = Some(false); variants.push(f); - let mut f = make(); f.mmq_min_batch = Some(2048); variants.push(f); - let mut f = make(); f.qkvza_split_tail = true; variants.push(f); - let mut f = make(); f.wo_wmma_variant = Some("k4".into()); variants.push(f); + let mut f = make(); + f.fp16_disabled = true; + variants.push(f); + let mut f = make(); + f.rocblas_all_archs = true; + variants.push(f); + let mut f = make(); + f.gemv_dp4a = Some(true); + variants.push(f); + let mut f = make(); + f.mmq_override = Some(false); + variants.push(f); + let mut f = make(); + f.mmq_min_batch = Some(2048); + variants.push(f); + let mut f = make(); + f.qkvza_split_tail = true; + variants.push(f); + let mut f = make(); + f.wo_wmma_variant = Some("k4".into()); + variants.push(f); // Packed FFN shares the budgeted MMQ slot and keeps native tails. let mut packed = make(); packed.packed_mq4_prefill = true; assert!(native_mq4_widened_flags_admitted("gfx1100", &packed, false)); packed.mmq_override = Some(false); - assert!(!native_mq4_widened_flags_admitted("gfx1100", &packed, false)); + assert!(!native_mq4_widened_flags_admitted( + "gfx1100", &packed, false + )); for f in variants { assert!(!native_mq4_widened_flags_admitted("gfx1100", &f, false)); } @@ -15789,14 +15965,27 @@ mod tests { let f16 = 127 * 17408 * 2; let partials = 127 * 5120 * 4 * 4; assert_eq!(mmq, 153 * 1024 * 1024); - assert_eq!(native_mq4_projection_deficit(&config, 8192, 0, 0, 0), - Some(mmq + f16 + partials)); - assert_eq!(native_mq4_projection_deficit(&config, 8192, mmq, f16, partials), Some(0)); - assert_eq!(native_mq4_projection_deficit(&config, 8192, mmq + 1, 0, partials), Some(f16)); - assert_eq!(native_mq4_projection_deficit(&config, usize::MAX, 0, 0, 0), None); + assert_eq!( + native_mq4_projection_deficit(&config, 8192, 0, 0, 0), + Some(mmq + f16 + partials) + ); + assert_eq!( + native_mq4_projection_deficit(&config, 8192, mmq, f16, partials), + Some(0) + ); + assert_eq!( + native_mq4_projection_deficit(&config, 8192, mmq + 1, 0, partials), + Some(f16) + ); + assert_eq!( + native_mq4_projection_deficit(&config, usize::MAX, 0, 0, 0), + None + ); for rows in [512, 1024, 2048, 4096, 8192] { - assert_eq!(native_mq4_projection_deficit(&config, rows, 0, 0, 0), - Some(rows * 19584 + f16 + partials)); + assert_eq!( + native_mq4_projection_deficit(&config, rows, 0, 0, 0), + Some(rows * 19584 + f16 + partials) + ); } } diff --git a/crates/hipfire-arch-qwen35/src/speculative.rs b/crates/hipfire-arch-qwen35/src/speculative.rs index 0360b0880c..2117772fd4 100644 --- a/crates/hipfire-arch-qwen35/src/speculative.rs +++ b/crates/hipfire-arch-qwen35/src/speculative.rs @@ -475,6 +475,38 @@ fn dflash_download_verify_argmax( Ok(host_idx.into_iter().map(|idx| idx as u32).collect()) } +/// Finish a retained/recorded verify forward through the same lm-head and +/// argmax route as the ordinary greedy DFlash path. +/// +/// Redline records only the target forward. The head intentionally remains +/// outside the retained tape, so its recorded-HIP oracle must call this helper +/// instead of substituting per-row GEMV (whose reduction order can differ from +/// the product batched head even when `final_hidden` is byte-identical). +pub fn dflash_finish_retained_lm_head_argmax( + gpu: &mut Gpu, + w_out: &llama::WeightTensor, + final_hidden: &GpuTensor, + verify_scratch: &VerifyScratch, + b: usize, + vocab: usize, +) -> HipResult> { + if dflash_batched_lm_head_supported(w_out.gpu_dtype) { + dflash_enqueue_verify_lm_head_argmax(gpu, w_out, final_hidden, verify_scratch, b, vocab)?; + return dflash_download_verify_argmax(gpu, verify_scratch, b); + } + + let dim = w_out.k; + let mut argmax = Vec::with_capacity(b); + for i in 0..b { + let hidden_row = final_hidden.sub_offset(i * dim, dim); + let logits_row = verify_scratch.logits.sub_offset(i * vocab, vocab); + llama::weight_gemv(gpu, w_out, &hidden_row, &logits_row)?; + let row = gpu.download_f32(&logits_row)?; + argmax.push(argmax_u32(&row)); + } + Ok(argmax) +} + /// Fold a DFlash2 candidate-selector proposal into the chain draft buffers. /// /// Greedy: tokens only — no full-vocab D2H. Temperature: each sparse q row @@ -1414,7 +1446,10 @@ impl DeltaNetSnapshot { pub fn mirrors(&self, state: &DeltaNetState) -> bool { fn same(live: &[GpuTensor], backs: &[DeviceBuffer]) -> bool { live.len() == backs.len() - && live.iter().zip(backs).all(|(t, b)| t.buf.size() == b.size()) + && live + .iter() + .zip(backs) + .all(|(t, b)| t.buf.size() == b.size()) } same(&state.s_matrices, &self.s_matrix_bufs) && same(&state.s_scales, &self.s_scale_bufs) @@ -1981,29 +2016,38 @@ impl GdnTape { } h }; - let stale = ml.from.as_ref().and_then(|f| f.as_ref()).map(|f| f.fingerprint) != Some(Some(fp)); + let stale = ml + .from + .as_ref() + .and_then(|f| f.as_ref()) + .map(|f| f.fingerprint) + != Some(Some(fp)); if stale { let (pre, gdn) = self.replay_ml_rows(ml, weights, config, dn_state); let pre: Vec<_> = pre .into_iter() .enumerate() - .map(|(la, base)| rdna_compute::dflash_gdn_replay::DflashReplayPreLayerFrom { - base, - conv_state_src: snap.conv_state_bufs[la].as_ptr() as u64, - }) + .map( + |(la, base)| rdna_compute::dflash_gdn_replay::DflashReplayPreLayerFrom { + base, + conv_state_src: snap.conv_state_bufs[la].as_ptr() as u64, + }, + ) .collect(); let gdn: Vec<_> = gdn .into_iter() .enumerate() - .map(|(la, base)| rdna_compute::dflash_gdn_replay::GdnLayerTableFrom { - base, - s_q8_src: snap.s_matrix_bufs[la].as_ptr() as u64, - s_scales_src: snap.s_scale_bufs[la].as_ptr() as u64, - ef_src: snap - .s_ef_residual_bufs - .get(la) - .map_or(0, |b| b.as_ptr() as u64), - }) + .map( + |(la, base)| rdna_compute::dflash_gdn_replay::GdnLayerTableFrom { + base, + s_q8_src: snap.s_matrix_bufs[la].as_ptr() as u64, + s_scales_src: snap.s_scale_bufs[la].as_ptr() as u64, + ef_src: snap + .s_ef_residual_bufs + .get(la) + .map_or(0, |b| b.as_ptr() as u64), + }, + ) .collect(); let Some(Some(from)) = ml.from.as_mut() else { unreachable!("armed above") @@ -2342,8 +2386,8 @@ impl GdnReplayMlFrom { if gpu.ensure_dflash_gdn_replay_ml_from().is_err() { return None; } - let pre_bytes = n_la - * std::mem::size_of::(); + let pre_bytes = + n_la * std::mem::size_of::(); let gdn_bytes = n_la * std::mem::size_of::(); let pre_table = gpu.hip.malloc(pre_bytes).ok()?; @@ -2829,6 +2873,14 @@ impl HiddenStateRingBuffer { ); let row_bytes = self.hidden_dim * 4; let bytes = n * row_bytes; + // The software Redline recorder retains typed kernel launches, not HIP + // memcpy commands. During recording, express this byte-identical F32 + // copy as a blob-backed kernel so recorded-HIP and PM4 refresh staging + // on every replay instead of committing stale rows to the hidden ring. + // HipGraph capture keeps the native async memcpy node below. + if gpu.replay.is_recording() { + return gpu.copy_f32_buffer(&self.staging_bufs[extract_idx], src, n * self.hidden_dim); + } if let Some(stream) = gpu.active_stream.as_ref() { gpu.hip.memcpy_dtod_async_at( &self.staging_bufs[extract_idx].buf, @@ -5881,7 +5933,8 @@ pub fn spec_step_dflash( let verify_out = match verify_pm4 { Some(route) => { let replay_failures_before = route.counters().replay_failures; - let hip_windows = |r: &DflashVerifyPm4| r.counters().full_hip + r.counters().partial_hip; + let hip_windows = + |r: &DflashVerifyPm4| r.counters().full_hip + r.counters().partial_hip; let hip_windows_before = hip_windows(route); match verify_dflash_block_retained( gpu, @@ -9211,7 +9264,9 @@ impl SeedPrefill { ) { Ok(limit) => limit, Err(e) => { - eprintln!("dflash seed: chunk-limit query failed ({e}); keeping legacy ceiling"); + eprintln!( + "dflash seed: chunk-limit query failed ({e}); keeping legacy ceiling" + ); qwen35::prefill_max_batch(gpu) } } @@ -9236,9 +9291,12 @@ impl SeedPrefill { return remaining; } let stride = qwen35::prefill::WIDENED_COMMIT_ROWS; - let cap = if ring >= stride { ring / stride * stride } else { ring }; - qwen35::prefill::next_exact_prefill_chunk_len(remaining, cap) - .unwrap_or(remaining.min(cap)) + let cap = if ring >= stride { + ring / stride * stride + } else { + ring + }; + qwen35::prefill::next_exact_prefill_chunk_len(remaining, cap).unwrap_or(remaining.min(cap)) } /// Prefill one piece at `pos`, extracting hidden rows into the ring. @@ -9275,7 +9333,12 @@ impl SeedPrefill { } if self.pbs.is_none() { let rows = self.ceiling.min(total).max(2); - self.pbs = Some(qwen35::PrefillBatchScratch::new_opt(gpu, &target.config, rows, false)?); + self.pbs = Some(qwen35::PrefillBatchScratch::new_opt( + gpu, + &target.config, + rows, + false, + )?); } qwen35::forward_prefill_batch_with_pbs( gpu, @@ -9413,7 +9476,14 @@ pub fn seed_target_hidden_suffix_abortable( let end = off + seed.next_len(suffix.len() - off); while off < end { let piece = SeedPrefill::piece_len(hidden_rb, end - off); - seed.forward(gpu, target, hidden_rb, &suffix[off..off + piece], pos, suffix.len())?; + seed.forward( + gpu, + target, + hidden_rb, + &suffix[off..off + piece], + pos, + suffix.len(), + )?; pos += piece; off += piece; } diff --git a/crates/hipfire-generate/map.md b/crates/hipfire-generate/map.md index 912a751863..426558c258 100644 --- a/crates/hipfire-generate/map.md +++ b/crates/hipfire-generate/map.md @@ -29,7 +29,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/img.rs`](src/img.rs) | 346 | 1 | 0 | | [`src/lib.rs`](src/lib.rs) | 61 | 8 | 0 | | [`src/qwen.rs`](src/qwen.rs) | 8,196 | 70 | 6 | -| [`src/redline.rs`](src/redline.rs) | 6,144 | 52 | 2 | +| [`src/redline.rs`](src/redline.rs) | 6,182 | 52 | 2 | | [`src/vision.rs`](src/vision.rs) | 3,660 | 9 | 8 | ### Public API surface @@ -57,6 +57,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 9 modules · 41,832 lines · 396 public items · 330 tests · 0 examples +- 9 modules · 41,870 lines · 396 public items · 330 tests · 0 examples diff --git a/crates/hipfire-generate/src/redline.rs b/crates/hipfire-generate/src/redline.rs index 001581c16a..b3e1528c26 100644 --- a/crates/hipfire-generate/src/redline.rs +++ b/crates/hipfire-generate/src/redline.rs @@ -1212,7 +1212,8 @@ pub fn redline_bench_decode_deepseek4( let g0 = G0Arm::from_request(msg)?; if g0.is_some() && (capture || product_route || iterations != 1) { return Err( - "g0 requires iterations == 1 and excludes redline_capture/redline_product_route".to_string(), + "g0 requires iterations == 1 and excludes redline_capture/redline_product_route" + .to_string(), ); } if context == 0 || iterations == 0 { @@ -1237,7 +1238,8 @@ pub fn redline_bench_decode_deepseek4( .map_err(|error| format!("bench_decode prefill prime failed: {error}"))?; loaded.seq_pos = context; - if capture || g0.is_some() || (product_route && gpu.replay.prepared_route_identity().is_some()) { + if capture || g0.is_some() || (product_route && gpu.replay.prepared_route_identity().is_some()) + { // Manual capture and prepared product routes are already warm paths. // The first product warmup must still materialize lazy allocations and // record the route; later requests replay from their first timed token. @@ -1255,14 +1257,21 @@ pub fn redline_bench_decode_deepseek4( gpu.replay.begin_replay_observation_window(); } let replay_before = gpu.replay.replay_observation(); - let settled = gpu.hip.device_synchronize().map_err(|error| error.to_string()); + let settled = gpu + .hip + .device_synchronize() + .map_err(|error| error.to_string()); let started = Instant::now(); let run = settled .and_then(|()| { redline_run_deepseek4_decode(gpu, bundle, context, iterations) .map_err(|error| format!("bench_decode forward failed: {error}")) }) - .and_then(|()| gpu.hip.device_synchronize().map_err(|error| error.to_string())); + .and_then(|()| { + gpu.hip + .device_synchronize() + .map_err(|error| error.to_string()) + }); let elapsed = started.elapsed().as_secs_f64(); // A G0 arm is always closed, so a failed forward cannot leave the // controller observing (or recording) every later launch. @@ -2702,10 +2711,7 @@ fn redline_dflash_kv_regions( continue; } if validate { - let expect = slot - .kv_cache - .physical_cap - .saturating_mul(bytes_per_pos); + let expect = slot.kv_cache.physical_cap.saturating_mul(bytes_per_pos); if bytes.len() < expect || bytes.len() > expect + 3 { return Err(format!( "DFlash KV guard: plane bytes {} != physical_cap {} * native stride {bytes_per_pos}", @@ -2892,25 +2898,16 @@ fn redline_dflash_recorded_hip( .hidden_rb .commit_staging_to_ring(gpu, b) .map_err(|e| e.to_string())?; - let mut argmax = Vec::with_capacity(b); - for i in 0..b { - let hidden_row = fixtures - .verify_scratch - .final_hidden - .sub_offset(i * dim, dim); - let logits_row = fixtures.verify_scratch.logits.sub_offset(i * vocab, vocab); - hipfire_runtime::llama::weight_gemv(gpu, &slot.weights.output, &hidden_row, &logits_row) - .map_err(|e| e.to_string())?; - let row = gpu.download_f32(&logits_row).map_err(|e| e.to_string())?; - argmax.push( - row.iter() - .enumerate() - .max_by(|a, b| a.1.total_cmp(b.1)) - .map(|(idx, _)| idx as u32) - .unwrap_or(0), - ); - } - Ok(argmax) + let final_hidden = fixtures.verify_scratch.final_hidden.sub_offset(0, b * dim); + hipfire_arch_qwen35::speculative::dflash_finish_retained_lm_head_argmax( + gpu, + &slot.weights.output, + &final_hidden, + &fixtures.verify_scratch, + b, + vocab, + ) + .map_err(|e| e.to_string()) } fn redline_dflash_run_window( @@ -3088,7 +3085,12 @@ pub fn redline_shadow_dflash_verify_pm4( return Err("DFlash shadow iterations must be non-zero".into()); } let count = iterations.max(12); - let base: usize = 112; + // The gfx1100 split-KV verifier is admitted only beyond its 4K crossover. + // Start the shadow at the 8K boundary so capture/record/PM4 evidence names + // the new partial + merge route rather than a short-context fallback tape. + let split_gfx1100 = gpu.arch == "gfx1100" + && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false); + let base: usize = if split_gfx1100 { 8176 } else { 112 }; let mut positions: Vec = Vec::with_capacity(count); for i in 0..count { positions.push(base + i * batch); @@ -3709,8 +3711,18 @@ fn redline_qwen4_snapshot( let mut lengths = Vec::with_capacity(state.qsa.len()); for (layer, qsa) in state.qsa.iter().enumerate() { let active = [ - ("full_keys", &qsa.full_keys, qsa.full_len, qsa.full_row_units), - ("full_values", &qsa.full_values, qsa.full_len, qsa.full_row_units), + ( + "full_keys", + &qsa.full_keys, + qsa.full_len, + qsa.full_row_units, + ), + ( + "full_values", + &qsa.full_values, + qsa.full_len, + qsa.full_row_units, + ), ("raw_keys", &qsa.raw_index_keys, qsa.raw_len, raw_width), ("pooled_keys", &qsa.pooled_keys, qsa.pooled_len, raw_width), ]; @@ -4093,7 +4105,9 @@ pub fn railgun_g0_dflash_cycle( context: usize, ) -> Result { if loaded.pp > 1 || loaded.ep.is_some() || (loaded.arch_id != 5 && loaded.arch_id != 6) { - return Err("railgun_g0_dflash_cycle requires a single-GPU Qwen3.5/3.8-family target".into()); + return Err( + "railgun_g0_dflash_cycle requires a single-GPU Qwen3.5/3.8-family target".into(), + ); } if hipfire_config::developer_var("HIPFIRE_DFLASH_VERIFY_PM4").as_deref() == Ok("1") { return Err("railgun_g0_dflash_cycle refuses HIPFIRE_DFLASH_VERIFY_PM4=1 (retained verify swaps controllers)".into()); @@ -5599,7 +5613,10 @@ fn redline_trace_dn_bytes( Ok(()) } -fn redline_trace_kv_bytes(gpu: &rdna_compute::Gpu, bundle: &Qwen35Bundle) -> Result, String> { +fn redline_trace_kv_bytes( + gpu: &rdna_compute::Gpu, + bundle: &Qwen35Bundle, +) -> Result, String> { let mut kv = Vec::new(); for tensor in redline_trace_kv_planes(bundle) { redline_append_mapped(gpu, &mut kv, tensor)?; @@ -5618,7 +5635,11 @@ fn redline_trace_upload( ) -> Result<(), String> { let mut offset = 0; for tensor in planes { - let size = if mapped { redline_mapped_len(gpu, tensor) } else { tensor.buf.size() }; + let size = if mapped { + redline_mapped_len(gpu, tensor) + } else { + tensor.buf.size() + }; let chunk = bytes .get(offset..offset + size) .ok_or_else(|| format!("state dump short: plane needs {size} B at {offset}"))?; @@ -5628,7 +5649,10 @@ fn redline_trace_upload( offset += size; } if offset != bytes.len() { - return Err(format!("state dump has {} B, planes took {offset} B", bytes.len())); + return Err(format!( + "state dump has {} B, planes took {offset} B", + bytes.len() + )); } Ok(()) } @@ -5722,7 +5746,10 @@ fn redline_greedy_trace( // `prompt_tokens_file`: raw little-endian u32 token ids. The prime uses // the first `context_tokens`, and the first decoded input defaults to the // file's next token; otherwise the shadow's synthetic prime and 101. - let prompt = match msg.get("prompt_tokens_file").and_then(|value| value.as_str()) { + let prompt = match msg + .get("prompt_tokens_file") + .and_then(|value| value.as_str()) + { None => None, Some(path) => { let bytes = std::fs::read(path).map_err(|error| format!("{path}: {error}"))?; @@ -5731,14 +5758,21 @@ fn redline_greedy_trace( .map(|word| u32::from_le_bytes(word.try_into().expect("4-byte chunk"))) .collect::>(); if tokens.len() < context { - return Err(format!("{path}: {} tokens < context {context}", tokens.len())); + return Err(format!( + "{path}: {} tokens < context {context}", + tokens.len() + )); } Some(tokens) } }; let first_token = number("first_token") .map(|token| token as u32) - .or_else(|| prompt.as_ref().and_then(|tokens| tokens.get(context).copied())) + .or_else(|| { + prompt + .as_ref() + .and_then(|tokens| tokens.get(context).copied()) + }) .unwrap_or(101); let fixed = match msg.get("mode").and_then(|value| value.as_str()) { None | Some("greedy") => false, @@ -5836,14 +5870,15 @@ fn redline_greedy_trace( } Some(source) => { let source = source.join(format!("seq{sequence}")); - let meta: serde_json::Value = serde_json::from_slice( - &std::fs::read(source.join("primed.json")).map_err(io)?, - ) - .map_err(|error| error.to_string())?; + let meta: serde_json::Value = + serde_json::from_slice(&std::fs::read(source.join("primed.json")).map_err(io)?) + .map_err(|error| error.to_string())?; if meta["context_tokens"].as_u64() != Some(context as u64) || meta["poison"].as_u64() != Some(u64::from(poison)) { - return Err(format!("primed state {meta} does not match ctx {context} poison {poison}")); + return Err(format!( + "primed state {meta} does not match ctx {context} poison {poison}" + )); } bundle .kv_cache @@ -5934,7 +5969,10 @@ fn redline_greedy_trace( let x_bytes = hidden.len(); redline_append_buffer(gpu, &mut hidden, &bundle.scratch.tmp.buf)?; if logits.len() < vocab * 4 { - return Err(format!("logits buffer {} < vocab {vocab} f32", logits.len())); + return Err(format!( + "logits buffer {} < vocab {vocab} f32", + logits.len() + )); } let mut best: Option<(usize, f32)> = None; let mut nonfinite = 0usize; diff --git a/crates/rdna-compute/map.md b/crates/rdna-compute/map.md index ec336caea5..9b68b4262a 100644 --- a/crates/rdna-compute/map.md +++ b/crates/rdna-compute/map.md @@ -24,7 +24,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | File | Lines | Public items | Tests | |---|---:|---:|---:| | [`src/arch_caps.rs`](src/arch_caps.rs) | 780 | 60 | 21 | -| [`src/attention.rs`](src/attention.rs) | 23,011 | 286 | 19 | +| [`src/attention.rs`](src/attention.rs) | 23,471 | 288 | 20 | | [`src/bin/hipfire-kernel-hash.rs`](src/bin/hipfire-kernel-hash.rs) | 142 | 0 | 0 | | [`src/bin/hipfire-kernel-manifest.rs`](src/bin/hipfire-kernel-manifest.rs) | 119 | 0 | 0 | | [`src/bin/hipfire-kernel-pack.rs`](src/bin/hipfire-kernel-pack.rs) | 101 | 0 | 0 | @@ -56,7 +56,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/kv_slots.rs`](src/kv_slots.rs) | 526 | 10 | 12 | | [`src/lib.rs`](src/lib.rs) | 120 | 47 | 1 | | [`src/moe.rs`](src/moe.rs) | 3,444 | 45 | 6 | -| [`src/mq_f16_producers.rs`](src/mq_f16_producers.rs) | 567 | 6 | 0 | +| [`src/mq_f16_producers.rs`](src/mq_f16_producers.rs) | 576 | 6 | 0 | | [`src/mq_f16_residual_producers.rs`](src/mq_f16_residual_producers.rs) | 766 | 7 | 0 | | [`src/norm.rs`](src/norm.rs) | 7,870 | 112 | 0 | | [`src/packed_mq4.rs`](src/packed_mq4.rs) | 222 | 3 | 1 | @@ -70,7 +70,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/rdna/gfx1201.rs`](src/rdna/gfx1201.rs) | 414 | 7 | 0 | | [`src/rdna/mod.rs`](src/rdna/mod.rs) | 11 | 1 | 0 | | [`src/replay/railgun_shadow.rs`](src/replay/railgun_shadow.rs) | 878 | 0 | 0 | -| [`src/replay.rs`](src/replay.rs) | 11,444 | 101 | 100 | +| [`src/replay.rs`](src/replay.rs) | 11,737 | 101 | 102 | | [`src/sampling.rs`](src/sampling.rs) | 1,928 | 25 | 4 | | [`src/scratch.rs`](src/scratch.rs) | 2,855 | 54 | 5 | | [`src/select_regrid.rs`](src/select_regrid.rs) | 210 | 10 | 0 | @@ -83,7 +83,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Public API surface - [`src/arch_caps.rs`](src/arch_caps.rs): `ArchCaps`, `note_process_gpu_arch`, `process_gpu_arch`, `new`, `should_use_mmq`, `is_gfx906`, `is_gfx908`, `is_gfx1010`, `is_gfx1011`, `is_gfx1012`, `is_gfx1030`, `is_gfx1031`, +48 more -- [`src/attention.rs`](src/attention.rs): `attention_q8_0_kv_independent_lds_bytes`, `attention_q8_0_kv_independent_max_lane_capacity`, `q8_flash_tile_size`, `GFX12_Q8_FA2_MAX_CTX`, `GFX12_QUERY16_MAX_CTX`, `fp8_e4m3_row_bytes`, `bf16_row_bytes`, `VerifyKv`, `VerifyWmmaPv`, `ALL`, `VerifyWmmaGeometry`, `dspark_stage_kv`, +274 more +- [`src/attention.rs`](src/attention.rs): `attention_q8_0_kv_independent_lds_bytes`, `attention_q8_0_kv_independent_max_lane_capacity`, `q8_flash_tile_size`, `GFX12_Q8_FA2_MAX_CTX`, `GFX12_QUERY16_MAX_CTX`, `fp8_e4m3_row_bytes`, `bf16_row_bytes`, `VerifyKv`, `VerifyWmmaPv`, `ALL`, `VerifyWmmaGeometry`, `dspark_stage_kv`, +276 more - [`src/bin/hipfire-kernel-hash.rs`](src/bin/hipfire-kernel-hash.rs): — - [`src/bin/hipfire-kernel-manifest.rs`](src/bin/hipfire-kernel-manifest.rs): — - [`src/bin/hipfire-kernel-pack.rs`](src/bin/hipfire-kernel-pack.rs): — @@ -152,6 +152,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 56 modules · 175,034 lines · 3999 public items · 476 tests · 224 examples +- 56 modules · 175,796 lines · 4001 public items · 479 tests · 224 examples diff --git a/crates/rdna-compute/src/attention.rs b/crates/rdna-compute/src/attention.rs index 32c1ee1a45..749cbe326f 100644 --- a/crates/rdna-compute/src/attention.rs +++ b/crates/rdna-compute/src/attention.rs @@ -149,7 +149,11 @@ pub const GFX12_QUERY16_MAX_CTX: usize = 32_768; #[derive(Clone, Copy)] enum QresidentOut<'a> { F32(&'a GpuTensor), - A4Slab { qgate: &'a GpuTensor, awq: &'a GpuTensor, x_i4: &'a GpuTensor }, + A4Slab { + qgate: &'a GpuTensor, + awq: &'a GpuTensor, + x_i4: &'a GpuTensor, + }, } const V_MODE_Q8: i32 = 8; @@ -490,7 +494,10 @@ const VERIFY_WMMA_GFX1151: VerifyWmmaKernels = VerifyWmmaKernels { module: "attention_verify_wmma_gfx1151", src: kernels::ATTENTION_VERIFY_WMMA_GFX1151_SRC, qk: "attention_verify_wmma_qk_gfx1151", - pv: ["attention_verify_wmma_pv1s_d4_gfx1151", "attention_verify_wmma_pv2_d2_gfx1151"], + pv: [ + "attention_verify_wmma_pv1s_d4_gfx1151", + "attention_verify_wmma_pv2_d2_gfx1151", + ], }; fn verify_wmma_kernels(gpu: &Gpu) -> Option<&'static VerifyWmmaKernels> { @@ -551,7 +558,11 @@ pub struct VerifyWmmaGeometry { impl Default for VerifyWmmaGeometry { fn default() -> Self { - Self { packed: true, qk_waves: 160, pv: None } + Self { + packed: true, + qk_waves: 160, + pv: None, + } } } @@ -585,6 +596,25 @@ fn q8_multirow_arch_supported(arch: &str) -> bool { matches!(arch, "gfx1100" | "gfx1151" | "gfx1201") } +#[inline] +fn gfx1100_q8_fa2_split_admitted( + enabled: bool, + arch: &str, + logical_ctx_len: usize, + n_heads: usize, + n_kv_heads: usize, + head_dim: usize, + batch_size: usize, +) -> bool { + enabled + && arch == "gfx1100" + && logical_ctx_len > 4096 + && n_heads == 24 + && n_kv_heads == 4 + && head_dim == 256 + && (4..=32).contains(&batch_size) +} + /// `head_dim` envelope of a batched flash-tile kernel, inclusive. /// /// The launcher-level check (`positive multiple of 32, <= 512`) describes the @@ -4374,7 +4404,11 @@ impl Gpu { kernels::kv_slot_desc_source(kernels::ATTENTION_Q8_0_FLASH_PREFILL_SRC, false) } ); - let module = if multi_slot { format!("{module}_paged") } else { module }; + let module = if multi_slot { + format!("{module}_paged") + } else { + module + }; self.ensure_kernel(&module, &src, func)?; } @@ -4570,8 +4604,11 @@ impl Gpu { /// would synchronize, free and malloc mid-capture). A cold capture keeps /// the incumbent route instead of failing the capture. fn gfx12_q8_fa2_capture_ready(&self, batch_size: usize) -> bool { - self.functions.contains_key("attention_q8_0_fa2_gqa_gfx1201") - && self.functions.contains_key("attention_fa2_q_preconvert_gfx1201") + self.functions + .contains_key("attention_q8_0_fa2_gqa_gfx1201") + && self + .functions + .contains_key("attention_fa2_q_preconvert_gfx1201") && !crate::scratch::scratch_will_grow( self.scratch.fa2_q16_scratch_bytes, self.scratch.fa2_q16_scratch.is_some(), @@ -4855,7 +4892,11 @@ impl Gpu { kernels::kv_slot_desc_source(kernel_src, false) } ); - let module = if multi_slot { format!("{module}_paged") } else { module }; + let module = if multi_slot { + format!("{module}_paged") + } else { + module + }; self.ensure_kernel(&module, &src, func)?; } const M_TILE: usize = 16; @@ -5930,8 +5971,18 @@ impl Gpu { self.qresident_launch( "attention_fp8_e4m3_fa2_gqa_qresident_v2_q8_gfx1201", kernels::ATTENTION_FP8_E4M3_FA2_GQA_QRESIDENT_V2_Q8_GFX1201_SRC, - QRESIDENT_V2_LDS_BYTES, true, q, k_cache, v_cache, QresidentOut::F32(out), positions, - n_heads, n_kv_heads, head_dim, max_ctx_len, batch_size, + QRESIDENT_V2_LDS_BYTES, + true, + q, + k_cache, + v_cache, + QresidentOut::F32(out), + positions, + n_heads, + n_kv_heads, + head_dim, + max_ctx_len, + batch_size, ) } @@ -5960,18 +6011,28 @@ impl Gpu { max_ctx_len: usize, batch_size: usize, ) -> HipResult<()> { - if q.dtype != crate::DType::Raw - || q.buf.size() < batch_size * n_heads * (head_dim + 4) - { - return Err(hip_bridge::HipError::new(0, "Q8 resident attention requires Raw codes and scales")); + if q.dtype != crate::DType::Raw || q.buf.size() < batch_size * n_heads * (head_dim + 4) { + return Err(hip_bridge::HipError::new( + 0, + "Q8 resident attention requires Raw codes and scales", + )); } self.ensure_mq_signs()?; self.qresident_launch( "attention_fp8_e4m3_fa2_gqa_qresident_v2_q8_a4epi_gfx1201", kernels::ATTENTION_FP8_E4M3_FA2_GQA_QRESIDENT_V2_Q8_A4EPI_GFX1201_SRC, - QRESIDENT_V2_LDS_BYTES, true, q, k_cache, v_cache, - QresidentOut::A4Slab { qgate, awq, x_i4 }, positions, - n_heads, n_kv_heads, head_dim, max_ctx_len, batch_size, + QRESIDENT_V2_LDS_BYTES, + true, + q, + k_cache, + v_cache, + QresidentOut::A4Slab { qgate, awq, x_i4 }, + positions, + n_heads, + n_kv_heads, + head_dim, + max_ctx_len, + batch_size, ) } @@ -6837,11 +6898,20 @@ impl Gpu { && self.arch == "gfx1151" && hipfire_config::developer_bool("HIPFIRE_GFX1151_FA2_TWIN", true); let (module, preconvert) = if r3 { - ("attention_q8_0_fa2_gqa_gfx1100", "attention_fa2_q_preconvert_gfx1100") + ( + "attention_q8_0_fa2_gqa_gfx1100", + "attention_fa2_q_preconvert_gfx1100", + ) } else if twin { - ("attention_q8_0_fa2_gqa_gfx1151", "attention_fa2_q_preconvert_gfx1151") + ( + "attention_q8_0_fa2_gqa_gfx1151", + "attention_fa2_q_preconvert_gfx1151", + ) } else { - ("attention_q8_0_fa2_gqa_gfx11", "attention_fa2_q_preconvert_gfx11") + ( + "attention_q8_0_fa2_gqa_gfx11", + "attention_fa2_q_preconvert_gfx11", + ) }; // KT32 pinned: the KT64/KT32 ABBA experiment selected KT32 // (32,768 B dynamic LDS, two resident WGs/CU) on both measured @@ -6965,6 +7035,314 @@ impl Gpu { result } + /// Experimental exact-gfx1100 split-KV twin of the tuned Q16/R3 FA2 + /// kernel. The partial grid partitions the live KT32 range across Z and + /// the merge performs a stable LSE-weighted reduction. Product dispatch + /// reaches this only through the fail-closed, opt-in admission in + /// [`Self::attention_flash_q8_0_rows_masked`]. + #[doc(hidden)] + #[allow(clippy::too_many_arguments)] + pub fn attention_q8_0_fa2_gqa_split_gfx1100_bench( + &mut self, + q: &GpuTensor, + k_cache: &GpuTensor, + v_cache: &GpuTensor, + out: &GpuTensor, + positions: &GpuTensor, + partials: &GpuTensor, + n_heads: usize, + n_kv_heads: usize, + head_dim: usize, + batch_size: usize, + n_splits: usize, + ) -> HipResult<()> { + self.bind_thread()?; + if self.arch != "gfx1100" { + return Err(hip_bridge::HipError::new( + 0, + &format!( + "attention_q8_0_fa2_gqa_split_gfx1100_bench requires gfx1100, got {}", + self.arch + ), + )); + } + if n_heads != 24 || n_kv_heads != 4 || head_dim != 256 { + return Err(hip_bridge::HipError::new( + 0, + &format!( + "attention_q8_0_fa2_gqa_split_gfx1100_bench requires H24/KV4/D256, got \ + H{n_heads}/KV{n_kv_heads}/D{head_dim}" + ), + )); + } + if batch_size == 0 || batch_size > 32 { + return Err(hip_bridge::HipError::new( + 0, + &format!( + "attention_q8_0_fa2_gqa_split_gfx1100_bench requires 1 <= batch <= 32, got {batch_size}" + ), + )); + } + if !matches!(n_splits, 1 | 2 | 4 | 8) { + return Err(hip_bridge::HipError::new( + 0, + &format!( + "attention_q8_0_fa2_gqa_split_gfx1100_bench requires 1/2/4/8 splits, got {n_splits}" + ), + )); + } + let need_qo = batch_size + .checked_mul(n_heads) + .and_then(|v| v.checked_mul(head_dim)) + .ok_or_else(|| { + hip_bridge::HipError::new( + 0, + "attention_q8_0_fa2_gqa_split_gfx1100_bench query size overflow", + ) + })?; + let records = batch_size + .checked_mul(n_heads) + .and_then(|v| v.checked_mul(n_splits)) + .ok_or_else(|| { + hip_bridge::HipError::new( + 0, + "attention_q8_0_fa2_gqa_split_gfx1100_bench record count overflow", + ) + })?; + let need_partials = records + .checked_mul(head_dim.checked_add(1).unwrap()) + .ok_or_else(|| { + hip_bridge::HipError::new( + 0, + "attention_q8_0_fa2_gqa_split_gfx1100_bench scratch size overflow", + ) + })?; + if q.numel() < need_qo + || out.numel() < need_qo + || positions.numel() < batch_size + || partials.numel() < need_partials + { + return Err(hip_bridge::HipError::new( + 0, + &format!( + "attention_q8_0_fa2_gqa_split_gfx1100_bench capacity mismatch: \ + q={} out={} positions={} partials={} (need qo>={need_qo}, \ + pos>={batch_size}, partials>={need_partials})", + q.numel(), + out.numel(), + positions.numel(), + partials.numel() + ), + )); + } + + const PARTIAL: &str = "attention_q8_0_fa2_gqa_partial_gfx1100"; + const MERGE: &str = "attention_q8_0_fa2_gqa_merge_gfx1100"; + const PRECONVERT: &str = "attention_fa2_q_preconvert_gfx1100"; + let src = format!( + "// HIPFIRE_COMPILER_FLAGS: -mcumode\n\ + #define HIPFIRE_FA2_KT 32\n\ + #define HIPFIRE_FA2_Q16 1\n\ + #define HIPFIRE_FA2_FILL 1\n\ + #define HIPFIRE_FA2_GFX1100 1\n{}", + kernels::ATTENTION_Q8_0_FA2_GQA_GFX11_SRC + ); + if !self.functions.contains_key(PARTIAL) { + self.ensure_kernel(PARTIAL, &src, PARTIAL)?; + } + if !self.functions.contains_key(PRECONVERT) { + self.ensure_kernel(PARTIAL, &src, PRECONVERT)?; + } + if !self.functions.contains_key(MERGE) { + self.ensure_kernel(MERGE, &src, MERGE)?; + } + + let need_q16_bytes = need_qo.checked_mul(2).ok_or_else(|| { + hip_bridge::HipError::new( + 0, + "attention_q8_0_fa2_gqa_split_gfx1100_bench Q16 size overflow", + ) + })?; + if crate::scratch::scratch_will_grow( + self.scratch.fa2_q16_scratch_bytes, + self.scratch.fa2_q16_scratch.is_some(), + need_q16_bytes, + ) { + self.invalidate_for_scratch_growth(); + } + let q16_ptr = self + .scratch + .ensure_fa2_q16_scratch(&self.hip, need_q16_bytes)?; + self.launch_fa2_q_preconvert_gfx11( + PRECONVERT, + q.buf.as_ptr(), + q16_ptr, + std::ptr::null(), + std::ptr::null(), + batch_size, + 0, + )?; + + let grid_x = batch_size.div_ceil(16) as u32; + let scale = 1.0f32 / (head_dim as f32).sqrt(); + let mut q16_arg = q16_ptr; + let mut k_ptr = k_cache.buf.as_ptr(); + let mut v_ptr = v_cache.buf.as_ptr(); + let mut p_ptr = partials.buf.as_ptr(); + let mut pos_ptr = positions.buf.as_ptr(); + let mut nh = n_heads as i32; + let mut nkv = n_kv_heads as i32; + let mut hd = head_dim as i32; + let mut bs = batch_size as i32; + let mut sc = scale; + let mut ns = n_splits as i32; + let mut params: Vec<*mut c_void> = vec![ + &mut q16_arg as *mut _ as *mut c_void, + &mut k_ptr as *mut _ as *mut c_void, + &mut v_ptr as *mut _ as *mut c_void, + &mut p_ptr as *mut _ as *mut c_void, + &mut pos_ptr as *mut _ as *mut c_void, + &mut nh as *mut _ as *mut c_void, + &mut nkv as *mut _ as *mut c_void, + &mut hd as *mut _ as *mut c_void, + &mut bs as *mut _ as *mut c_void, + &mut sc as *mut _ as *mut c_void, + &mut ns as *mut _ as *mut c_void, + ]; + self.launch_maybe_blob( + PARTIAL, + [grid_x, 4, n_splits as u32], + [256, 1, 1], + 32768, + &mut params, + || { + let mut b = hip_bridge::KernargBlob::new(); + b.push_ptr(q16_arg); + b.push_ptr(k_ptr); + b.push_ptr(v_ptr); + b.push_ptr(p_ptr); + b.push_ptr(pos_ptr); + b.push_i32(nh); + b.push_i32(nkv); + b.push_i32(hd); + b.push_i32(bs); + b.push_f32(sc); + b.push_i32(ns); + b + }, + )?; + + let mut pp_ptr = partials.buf.as_ptr(); + let mut o_ptr = out.buf.as_ptr(); + let mut mbs = batch_size as i32; + let mut mnh = n_heads as i32; + let mut mhd = head_dim as i32; + let mut mns = n_splits as i32; + let mut mparams: Vec<*mut c_void> = vec![ + &mut pp_ptr as *mut _ as *mut c_void, + &mut o_ptr as *mut _ as *mut c_void, + &mut mbs as *mut _ as *mut c_void, + &mut mnh as *mut _ as *mut c_void, + &mut mhd as *mut _ as *mut c_void, + &mut mns as *mut _ as *mut c_void, + ]; + let merge_grid_x = (batch_size * n_heads).div_ceil(8) as u32; + self.launch_maybe_blob( + MERGE, + [merge_grid_x, 1, 1], + [256, 1, 1], + 0, + &mut mparams, + || { + let mut b = hip_bridge::KernargBlob::new(); + b.push_ptr(pp_ptr); + b.push_ptr(o_ptr); + b.push_i32(mbs); + b.push_i32(mnh); + b.push_i32(mhd); + b.push_i32(mns); + b + }, + ) + } + + /// Whether the gfx1100 split-KV verifier can be entered while HipGraph or + /// retained-PM4 capture is active. Capture must never trigger JIT or move + /// the persistent Q16 scratch pointer; the direct warmup/prime window + /// materializes both before either recorder reaches this predicate. + fn gfx1100_q8_fa2_split_capture_ready(&self, batch_size: usize) -> bool { + const PARTIAL: &str = "attention_q8_0_fa2_gqa_partial_gfx1100"; + const MERGE: &str = "attention_q8_0_fa2_gqa_merge_gfx1100"; + const PRECONVERT: &str = "attention_fa2_q_preconvert_gfx1100"; + let need_q16_bytes = batch_size * 24 * 256 * 2; + self.functions.contains_key(PARTIAL) + && self.functions.contains_key(MERGE) + && self.functions.contains_key(PRECONVERT) + && !crate::scratch::scratch_will_grow( + self.scratch.fa2_q16_scratch_bytes, + self.scratch.fa2_q16_scratch.is_some(), + need_q16_bytes, + ) + } + + #[allow(clippy::too_many_arguments)] + fn try_attention_q8_0_fa2_gqa_split_gfx1100( + &mut self, + q: &GpuTensor, + k_cache: &GpuTensor, + v_cache: &GpuTensor, + out: &GpuTensor, + positions: &GpuTensor, + partials: &GpuTensor, + n_heads: usize, + n_kv_heads: usize, + head_dim: usize, + logical_ctx_len: usize, + batch_size: usize, + ) -> HipResult { + const SPLITS: usize = 8; + let capturing = self.graphs.capture_mode || self.replay.is_recording(); + let capture_ready = !capturing || self.gfx1100_q8_fa2_split_capture_ready(batch_size); + if !capture_ready + || !gfx1100_q8_fa2_split_admitted( + hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false), + &self.arch, + logical_ctx_len, + n_heads, + n_kv_heads, + head_dim, + batch_size, + ) + { + return Ok(false); + } + let Some(need_qo) = batch_size + .checked_mul(n_heads) + .and_then(|v| v.checked_mul(head_dim)) + else { + return Ok(false); + }; + let Some(need_partials) = batch_size + .checked_mul(n_heads) + .and_then(|v| v.checked_mul(SPLITS)) + .and_then(|v| v.checked_mul(head_dim + 1)) + else { + return Ok(false); + }; + if q.numel() < need_qo + || out.numel() < need_qo + || positions.numel() < batch_size + || partials.numel() < need_partials + { + return Ok(false); + } + self.attention_q8_0_fa2_gqa_split_gfx1100_bench( + q, k_cache, v_cache, out, positions, partials, n_heads, n_kv_heads, head_dim, + batch_size, SPLITS, + )?; + Ok(true) + } + /// gfx11 (RDNA3) GQA-fused FA2 prefill for fwht3 K (research opt-in). /// /// Same contract as [`Self::attention_q8_0_fa2_gqa_gfx11`] except K is @@ -8208,8 +8586,15 @@ impl Gpu { let nf = (6 * pack).div_ceil(VERIFY_WMMA_ROWS); // The S launch stages K with at least 4 waves (kernel VW_QK_MIN_WAVES). let qk_block = (32 * nf.max(4)) as u32; - let qk_splits = geom.qk_waves.div_ceil(row_groups * n_kv_heads * nf).clamp(1, t_stride); - let pv = geom.pv.unwrap_or(if nf >= 4 { VerifyWmmaPv::Chunk2Whole2 } else { VerifyWmmaPv::Chunk1Split4 }); + let qk_splits = geom + .qk_waves + .div_ceil(row_groups * n_kv_heads * nf) + .clamp(1, t_stride); + let pv = geom.pv.unwrap_or(if nf >= 4 { + VerifyWmmaPv::Chunk2Whole2 + } else { + VerifyWmmaPv::Chunk1Split4 + }); let pv_chunks = pv.chunks(); // One 8-dim V chunk per P.V thread: at least `pv_chunks` waves. let pv_block = (32 * nf.max(pv_chunks)) as u32; @@ -8458,11 +8843,49 @@ impl Gpu { ) } - /// One KV scan for `batch_size` query rows; `Ok(false)` = out of scope, caller must fall back. + /// Compatibility entry for eager callers whose KV stride equals the live + /// causal extent. Capture/retained callers must use + /// [`Self::attention_flash_q8_0_rows_masked_logical`] so an oversized + /// fixed KV allocation cannot satisfy a live-context crossover gate. + #[allow(clippy::too_many_arguments)] + pub fn attention_flash_q8_0_rows_masked( + &mut self, + q: &GpuTensor, + k_cache: &GpuTensor, + v_cache: &GpuTensor, + out: &GpuTensor, + positions: &GpuTensor, + n_heads: usize, + n_kv_heads: usize, + head_dim: usize, + max_ctx_len: usize, + batch_size: usize, + partials: &GpuTensor, + ) -> HipResult { + self.attention_flash_q8_0_rows_masked_logical( + q, + k_cache, + v_cache, + out, + positions, + n_heads, + n_kv_heads, + head_dim, + max_ctx_len, + max_ctx_len, + batch_size, + partials, + ) + } + + /// One KV scan for `batch_size` query rows; `Ok(false)` = out of scope, + /// caller must fall back. `max_ctx_len` is the KV allocation stride while + /// `logical_ctx_len` is the live causal extent used for crossover + /// admission; they intentionally differ during graph/retained capture. /// On arches with a VerifyAttn `rows_q8` twin (gfx1100) an admitted shape /// runs the byte-identical twin instead ([`Self::try_attention_verify_gqa_rows`]). #[allow(clippy::too_many_arguments)] - pub fn attention_flash_q8_0_rows_masked( + pub fn attention_flash_q8_0_rows_masked_logical( &mut self, q: &GpuTensor, k_cache: &GpuTensor, @@ -8473,12 +8896,10 @@ impl Gpu { n_kv_heads: usize, head_dim: usize, max_ctx_len: usize, + logical_ctx_len: usize, batch_size: usize, partials: &GpuTensor, ) -> HipResult { - if self.replay.is_recording() { - return Ok(false); - } if !q8_multirow_arch_supported(self.arch_caps.arch()) { return Ok(false); } @@ -8491,6 +8912,27 @@ impl Gpu { return Ok(false); } self.bind_thread()?; + if self.try_attention_q8_0_fa2_gqa_split_gfx1100( + q, + k_cache, + v_cache, + out, + positions, + partials, + n_heads, + n_kv_heads, + head_dim, + logical_ctx_len, + batch_size, + )? { + return Ok(true); + } + // The incumbent multi-row and VerifyGQA routes are not certified for + // retained recording. The split-KV route above is: it has fixed launch + // geometry, recorder-owned kernargs, and a warmup-stable scratch base. + if self.replay.is_recording() { + return Ok(false); + } if self.try_attention_verify_gqa_rows( q, k_cache, @@ -9488,7 +9930,11 @@ impl Gpu { ), )); } - self.ensure_kernel(TILE, kernels::ATTENTION_FLASH_Q8_0_TILE_GQA_GFX1100_SRC, TILE)?; + self.ensure_kernel( + TILE, + kernels::ATTENTION_FLASH_Q8_0_TILE_GQA_GFX1100_SRC, + TILE, + )?; self.ensure_kernel( REDUCE, kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_AWQ_DEC_GFX1100_SRC, @@ -15229,7 +15675,6 @@ impl Gpu { ) } - /// DFlash draft cross-attention: `B` queries attend to `L` keys/values /// with NO causal mask (bidirectional). Supports GQA; `n_heads` must be /// a multiple of `n_kv_heads`. See `kernels/src/attention_dflash.hip` @@ -22595,7 +23040,6 @@ fn pack_attention_q8_0_fa2_gqa_gfx11_kernarg( b } - /// `*_paged` symbol of a descriptor-aware kernel launched through the /// givens4/turbo assembler, or `None` when the kernel has no paged variant. fn kv_slot_paged_symbol(func: &str) -> Option<&'static str> { @@ -22626,9 +23070,9 @@ fn kv_slot_paged_symbol(func: &str) -> Option<&'static str> { mod tests { use super::{ flux_attn_dtype_error, flux_attn_dtype_suffix, flux_attn_route_dtypes, - flux_attn_route_name, pack_attention_q8_0_fa2_gqa_gfx11_kernarg, - q8_flash_default_tile_size, q8_flash_reduce_safe_tile_size, q8_multirow_arch_supported, - replay_stable_tile_count, + flux_attn_route_name, gfx1100_q8_fa2_split_admitted, + pack_attention_q8_0_fa2_gqa_gfx11_kernarg, q8_flash_default_tile_size, + q8_flash_reduce_safe_tile_size, q8_multirow_arch_supported, replay_stable_tile_count, }; use crate::DType; use std::ffi::c_void; @@ -22695,6 +23139,22 @@ mod tests { } } + #[test] + fn gfx1100_fa2_split_admission_uses_live_context_not_kv_capacity() { + let admitted = |logical_ctx_len, batch_size| { + gfx1100_q8_fa2_split_admitted(true, "gfx1100", logical_ctx_len, 24, 4, 256, batch_size) + }; + assert!(!admitted(4096, 16)); + assert!(admitted(4097, 16)); + assert!(!admitted(262_144, 3)); + assert!(!gfx1100_q8_fa2_split_admitted( + true, "gfx1201", 262_144, 24, 4, 256, 16, + )); + assert!(!gfx1100_q8_fa2_split_admitted( + false, "gfx1100", 262_144, 24, 4, 256, 16, + )); + } + /// The suffix table is the mapping from tensor dtypes to a kernel symbol. /// Get an arm wrong and the launcher asks for an entry that either does /// not exist (a load failure) or exists but reads the buffer at the wrong diff --git a/crates/rdna-compute/src/mq_f16_producers.rs b/crates/rdna-compute/src/mq_f16_producers.rs index d2ad5fea8f..bb3e7e4bb3 100644 --- a/crates/rdna-compute/src/mq_f16_producers.rs +++ b/crates/rdna-compute/src/mq_f16_producers.rs @@ -23,9 +23,10 @@ //! `fp16_x_source_ptr`. //! //! Route contract (mirrored by the prefill hook predicate): exact gfx1100, -//! `DflashFusionCtx::ChainVerify`, N<=16, MQ4G256V2 weights, graph-off and -//! no active replay recording, `HIPFIRE_MQ_F16_PROJECTION_OFF != 1`. Every -//! failed predicate runs the pre-change path; these entries return +//! `DflashFusionCtx::ChainVerify`, N<=16, MQ4G256V2 weights, and +//! `HIPFIRE_MQ_F16_PROJECTION_OFF != 1`. HipGraph/replay recording requires +//! the fixed-grid gfx1100 split-verifier experiment. Every failed predicate +//! runs the pre-change path; these entries return //! `Err` on a non-gfx1100 arch or non-F16 input rather than silently //! falling back. New kernels use `launch_maybe_blob` with the inline //! `KernargBlob` builder (capture-safe ABI, same as the baselines). @@ -228,9 +229,10 @@ impl Gpu { /// accounting) except `xp` is the validated F16 pointer — no /// `ensure_fp16_x`, no `fp16_x_source_ptr` traffic. The MMQ/BT perf /// policies of the base launcher are intentionally absent: callers - /// guarantee the exact route (gfx1100, N<=16, graph-off, no recording), - /// where the base launcher itself falls through to this same base - /// kernel. Calibration taps mirror the `FusedQkvzaMq4G256V2` run-arm. + /// guarantee the exact route (gfx1100, N<=16), where the base launcher + /// itself falls through to this same base kernel. Capture/recording is + /// admitted only by the fixed-grid split-verifier route. Calibration taps + /// mirror the `FusedQkvzaMq4G256V2` run-arm. pub fn gemm_qkvza_mq4g256v2_wmma_f16( &mut self, a_qkv: &GpuTensor, @@ -454,8 +456,10 @@ impl Gpu { /// copy of the `gemm_gate_up_mq4g256v2_wmma` path with the validated F16 /// pointer. Exact-gfx1100 eager HIP defaults to RAW-slab ldsstage when /// eligible (HIPFIRE_GATEUP_LDSSTAGE default-on, 1<=N<=16, K%512==0; `=0` - /// historical base); capture/replay keep base symbol/block32. Taps mirror - /// the `FusedGateUpMq4G256V2` run-arm. + /// historical base). Capture/replay use the same LDS-stage symbol only for + /// the fixed-grid gfx1100 split-verifier experiment; every other recorded + /// route keeps the historical base symbol/block32. Taps mirror the + /// `FusedGateUpMq4G256V2` run-arm. pub fn gemm_gate_up_mq4g256v2_wmma_f16( &mut self, a_gate: &GpuTensor, @@ -484,8 +488,13 @@ impl Gpu { self.maybe_capture_activation(a_up, x_f16, batch_size, k); self.bind_thread()?; // Same guarded tuple as gemm_gate_up_mq4g256v2_wmma small-N eager HIP. - let (kname, ksrc, block_x) = if !self.replay.is_recording() - && !self.graphs.capture_mode + // The fixed-grid split-verifier experiment also supplies replay resource + // contracts for this exact launch, so its capture must preserve the eager + // reduction route instead of silently switching to the base kernel. + let recording = self.replay.is_recording() || self.graphs.capture_mode; + let recording_supported = + !recording || hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false); + let (kname, ksrc, block_x) = if recording_supported && self.arch_caps.is_gfx1100() && self.arch == "gfx1100" && (1..=16).contains(&batch_size) diff --git a/crates/rdna-compute/src/replay.rs b/crates/rdna-compute/src/replay.rs index 23652a06db..d82e80d469 100644 --- a/crates/rdna-compute/src/replay.rs +++ b/crates/rdna-compute/src/replay.rs @@ -25,10 +25,10 @@ use radiowave::{CodeObjectCertification, KernelArgumentAccess, MutableReadCache} use redline_dispatch::aql::{ load_symbols, BatchFencePolicy, Executable, FenceScope, Gfx10DispatchInitiatorPolicy, Gfx10Pm4CommandBuffer, Gfx10SetShRegRecord, Gfx11ComputeResourceLimitsPolicy, - Gfx11DispatchInterleave, Gfx12DispatchPacing, Gfx12Pm4CommandBuffer, Gfx12RmwAcquirePolicy, GpuBatchTiming, GpuDevice, GpuMultiQueueTiming, - GpuSelector, HeaderPolicy, KernargBuffer, KernargPool, Kernel, LaunchGeometry, - PhasedMultiQueuePm4Ib, QueuePolicy, Quiescence, RecordedDispatch, Runtime, - SingleQueueBatchGraph, SingleQueuePm4Ib, + Gfx11DispatchInterleave, Gfx12DispatchPacing, Gfx12Pm4CommandBuffer, Gfx12RmwAcquirePolicy, + GpuBatchTiming, GpuDevice, GpuMultiQueueTiming, GpuSelector, HeaderPolicy, KernargBuffer, + KernargPool, Kernel, LaunchGeometry, PhasedMultiQueuePm4Ib, QueuePolicy, Quiescence, + RecordedDispatch, Runtime, SingleQueueBatchGraph, SingleQueuePm4Ib, }; use redline_dispatch::{ AllocationPolicy, BindingRevision, KernargAbi, KernargField, Recorder, ReplayBindings, @@ -741,20 +741,54 @@ fn pointer_effects(kernel: &str) -> Option> { // The repacker completely overwrites all three planes on each launch. // ADD is a read-modify-write of Y0, represented conservatively as Write. if kernel == "mq4v2_fp8_fragment_repack_gfx1201" { - return Some(vec![read(0), read(8), read(16), read(24), - write(32), write(40), write(48)]); + return Some(vec![ + read(0), + read(8), + read(16), + read(24), + write(32), + write(40), + write(48), + ]); } match kernel { "gemm_mq4g256v2_fp8_set_row_b1" | "gemm_mq4g256v2_fp8_add_row_b1" - | "gemm_mq4g256v2_fp8_silu_row_b1" => - return Some(vec![read(0), read(8), read(16), read(24), read(32), write(40)]), - "gemm_mq4g256v2_fp8_qkv_row_b1" => - return Some(vec![read(0), read(8), read(16), read(24), read(32), - write(40), write(48), write(56)]), - "gemm_mq4g256v2_fp8_qkvza_row_b1" => - return Some(vec![read(0), read(8), read(16), read(24), read(32), - write(40), write(48), write(56), write(64)]), + | "gemm_mq4g256v2_fp8_silu_row_b1" => { + return Some(vec![ + read(0), + read(8), + read(16), + read(24), + read(32), + write(40), + ]) + } + "gemm_mq4g256v2_fp8_qkv_row_b1" => { + return Some(vec![ + read(0), + read(8), + read(16), + read(24), + read(32), + write(40), + write(48), + write(56), + ]) + } + "gemm_mq4g256v2_fp8_qkvza_row_b1" => { + return Some(vec![ + read(0), + read(8), + read(16), + read(24), + read(32), + write(40), + write(48), + write(56), + write(64), + ]) + } _ => {} } if matches!( @@ -995,6 +1029,16 @@ fn pointer_effects(kernel: &str) -> Option> { ) { return Some(vec![read(0), read(8), write(16)]); } + // Exact-gfx1100 small-N gate/up LDS-stage GEMM. Five pointers followed by + // gate_m, up_m, K, N: weights and X are read, split outputs are overwritten. + if kernel == "gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage" { + return Some(vec![read(0), read(8), read(16), write(24), write(32)]); + } + // Typed F32 copy used when hidden-state staging is recorded for retained + // replay: dst, src, n. + if kernel == "copy_f32_buffer" { + return Some(vec![write(0), read(8)]); + } // F16 dense batched GEMM (Maple router + DeepSeek compressor shapes). 3 pointers // + 3 i32 (M,K,B) = 36 explicit bytes. A@0 and X@8 are reads; Y@16 is a pure // overwrite (`Y[...] = acc`), so write — never an RMW. gfx11 and gfx12 are @@ -1078,11 +1122,20 @@ fn pointer_effects(kernel: &str) -> Option> { match kernel { "add_inplace_f32" => Some(vec![write(0), read(8)]), "fused_rmsnorm_mq_rotate" + | "fused_rmsnorm_mq_rotate_f16" | "fused_rmsnorm_mq_rotate_vecsum" | "fused_rmsnorm_mq_rotate_vecsum_sign_const" | "fused_rmsnorm_mq_rotate_vecsum_sign_lds" => { Some(vec![read(0), read(8), read(16), read(24), write(32)]) } + "fused_rmsnorm_mq_rotate_awq_f16" => Some(vec![ + read(0), + read(8), + read(16), + read(24), + read(32), + write(40), + ]), "fused_rmsnorm_mq_rotate_wavegrid" => Some(vec![ read(0), read(8), @@ -1377,6 +1430,14 @@ fn pointer_effects(kernel: &str) -> Option> { read(40), read(48), ]), + // gfx1100 split-KV verifier: preconvert writes persistent Q16 scratch; + // the partial pass scans Q16/K/V/positions into split records; merge + // reads those records and writes the final attention output. + "attention_fa2_q_preconvert_gfx1100" => Some(vec![read(0), write(8), read(16), read(24)]), + "attention_q8_0_fa2_gqa_partial_gfx1100" => { + Some(vec![read(0), read(8), read(16), write(24), read(32)]) + } + "attention_q8_0_fa2_gqa_merge_gfx1100" => Some(vec![read(0), write(8)]), // The gfx1201 GQA fp8, gfx1100 GQA Q8_0 and gfx1151 GQA Q8_0 decode // tiles and the head-dim-split reduces keep their reference twins' // 13/7-argument ABIs and pointer effects. @@ -1406,14 +1467,23 @@ fn pointer_effects(kernel: &str) -> Option> { } fn expected_kernarg_bytes(kernel: &str) -> Option { + if kernel == "fused_rmsnorm_mq_rotate_f16" { + return Some(48); + } + if kernel == "fused_rmsnorm_mq_rotate_awq_f16" { + return Some(64); + } if kernel == "mq4v2_fp8_fragment_repack_gfx1201" { return Some(80); } - if matches!(kernel, "gemm_mq4g256v2_fp8_set_row_b1" - | "gemm_mq4g256v2_fp8_add_row_b1" - | "gemm_mq4g256v2_fp8_silu_row_b1" - | "gemm_mq4g256v2_fp8_qkv_row_b1" - | "gemm_mq4g256v2_fp8_qkvza_row_b1") { + if matches!( + kernel, + "gemm_mq4g256v2_fp8_set_row_b1" + | "gemm_mq4g256v2_fp8_add_row_b1" + | "gemm_mq4g256v2_fp8_silu_row_b1" + | "gemm_mq4g256v2_fp8_qkv_row_b1" + | "gemm_mq4g256v2_fp8_qkvza_row_b1" + ) { return Some(96); } if matches!( @@ -1692,6 +1762,15 @@ fn expected_kernarg_bytes(kernel: &str) -> Option { ) { return Some(48); } + // Five pointers + four i32 = 56 explicit bytes, padded to the recorder's + // required 16-byte boundary. + if kernel == "gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage" { + return Some(64); + } + // Two pointers + one i32 = 20 explicit bytes, padded to 32. + if kernel == "copy_f32_buffer" { + return Some(32); + } // F16 dense batched GEMM: 3 ptr + M,K,B = 36 → 48 padded. gfx11 and gfx12 // share one ABI — see `Gpu::gemm_f16_x_f16_wmma`, whose blob builder pushes // the same 3 ptr + 3 i32 on both paths before the record path's pad_to(16). @@ -1752,6 +1831,9 @@ fn expected_kernarg_bytes(kernel: &str) -> Option { | "hc_input_map_4stream" | "sigmoid_mul_f32" => Some(32), "gemma4_ple_gelu_mul_strided_f32" => Some(48), + "attention_fa2_q_preconvert_gfx1100" => Some(48), + "attention_q8_0_fa2_gqa_partial_gfx1100" => Some(64), + "attention_q8_0_fa2_gqa_merge_gfx1100" => Some(32), "attention_flash_q8_0_reduce" | "attention_flash_reduce_dsplit_gfx1201" | "attention_flash_reduce_dsplit_gfx1151" @@ -1990,15 +2072,22 @@ fn encode_bound_kernarg( let mut encoded = snapshot.to_vec(); for slot in &layout.slots { let binding = bindings.resource(slot.resource).ok_or_else(|| { - format!("{kernel}: pointer slot at offset {} has no bound resource", slot.offset) + format!( + "{kernel}: pointer slot at offset {} has no bound resource", + slot.offset + ) })?; let base = binding.base().as_ptr() as usize as u64; let address = base.checked_add(slot.interior_offset).ok_or_else(|| { - format!("{kernel}: pointer slot at offset {} base overflows", slot.offset) - })?; - let end = slot.offset.checked_add(8).ok_or_else(|| { - format!("{kernel}: pointer slot at offset {} overflows", slot.offset) + format!( + "{kernel}: pointer slot at offset {} base overflows", + slot.offset + ) })?; + let end = slot + .offset + .checked_add(8) + .ok_or_else(|| format!("{kernel}: pointer slot at offset {} overflows", slot.offset))?; if end > encoded.len() { return Err(format!( "{kernel}: pointer slot at offset {} out of bounds (len {})", @@ -2016,11 +2105,7 @@ fn encode_bound_kernarg( /// except the moved pointer slots — checked by the refresh path with its own /// slot mask). `debug_assert` covers debug builds; `HIPFIRE_REPLAY_BINDINGS_VERIFY=1` /// promotes the check to a hard error in release for the harness. -fn verify_bound_kernarg( - kernel: &str, - snapshot: &[u8], - encoded: &[u8], -) -> Result<(), String> { +fn verify_bound_kernarg(kernel: &str, snapshot: &[u8], encoded: &[u8]) -> Result<(), String> { let mismatch = bound_kernarg_mismatch_message(kernel, snapshot, encoded); debug_assert!( mismatch.is_none(), @@ -2038,11 +2123,7 @@ fn verify_bound_kernarg( /// Pure mismatch description for the snapshot gate: `None` when the /// re-encoded segment equals the snapshot, else the fail-closed message with /// the launch name and the first differing offset. -fn bound_kernarg_mismatch_message( - kernel: &str, - snapshot: &[u8], - encoded: &[u8], -) -> Option { +fn bound_kernarg_mismatch_message(kernel: &str, snapshot: &[u8], encoded: &[u8]) -> Option { if snapshot == encoded { return None; } @@ -4596,7 +4677,10 @@ impl PreparedPm4Replay { if let Ok(timing) = &result { crate::gap_timing::add_ns( crate::gap_timing::Slot::GpuSpan, - timing.last_end.saturating_sub(timing.first_start).saturating_mul(1_000_000_000) + timing + .last_end + .saturating_sub(timing.first_start) + .saturating_mul(1_000_000_000) / timing.frequency_hz.max(1), ); } @@ -5901,9 +5985,14 @@ impl ReplayController { } if shadow_on { shadow_decisions.push(railgun::shadow::RedlineDecision { - effects: redline_effect_source(&self.radiowave_effect_certifications, current_launch), + effects: redline_effect_source( + &self.radiowave_effect_certifications, + current_launch, + ), resource_independent: resources_independent, - name_acquire: self.pm4_mid_acquire_policy.acquire_between(previous, current), + name_acquire: self + .pm4_mid_acquire_policy + .acquire_between(previous, current), pre_dispatch_name: gfx12_pre_dispatch_acquire, }); } @@ -5911,7 +6000,10 @@ impl ReplayController { resource_frontier.advance(&self.recorded[index], false); if shadow_on { shadow_decisions.push(railgun::shadow::RedlineDecision { - effects: redline_effect_source(&self.radiowave_effect_certifications, &self.recorded[index]), + effects: redline_effect_source( + &self.radiowave_effect_certifications, + &self.recorded[index], + ), ..Default::default() }); } @@ -6061,7 +6153,9 @@ impl ReplayController { match executable { None => {} Some(Ok(railgun_commands)) => commands = railgun_commands, - Some(Err(reason)) => return Err(format!("railgun backend refused: {reason}")), + Some(Err(reason)) => { + return Err(format!("railgun backend refused: {reason}")) + } } } } @@ -6100,8 +6194,15 @@ impl ReplayController { )?; (PreparedPm4Graph::Single(graph), command_dwords) } else { - if self.railgun_shadow.as_ref().is_some_and(railgun_shadow::RailgunShadow::backend_railgun) { - return Err("railgun backend refused: multi-queue PM4 tapes are not lowered by railgun".to_owned()); + if self + .railgun_shadow + .as_ref() + .is_some_and(railgun_shadow::RailgunShadow::backend_railgun) + { + return Err( + "railgun backend refused: multi-queue PM4 tapes are not lowered by railgun" + .to_owned(), + ); } let min_parallel_width = pm4_min_parallel_width_from_config(); let min_parallel_workgroups = pm4_min_parallel_workgroups_from_config(); @@ -6467,9 +6568,16 @@ impl ReplayController { pub(crate) fn prepared_pm4_submitted_kernargs(&mut self) -> Result>, String> { let prepared = self.prepared_pm4.as_mut().ok_or("no prepared PM4 replay")?; if !prepared.dynamic_grids.is_empty() { - return Err(format!("{} dynamic grid patches", prepared.dynamic_grids.len())); + return Err(format!( + "{} dynamic grid patches", + prepared.dynamic_grids.len() + )); } - Ok(prepared.kernargs.iter_mut().map(|k| k.as_mut_bytes().to_vec()).collect()) + Ok(prepared + .kernargs + .iter_mut() + .map(|k| k.as_mut_bytes().to_vec()) + .collect()) } /// Generation of the prepared PM4 program (see `pm4_generation`). @@ -6479,7 +6587,9 @@ impl ReplayController { /// The prepared railgun program's surfaces (`None` without the shadow). pub(crate) fn railgun_program_surfaces(&self) -> Option { - self.railgun_shadow.as_ref().and_then(|s| s.surfaces().cloned()) + self.railgun_shadow + .as_ref() + .and_then(|s| s.surfaces().cloned()) } pub fn prepared_pm4_dispatch_boundaries(&self) -> Option<&[Pm4DispatchBoundary]> { @@ -6882,7 +6992,8 @@ impl ReplayController { // tape-global `ResourceId` now, while the allocations are live. Any // unresolvable slot (or unknown kernel) leaves the launch untyped on // the raw snapshot path — never a partial slot set. - let binding_layout = self.build_binding_layout(hip, kernel, kernarg, certified_effects.as_deref()); + let binding_layout = + self.build_binding_layout(hip, kernel, kernarg, certified_effects.as_deref()); let before = self.recorded.len(); self.record_hip_launch_with_accesses( kernel, @@ -6898,8 +7009,21 @@ impl ReplayController { ); if self.recorded.len() > before { if let Some(shadow) = self.railgun_shadow.as_mut() { - let artifact = self.recorded.last().and_then(|launch| launch.artifact.as_deref()); - shadow.observe(hip, compiler, kernel, artifact, grid, block, shared_mem, kernarg, declared_words); + let artifact = self + .recorded + .last() + .and_then(|launch| launch.artifact.as_deref()); + shadow.observe( + hip, + compiler, + kernel, + artifact, + grid, + block, + shared_mem, + kernarg, + declared_words, + ); } } } @@ -7133,7 +7257,9 @@ impl ReplayController { .as_mut() .expect("prepared PM4 route checked above"); if prepared.bound_explicit_lens.len() != prepared.kernargs.len() { - return Err("retained PM4 binding cache disagrees with prepared kernarg count".to_owned()); + return Err( + "retained PM4 binding cache disagrees with prepared kernarg count".to_owned(), + ); } for (index, kernarg) in prepared.kernargs.iter_mut().enumerate() { let explicit_len = prepared.bound_explicit_lens[index]; @@ -7143,8 +7269,12 @@ impl ReplayController { let Some(layout) = launch.binding_layout.as_ref() else { continue; }; - let encoded = - encode_bound_kernarg(&launch.kernarg, layout, &self.replay_bindings, &launch.kernel)?; + let encoded = encode_bound_kernarg( + &launch.kernarg, + layout, + &self.replay_bindings, + &launch.kernel, + )?; let bytes = kernarg.as_mut_bytes(); if explicit_len > encoded.len() || explicit_len > bytes.len() { return Err(format!( @@ -7496,8 +7626,12 @@ mod tests { snapshot[16..24].copy_from_slice(&base_b.to_ne_bytes()); snapshot[24..28].copy_from_slice(&0x0000_0007u32.to_ne_bytes()); let mut issuer = Recorder::new(); - let resource_a = issuer.resource("fixture-a", 0x1_0000).expect("valid resource"); - let resource_b = issuer.resource("fixture-b", 0x2_0000).expect("valid resource"); + let resource_a = issuer + .resource("fixture-a", 0x1_0000) + .expect("valid resource"); + let resource_b = issuer + .resource("fixture-b", 0x2_0000) + .expect("valid resource"); let mut bindings = ReplayBindings::new(); // SAFETY: synthetic non-deref'd addresses used only as binding identity // in unit tests; sizes match the fixture resources; never launched. @@ -8354,6 +8488,152 @@ mod tests { } } + #[test] + fn gfx1100_fa2_split_verify_keeps_padded_replay_contracts() { + use RecordedAccessMode::{Read, Write}; + + let preconvert = "attention_fa2_q_preconvert_gfx1100"; + let mut blob = hip_bridge::KernargBlob::new(); + for _ in 0..4 { + blob.push_ptr(std::ptr::null()); + } + blob.push_i32(0); + blob.push_i32(0); + blob.pad_to(16); + assert_eq!(expected_kernarg_bytes(preconvert), Some(blob.len())); + assert_eq!( + pointer_effects(preconvert) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![(0, Read), (8, Write), (16, Read), (24, Read)] + ); + + let partial = "attention_q8_0_fa2_gqa_partial_gfx1100"; + let mut blob = hip_bridge::KernargBlob::new(); + for _ in 0..5 { + blob.push_ptr(std::ptr::null()); + } + for _ in 0..4 { + blob.push_i32(0); + } + blob.push_f32(0.0); + blob.push_i32(0); + blob.pad_to(16); + assert_eq!(expected_kernarg_bytes(partial), Some(blob.len())); + assert_eq!( + pointer_effects(partial) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![(0, Read), (8, Read), (16, Read), (24, Write), (32, Read)] + ); + + let merge = "attention_q8_0_fa2_gqa_merge_gfx1100"; + let mut blob = hip_bridge::KernargBlob::new(); + blob.push_ptr(std::ptr::null()); + blob.push_ptr(std::ptr::null()); + for _ in 0..4 { + blob.push_i32(0); + } + blob.pad_to(16); + assert_eq!(expected_kernarg_bytes(merge), Some(blob.len())); + assert_eq!( + pointer_effects(merge) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![(0, Read), (8, Write)] + ); + } + + #[test] + fn gfx1100_f16_projection_producers_keep_replay_contracts() { + use RecordedAccessMode::{Read, Write}; + + let base = "fused_rmsnorm_mq_rotate_f16"; + let mut blob = hip_bridge::KernargBlob::new(); + for _ in 0..5 { + blob.push_ptr(std::ptr::null()); + } + blob.push_i32(0); + blob.push_f32(0.0); + blob.pad_to(16); + assert_eq!(expected_kernarg_bytes(base), Some(blob.len())); + assert_eq!( + pointer_effects(base) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![(0, Read), (8, Read), (16, Read), (24, Read), (32, Write)] + ); + + let awq = "fused_rmsnorm_mq_rotate_awq_f16"; + let mut blob = hip_bridge::KernargBlob::new(); + for _ in 0..6 { + blob.push_ptr(std::ptr::null()); + } + blob.push_i32(0); + blob.push_f32(0.0); + blob.pad_to(16); + assert_eq!(expected_kernarg_bytes(awq), Some(blob.len())); + assert_eq!( + pointer_effects(awq) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![ + (0, Read), + (8, Read), + (16, Read), + (24, Read), + (32, Read), + (40, Write), + ] + ); + + let gate_up = "gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage"; + let mut blob = hip_bridge::KernargBlob::new(); + for _ in 0..5 { + blob.push_ptr(std::ptr::null()); + } + for _ in 0..4 { + blob.push_i32(0); + } + blob.pad_to(16); + assert_eq!(blob.len(), 64); + assert_eq!(expected_kernarg_bytes(gate_up), Some(blob.len())); + assert_eq!( + pointer_effects(gate_up) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![(0, Read), (8, Read), (16, Read), (24, Write), (32, Write)] + ); + + let copy = "copy_f32_buffer"; + let mut blob = hip_bridge::KernargBlob::new(); + blob.push_ptr(std::ptr::null()); + blob.push_ptr(std::ptr::null()); + blob.push_i32(0); + blob.pad_to(16); + assert_eq!(expected_kernarg_bytes(copy), Some(blob.len())); + assert_eq!( + pointer_effects(copy) + .unwrap() + .iter() + .map(|effect| (effect.offset, effect.mode)) + .collect::>(), + vec![(0, Write), (8, Read)] + ); + } + #[test] fn gfx1201_qwen36_27b_decode_fusions_keep_padded_replay_contract() { use RecordedAccessMode::{Read, Write}; @@ -8401,7 +8681,15 @@ mod tests { .collect(); assert_eq!( modes, - vec![(0, Read), (8, Write), (16, Write), (24, Write), (32, Read), (40, Read), (48, Read)] + vec![ + (0, Read), + (8, Write), + (16, Write), + (24, Write), + (32, Read), + (40, Read), + (48, Read) + ] ); } @@ -10257,8 +10545,13 @@ mod tests { assert!(controller.recorded_launches().is_empty()); let seen = controller.finish_g0_observation().unwrap(); assert_eq!( - seen.iter().map(|l| (l.kernel.as_str(), l.grid, l.shared_mem, l.kernarg.clone())).collect::>(), - vec![("a", [1, 2, 3], 0, vec![1, 2]), ("b", [4, 1, 1], 128, vec![3])] + seen.iter() + .map(|l| (l.kernel.as_str(), l.grid, l.shared_mem, l.kernarg.clone())) + .collect::>(), + vec![ + ("a", [1, 2, 3], 0, vec![1, 2]), + ("b", [4, 1, 1], 128, vec![3]) + ] ); assert!(controller.finish_g0_observation().is_err()); @@ -10566,7 +10859,7 @@ mod tests { None, &declared, None, - None, + None, ); let snapshot = earlier.snapshot_recorded_kernargs(); @@ -10584,7 +10877,7 @@ mod tests { None, &declared, None, - None, + None, ); let synthesized = current @@ -10624,7 +10917,7 @@ mod tests { None, &declared, None, - None, + None, ); let launches = controller.recorded_launches().to_vec(); let mut bindings: Vec<(usize, ReplayKernargBinding)> = vec![( @@ -10693,7 +10986,7 @@ mod tests { None, &declared, None, - None, + None, ); let launches = controller.recorded_launches().to_vec(); let mut bindings: Vec<(usize, ReplayKernargBinding)> = Vec::new(); @@ -10719,7 +11012,7 @@ mod tests { None, declared, None, - None, + None, ); replay_sequence_hash(controller.recorded_launches()) }; diff --git a/docs/env-vars.md b/docs/env-vars.md index a4b923e318..d67aa9b121 100644 --- a/docs/env-vars.md +++ b/docs/env-vars.md @@ -216,6 +216,7 @@ Read only by the Qwen4 carrier and its kernels; no other model reads them. | `HIPFIRE_GFX11_FA2_PREFILL` | GQA-fused FA2 prefill on gfx1100/gfx1151 (Qwen NH24/NKV4/HD256, N 64..512 step 16, ctx 64..32768) — default ON (`kernel.gfx11_fa2_prefill`); `=0` opts out toward the byte-identical incumbent | | `HIPFIRE_FA2_FILL` | Warp-specialized K/V fill in that FA2 kernel on gfx1100/gfx1151 (bit-exact; helper waves dequantize the next K/V tile while compute waves run QK/PV) — default ON; `=0` restores the all-wave per-tile fill | | `HIPFIRE_GFX1100_FA2_R3` | Exact-gfx1100 variant of that FA2 fill body (bit-exact; CU mode, bank-conflict-free helper plane stores, O rescale skipped when alpha is exactly 1, heaviest q tiles first; symbols `attention_q8_0_fa2_gqa_gfx1100` / `attention_fa2_q_preconvert_gfx1100`) — default ON; `=0` restores the shared gfx11 body | +| `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | Experimental, default off, exact gfx1100: `1` replaces the Q8 DFlash verifier's eager R4/R8 attention with the FA2 split-KV S8 route for dense Qwen H24/NKV4/HD256 batches 4..32 once the **live logical context** exceeds 4,096. HipGraph/retained recording additionally require precompiled kernels and fully materialized fixed-address Q16 scratch; unsupported or not-yet-ready shapes fail closed to the established batched route. The Redline/PM4 product admission remains B=16 and inherits its existing single-GPU/state guards. | | `HIPFIRE_GFX1151_FA2_TWIN` | Exact-gfx1151 twin of that FA2 fill kernel (CU mode, heaviest q-tile first, conflict-free helper V stores; bit-exact) — default ON; `=0` restores the gfx11 module | | `HIPFIRE_GFX12_FA2_PREFILL` | GQA-fused FA2 prefill on exact gfx1201 (same Qwen NH24/NKV4/HD256 envelope) — default ON (`kernel.gfx12_fa2_prefill`); `=0` opts out toward the byte-identical incumbent | | `HIPFIRE_GFX12_FA_PACKET` | Packet-minimal Q128 FA2 body on exact gfx1201 (same Qwen envelope as `HIPFIRE_GFX12_FA2_PREFILL`) — default ON (`kernel.gfx12_fa_packet`); `=0` opts out to the byte-identical route-N body | @@ -1050,6 +1051,7 @@ Presence in the inventory means the token appears in source; it does **not** mea | `HIPFIRE_GFX1100_DENSE_GATE_UP_SETPRIO` | crates/rdna-compute/src/gemm.rs | developer | | `HIPFIRE_GFX1100_DENSE_GATE_UP_STAGE_X32` | crates/rdna-compute/src/gemm.rs | developer | | `HIPFIRE_GFX1100_FA2_R3` | crates/rdna-compute/src/attention.rs | developer | +| `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | crates/hipfire-arch-qwen35/src/dflash_spec.rs, crates/hipfire-arch-qwen35/src/qwen35/prefill.rs | developer | | `HIPFIRE_GFX1100_FA_PREP` | crates/hipfire-arch-qwen35/src/qwen35/prefill.rs | developer | | `HIPFIRE_GFX1100_GATED_NORM_V2` | crates/rdna-compute/src/gemv.rs, crates/rdna-compute/src/kernels.rs | developer | | `HIPFIRE_GFX1100_MQ4V2_NINEPATH_RPB8` | crates/hipfire-runtime/examples/mq4v2_fused_parity.rs | harness | diff --git a/docs/perf-checkpoints/2026-10-05-gfx1100-fa2-splitkv-verifier-graph-pm4.md b/docs/perf-checkpoints/2026-10-05-gfx1100-fa2-splitkv-verifier-graph-pm4.md new file mode 100644 index 0000000000..dce8acb24b --- /dev/null +++ b/docs/perf-checkpoints/2026-10-05-gfx1100-fa2-splitkv-verifier-graph-pm4.md @@ -0,0 +1,128 @@ +# gfx1100 FA2 split-KV verifier with HipGraph and Redline/PM4 — 2026-10-05 + +**Lifecycle:** `historical` + +**Disposition:** measured opt-in candidate evidence; not a product default, +current baseline, or admission decision. + +## Question and scope + +Measure an exact-gfx1100 Q8 DFlash verifier route that keeps the packed-KV +front end and replaces the established R4/R8 attention step with an FA2 +split-KV S8 partial/merge back end. Then determine whether the same fixed-grid +route can be safely retained by HipGraph and Redline/PM4. + +The candidate is default-off behind +`HIPFIRE_GFX1100_FA2_SPLIT_VERIFY=1`. It fails closed unless the target is +dense Qwen H24/NKV4/HD256, the verify batch is 4..32, Q8 KV is active, and the +live logical context exceeds 4,096. Capture additionally requires precompiled +kernels and fully materialized fixed-address Q16 scratch. Redline product +admission retains its existing B=16, single-GPU, Q8-state, and route guards. + +Source base: `d5305333d4c609848f09ce1cd98502bb2ee2fe82` (`warpfront/beta`). + +## Fixture identity + +- Host GPU: Radeon Pro W7900, exact `gfx1100`, HIP 7.15, GPU0. +- Target: `qwen3.8-27b.mq4-xt`, SHA-256 + `9f91556f7e0431a077d03756a7102d0154108757289e6e5fe9a2d204c0c9eeb7`, + MD5 `e45d15bfe0c9a87132697101d17cbed6`. +- Draft: `qwen38-27b-dflash-mq4.hfq`, SHA-256 + `d0a74a232a0e2166d889f823e91e0fbf778d21dd9668d7de055cdecb065401bc`, + MD5 `013395583cd04206c8aa68f4d061983d`. +- Prompt: `benchmarks/prompts/qwen38_issue693_longcode_20676.txt`, 21,550 + actual tokens, MD5 `b4d0b63cddcac872648ddf3cdd92cac2`. +- Product settings: Q8 VMM KV, DFlash, greedy sampling, 200 output tokens, + `max_seq=65536`. + +## Fresh-process product A/B + +Six graph-off fresh processes ran in declared order +`off,on,on,off,off,on`, with one unrecorded warmup before each measured run. +The baseline is the established gfx1100 R4/R8 verifier; the candidate changes +only the split-verifier route. + +- CLI MD5: `9bbfbbaac68ed262867a6e7136485081`. +- Daemon MD5: `f0c79a77e2cb1b019ee58bbae96de113`. + +| route | decode samples (tok/s) | median | tau / cycles | delta | +|---|---|---:|---:|---:| +| established R4/R8 | 43.5, 42.6, 42.5 | 42.6 | 1.97 / 67 | — | +| FA2 split-KV S8 | 48.5, 48.3, 47.8 | 48.3 | 1.97 / 67 | +13.38% | + +The unchanged tau and cycle count isolate the speedup to verifier execution, +not improved draft acceptance. + +## Kernel screen + +Batch 16 used 10 warmups and 30 measured repetitions per cell. + +| logical context | established R4/R8 | split-KV S8 | speedup | +|---:|---:|---:|---:| +| 8,192 | 368.80 us | 177.36 us | 2.079x | +| 20,676 | 828.45 us | 410.72 us | 2.017x | +| 32,768 | 1,369.01 us | 648.21 us | 2.112x | + +S1 was bit-identical to direct FA2. S8 relative L2 against direct FA2 was +about `2.46e-4`, `2.50e-4`, and `2.58e-4` respectively; cosine was about +`0.999999970`. The direct and partial kernels compiled at 254/255 VGPR, the +merge at 18 VGPR, with no spills or private scratch. + +## HipGraph validation + +The warpfront-beta graph-validation binaries were CLI MD5 +`9bbfbbaac68ed262867a6e7136485081` and daemon MD5 +`f0c79a77e2cb1b019ee58bbae96de113`. + +- B=16 capture retained 706 launch blobs. +- Graph off/on decoded at 4.8/4.9 tok/s in this separate cold + `serve_harness.py` comparison, with `tau=2.06` and 200 generated tokens. +- Graph off/on transcripts were byte-identical, MD5 + `7b8dc5b28daef60f803fe2a466c888b2`. + +The two beta runs used the same prompt MD5 and request MD5 +`8a54e9aa236f678362f89864bf000125`. Both ended inside the model's hidden +thought channel and therefore reported `RUNAWAY,EMPTY`; their rates are not a +performance claim and the run is route/output parity evidence, not +answer-quality evidence. The fresh-process native bench above remains the +performance measurement. + +## Redline/PM4 validation + +The route-scoped daemon harness exercised 12 consecutive B=16 windows from +positions 8,176 through 8,352 across HipAuto, capture-safe direct HIP, +recorded HIP, and PM4. + +- Result: PASS; backend `pm4_ib`. +- Tape: 711 launches/dispatches, 18 unique typed AQL kernel contracts, + one packet, one queue, and one phase; dispatch count matched launch count. +- Retained route: ready; 17 successful replays; zero contract, preparation, + or replay failures. +- All four arms agreed exactly in every window on tokens, argmax, hidden + staging/ring, final hidden, logits, active KV hashes, GDN intermediates, + recurrent state after forward, and recurrent state after rollback. +- Five interleaved timing windows: HipGraph median 46.413 ms, PM4 median + 43.630 ms, delta -2.783 ms (-6.00%); p95 delta -2.652 ms. + +The five-window PM4 timing is a retained-route smoke measurement, not a +standalone headline throughput claim. + +## Validation and interpretation + +- `test_kernels`: 17 passed on the W7900. +- `rdna-compute`: 438 passed, 12 ignored. +- `hipfire-arch-qwen35`: 247 passed, 23 ignored. +- `hipfire-generate`: 54 passed. +- `scripts/no-gpu-ci.sh`: Rust/main gates passed. Its Python stage reported + 21 failures and 541 passes because `scripts/hw-gate/select.py` shadows the + standard-library `select` module when `scripts/hw-gate/review.py` imports + `subprocess`. This reproduces on the unchanged beta files; this PR does not + touch `scripts/hw-gate`, `tests`, or `scripts/no-gpu-ci.sh`. +- Crate maps, lifecycle/env inventory, changed-file formatting, and diff + whitespace checks passed. + +These results support review of the narrowly gated gfx1100 candidate and its +HipGraph/Redline retention contracts. They do not transfer to other GPUs, KV +formats, model shapes, prompts, drafts, or sampling modes. The route remains +default-off, and low draft-target agreement can still make DFlash slower than +plain autoregressive decode. diff --git a/kernels/src/attention_q8_0_fa2_gqa.gfx11.hip b/kernels/src/attention_q8_0_fa2_gqa.gfx11.hip index 8862942416..7c1120d5e3 100644 --- a/kernels/src/attention_q8_0_fa2_gqa.gfx11.hip +++ b/kernels/src/attention_q8_0_fa2_gqa.gfx11.hip @@ -508,12 +508,16 @@ __device__ __forceinline__ void fa2_fill_softmax_pv( } #endif -// Shared workgroup body (direct only: writes normalized O to `out`). +// Shared workgroup body. PARTIAL=false writes normalized O to `out`; +// PARTIAL=true writes one LSE scalar plus 256 normalized output values for +// each (query, head, split). `split` owns a disjoint range of the live KT32 +// tiles derived from positions[], so every split scans only useful KV rows. // `kv_h` selects the KV head (blockIdx.y); `q_base` the 16-position tile // (blockIdx.x). Q arrives f16 pre-converted row-major [batch, 24, 256] // (see attention_fa2_q_preconvert_gfx11); the body never touches f32 Q. // Causal bounds come from positions[] (mailbox-reduced to gmin/gmax); // the tile loop clamps to the derived seq_len. +template __device__ __forceinline__ void fa2_gqa_body( const _Float16* __restrict__ q16, const unsigned char* __restrict__ k_cache, @@ -523,7 +527,9 @@ __device__ __forceinline__ void fa2_gqa_body( int batch_size, float scale_attn, int kv_h, - int q_base) + int q_base, + int split, + int n_splits) { constexpr int KT = HIPFIRE_FA2_KT; constexpr int KT_SUBS = KT / 16; @@ -610,6 +616,10 @@ __device__ __forceinline__ void fa2_gqa_body( const int gmax = __builtin_amdgcn_readfirstlane((int)LDS[LDS_DWORDS - 2]); const int gmin = __builtin_amdgcn_readfirstlane((int)LDS[LDS_DWORDS - 1]); const int seq_len = gmax + 1; + const int tiles_total = (seq_len + KT - 1) / KT; + const int tiles_per_split = (tiles_total + n_splits - 1) / n_splits; + const int tile0 = split * tiles_per_split; + const int tile1 = min(tile0 + tiles_per_split, tiles_total); #if !FA2_FILL // The mailbox dwords are the V plane's last two, and the tile loop's // first V fill (Transition 2) writes them. Without this barrier a wave @@ -631,9 +641,9 @@ __device__ __forceinline__ void fa2_gqa_body( #pragma unroll for (int s = 0; s < FA2_FILL_HV; ++s) #if FA2_R3 - fa2_fill_v_load(v_cache, 0, seq_len, kv_blk, fa2_r3_vslot(hl, s), + fa2_fill_v_load(v_cache, tile0 * KT, seq_len, kv_blk, fa2_r3_vslot(hl, s), #else - fa2_fill_v_load(v_cache, 0, seq_len, kv_blk, hl + s * FA2_FILL_HT, + fa2_fill_v_load(v_cache, tile0 * KT, seq_len, kv_blk, hl + s * FA2_FILL_HT, #endif rv_s[s], rv_u + s * 4); } @@ -641,15 +651,13 @@ __device__ __forceinline__ void fa2_gqa_body( { uint32_t su; uint32_t uc[8]; - fa2_fill_k_load(k_cache, 0, seq_len, kv_blk, tid, su, uc); + fa2_fill_k_load(k_cache, tile0 * KT, seq_len, kv_blk, tid, su, uc); fa2_fill_k_store(Kdw, tid, su, uc); } __syncthreads(); if (fill_compute) { - for (int tile = 0; ; ++tile) { + for (int tile = tile0; tile < tile1; ++tile) { const int ktile = tile * KT; - if (ktile >= seq_len) - break; const bool do0 = ktile <= gmax; const bool do1 = ktile + 16 <= gmax; float8_t sacc0, sacc1; @@ -681,11 +689,9 @@ __device__ __forceinline__ void fa2_gqa_body( #endif uint32_t rk_s[FA2_FILL_HK]; uint32_t rk_u[FA2_FILL_HK * 8]; - for (int tile = 0; ; ++tile) { + for (int tile = tile0; tile < tile1; ++tile) { const int ktile = tile * KT; - if (ktile >= seq_len) - break; - const bool next = ktile + KT < seq_len; + const bool next = tile + 1 < tile1; if (next) { #pragma unroll for (int t = 0; t < FA2_FILL_HK; ++t) @@ -740,10 +746,8 @@ __device__ __forceinline__ void fa2_gqa_body( #endif #endif // ---- Transitions 2-4: KT tiles ---- - for (int tile = 0; ; ++tile) { + for (int tile = tile0; tile < tile1; ++tile) { const int ktile = tile * KT; - if (ktile >= seq_len) - break; #if HIPFIRE_FA2_KT == 32 // F6: tiles after the first consume the raw bytes prefetched into // registers during the previous tile's compute phase (see below); @@ -751,7 +755,7 @@ __device__ __forceinline__ void fa2_gqa_body( // divergence. tile > 0 implies the prefetch ran: the previous // iteration saw ktile_prev + KT < seq_len, else this iteration // would have broken out above. - const bool have_pf = (tile > 0); + const bool have_pf = (tile > tile0); #else constexpr bool have_pf = false; #endif @@ -966,7 +970,7 @@ __device__ __forceinline__ void fa2_gqa_body( // sched_barrier (and never an asm memory clobber) is used. { const int ktile_next = ktile + KT; - if (ktile_next < seq_len) { + if (tile + 1 < tile1) { #if HIPFIRE_FA2_KMODE == 3 #pragma unroll 1 for (int t = 0; t < KT_KITERS; ++t) { @@ -1378,23 +1382,45 @@ __device__ __forceinline__ void fa2_gqa_body( } #endif - // ---- Transition 4: completion (direct epilogue) ---- + // ---- Transition 4: completion ---- // Ofr[dc][j] at lane (ml,half) = O[query=ml][dim=dc*16+2*j+half]: // the pair covers every output dimension exactly once (even/odd // split), and the pair-local l_val normalizes. Each valid // out[query,head,dim] has exactly one writer. if (compute) { - if (qok_ml) { + if (!PARTIAL) { + if (qok_ml) { + const float inv = l_val > 0.0f ? 1.0f / l_val : 0.0f; + // 32-bit row base (whole O tensor < 10 MiB). + const unsigned bo = + (unsigned)qr_ml * 6144u + (unsigned)h_ml * 256u; +#pragma unroll + for (int dc = 0; dc < 16; ++dc) { + const unsigned bd = bo + (unsigned)(dc * 16 + half); +#pragma unroll + for (int j = 0; j < 8; ++j) + out[bd + (unsigned)(j * 2)] = Ofr[dc][j] * inv; + } + } + } else if (qok_ml) { + const unsigned records = + (unsigned)batch_size * 24u * (unsigned)n_splits; + float* lse = out; + float* oacc = out + records; + const unsigned record = + ((unsigned)qr_ml * 24u + (unsigned)h_ml) + * (unsigned)n_splits + (unsigned)split; const float inv = l_val > 0.0f ? 1.0f / l_val : 0.0f; - // 32-bit row base (whole O tensor < 10 MiB). - const unsigned bo = - (unsigned)qr_ml * 6144u + (unsigned)h_ml * 256u; + if (half == 0) + lse[record] = l_val > 0.0f + ? m_old + __logf(l_val) + : -INFINITY; + float* rec = oacc + record * 256u; #pragma unroll for (int dc = 0; dc < 16; ++dc) { - const unsigned bd = bo + (unsigned)(dc * 16 + half); #pragma unroll for (int j = 0; j < 8; ++j) - out[bd + (unsigned)(j * 2)] = Ofr[dc][j] * inv; + rec[(unsigned)(dc * 16 + j * 2 + half)] = Ofr[dc][j] * inv; } } } @@ -1514,9 +1540,80 @@ extern "C" __global__ __launch_bounds__(HIPFIRE_FA2_Q16 ? 256 : 128, 1) void HIP #endif if (q_base >= batch_size) return; - fa2_gqa_body(q16, k_cache, v_cache, out, positions, batch_size, - scale_attn, kv_h, q_base); + fa2_gqa_body(q16, k_cache, v_cache, out, positions, batch_size, + scale_attn, kv_h, q_base, 0, 1); +} + +#if HIPFIRE_FA2_KMODE == 0 && FA2_R3 +extern "C" __global__ __launch_bounds__(256, 1) void attention_q8_0_fa2_gqa_partial_gfx1100( + const _Float16* __restrict__ q16, + const unsigned char* __restrict__ k_cache, + const unsigned char* __restrict__ v_cache, + float* __restrict__ partials, + const int* __restrict__ positions, + int n_heads, + int n_kv_heads, + int head_dim, + int batch_size, + float scale_attn, + int n_splits) +{ + if (n_heads != 24 || n_kv_heads != 4 || head_dim != 256) + return; + if (n_splits < 1 || n_splits > 8) + return; + const unsigned lin = blockIdx.x + gridDim.x * blockIdx.y; + const int kv_h = (int)(lin & 3u); + const int q_base = (int)(gridDim.x - 1u - (lin >> 2)) * 16; + if (q_base >= batch_size) + return; + const int split = blockIdx.z; + if (split >= n_splits) + return; + fa2_gqa_body(q16, k_cache, v_cache, partials, positions, batch_size, + scale_attn, kv_h, q_base, split, n_splits); +} + +extern "C" __global__ __launch_bounds__(256) void attention_q8_0_fa2_gqa_merge_gfx1100( + const float* __restrict__ partials, + float* __restrict__ out, + int batch_size, + int n_heads, + int head_dim, + int n_splits) +{ + if (n_heads != 24 || head_dim != 256) + return; + if (n_splits < 1 || n_splits > 8) + return; + const int tid = (int)threadIdx.x; + const int wave = tid >> 5; + const int lane = tid & 31; + const int rec = blockIdx.x * 8 + wave; + if (rec >= batch_size * n_heads) + return; + const long long records = + (long long)batch_size * n_heads * n_splits; + const float* lse = partials; + const float* oacc = partials + records; + float m = -INFINITY; + for (int s = 0; s < n_splits; ++s) + m = fmaxf(m, lse[(long long)rec * n_splits + s]); + for (int d = lane; d < head_dim; d += 32) { + float num = 0.0f; + float den = 0.0f; + for (int s = 0; s < n_splits; ++s) { + const long long record = (long long)rec * n_splits + s; + const float split_lse = lse[record]; + const float w = isfinite(split_lse) ? __expf(split_lse - m) : 0.0f; + den += w; + num += w * oacc[record * head_dim + d]; + } + out[(long long)rec * head_dim + d] = + den > 0.0f ? num / den : 0.0f; + } } +#endif // ---- KMODE=3 entry: fwht3 K, pre-rotated + pre-converted f16 Q ---- // K is stored FWHT-rotated, so Q must be rotated before QK (orthogonal @@ -1547,7 +1644,7 @@ extern "C" __global__ __launch_bounds__(128, 1) void attention_q8_0_fa2_gqa_fwht const int q_base = blockIdx.x * 8; if (q_base >= batch_size) return; - fa2_gqa_body(q16, k_cache, v_cache, out, positions, batch_size, - scale_attn, kv_h, q_base); + fa2_gqa_body(q16, k_cache, v_cache, out, positions, batch_size, + scale_attn, kv_h, q_base, 0, 1); } #endif diff --git a/kernels/src/gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage.hip b/kernels/src/gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage.hip index 533b30f093..e5d67480e4 100644 --- a/kernels/src/gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage.hip +++ b/kernels/src/gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage.hip @@ -2,7 +2,9 @@ // Copyright (c) 2026 Kaden Schutt // hipfire — see LICENSE and NOTICE in the project root. // -// Default for eligible exact-gfx1100 eager HIP launches; capture keeps base. +// Default for eligible exact-gfx1100 eager HIP launches. The fixed-grid +// split-verifier route may also retain this symbol under HipGraph/PM4 capture +// once its resource and kernarg contracts are registered. // // gfx1100 (RDNA3) RAW-slab LDS-stage gate+up for MQ4G256V2 (qt=44), DFlash // N<=16 tier. Composes: From 4f07ac0bd26d480d0279dc1363bb3da444204bd1 Mon Sep 17 00:00:00 2001 From: HUSRCF Date: Tue, 6 Oct 2026 17:03:21 +0800 Subject: [PATCH 2/2] fix(gfx1100): default split verify inside retained envelope Make the exact dense Q8 H24/KV4/HD256 long-context route default-on while retaining the parent VerifyAttn and route-specific opt-outs. Align eager, HipGraph, and Redline admission; package every split symbol in the compiler-free registry; and fail closed when kernels, fixed Q16 scratch, or the caller's S8 partial workspace are not ready.\n\nThe PM4 loader now includes the target flash-partials capacity in its fixed-B16 admission, so a batched fallback cannot be mislabeled as a retained split tape. --- ...est_mq_f16_projection_producers_gfx1100.rs | 1 + crates/hipfire-arch-qwen35/map.md | 14 +- crates/hipfire-arch-qwen35/src/dflash_spec.rs | 28 +- .../hipfire-arch-qwen35/src/qwen35/config.rs | 16 + .../hipfire-arch-qwen35/src/qwen35/prefill.rs | 173 +- crates/hipfire-arch-qwen35/src/speculative.rs | 20 +- crates/hipfire-config/map.md | 4 +- crates/hipfire-config/src/lib.rs | 43 +- crates/hipfire-generate/map.md | 4 +- crates/hipfire-generate/src/redline.rs | 31 +- crates/hipfire-loader/map.md | 4 +- crates/hipfire-loader/src/lib.rs | 21 +- crates/rdna-compute/map.md | 10 +- crates/rdna-compute/src/attention.rs | 179 +- crates/rdna-compute/src/feature_flags.rs | 195 +- crates/rdna-compute/src/kernel_registry.rs | 3940 +++++++++++++---- crates/rdna-compute/src/mq_f16_producers.rs | 4 +- docs/CONFIG.md | 3 +- docs/env-vars.md | 20 +- 19 files changed, 3820 insertions(+), 890 deletions(-) diff --git a/crates/hipfire-arch-qwen35/examples/test_mq_f16_projection_producers_gfx1100.rs b/crates/hipfire-arch-qwen35/examples/test_mq_f16_projection_producers_gfx1100.rs index 60dac92339..04f89be96d 100644 --- a/crates/hipfire-arch-qwen35/examples/test_mq_f16_projection_producers_gfx1100.rs +++ b/crates/hipfire-arch-qwen35/examples/test_mq_f16_projection_producers_gfx1100.rs @@ -527,6 +527,7 @@ fn main() { up_m, k, n, + false, ) .expect("new gate_up gemm"); gpu.hip.device_synchronize().expect("sync gate_up"); diff --git a/crates/hipfire-arch-qwen35/map.md b/crates/hipfire-arch-qwen35/map.md index 22c858bc62..887fd3ed66 100644 --- a/crates/hipfire-arch-qwen35/map.md +++ b/crates/hipfire-arch-qwen35/map.md @@ -28,7 +28,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/carrier.rs`](src/carrier.rs) | 854 | 4 | 7 | | [`src/checkpoint.rs`](src/checkpoint.rs) | 108 | 3 | 0 | | [`src/dflash_slot.rs`](src/dflash_slot.rs) | 783 | 12 | 0 | -| [`src/dflash_spec.rs`](src/dflash_spec.rs) | 2,080 | 19 | 17 | +| [`src/dflash_spec.rs`](src/dflash_spec.rs) | 2,088 | 19 | 17 | | [`src/dflash_verify_pm4.rs`](src/dflash_verify_pm4.rs) | 780 | 37 | 10 | | [`src/forward_slots.rs`](src/forward_slots.rs) | 4,346 | 24 | 6 | | [`src/grammar_config.rs`](src/grammar_config.rs) | 143 | 2 | 4 | @@ -42,18 +42,18 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/mtp_speculator.rs`](src/mtp_speculator.rs) | 847 | 4 | 0 | | [`src/paro_moe.rs`](src/paro_moe.rs) | 575 | 0 | 0 | | [`src/qwen35/batch.rs`](src/qwen35/batch.rs) | 2,132 | 17 | 2 | -| [`src/qwen35/config.rs`](src/qwen35/config.rs) | 1,947 | 43 | 25 | +| [`src/qwen35/config.rs`](src/qwen35/config.rs) | 1,963 | 44 | 25 | | [`src/qwen35/ep_batch.rs`](src/qwen35/ep_batch.rs) | 5,641 | 23 | 8 | | [`src/qwen35/forward.rs`](src/qwen35/forward.rs) | 7,697 | 33 | 12 | | [`src/qwen35/load.rs`](src/qwen35/load.rs) | 7,635 | 16 | 14 | | [`src/qwen35/oracle.rs`](src/qwen35/oracle.rs) | 1,243 | 20 | 0 | -| [`src/qwen35/prefill.rs`](src/qwen35/prefill.rs) | 17,348 | 18 | 71 | +| [`src/qwen35/prefill.rs`](src/qwen35/prefill.rs) | 17,495 | 20 | 74 | | [`src/qwen35/weights.rs`](src/qwen35/weights.rs) | 2,991 | 43 | 11 | | [`src/qwen35.rs`](src/qwen35.rs) | 69 | 8 | 0 | | [`src/serve_engine.rs`](src/serve_engine.rs) | 7,434 | 12 | 33 | | [`src/spec_emit.rs`](src/spec_emit.rs) | 1,042 | 5 | 14 | | [`src/spec_impl.rs`](src/spec_impl.rs) | 643 | 1 | 0 | -| [`src/speculative.rs`](src/speculative.rs) | 10,002 | 88 | 14 | +| [`src/speculative.rs`](src/speculative.rs) | 10,016 | 88 | 14 | ### Public API surface @@ -76,12 +76,12 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside - [`src/mtp_speculator.rs`](src/mtp_speculator.rs): `Qwen35MtpDrafter`, `new`, `mtp_live_state`, `build_qwen35_mtp_speculator` - [`src/paro_moe.rs`](src/paro_moe.rs): — - [`src/qwen35/batch.rs`](src/qwen35/batch.rs): `PrefillBatchScratch`, `new`, `new_opt`, `new_opt_lean`, `free_gpu`, `Qwen35DecodeBatchState`, `reset`, `reset_lane`, `prefill_lane`, `sample`, `sample_product`, `sample_lane`, +5 more -- [`src/qwen35/config.rs`](src/qwen35/config.rs): `LayerType`, `MaskEmbedOverride`, `DflashFusionCtx`, `TreeVerifyCtx`, `Qwen35Config`, `DenseTpRankLayout`, `dense_tp_rank_layouts`, `validate_dense_tp`, `local_dense_tp_config`, `Qwen35EpReduce`, `Qwen35BatchParallelism`, `Qwen35EpBatchReceipt`, +31 more +- [`src/qwen35/config.rs`](src/qwen35/config.rs): `LayerType`, `MaskEmbedOverride`, `DflashFusionCtx`, `fn`, `TreeVerifyCtx`, `Qwen35Config`, `DenseTpRankLayout`, `dense_tp_rank_layouts`, `validate_dense_tp`, `local_dense_tp_config`, `Qwen35EpReduce`, `Qwen35BatchParallelism`, +32 more - [`src/qwen35/ep_batch.rs`](src/qwen35/ep_batch.rs): `validate_ep_batch_compatibility`, `Qwen35DecodeBatchEpState`, `max_batch`, `sequential_prefill_scratch_mut`, `peer_reduce_lease`, `lane_capacity`, `epoch`, `poison_mask`, `lane_state`, `download_rank_outputs`, `new`, `reset_all`, +11 more - [`src/qwen35/forward.rs`](src/qwen35/forward.rs): `dump_expert_stats`, `forward`, `Qwen35Scratch`, `new`, `new_with_kv_max`, `free_gpu`, `Qwen35ScratchSet`, `new_with_kv_max_multi`, `free_gpu_multi`, `forward_scratch`, `prepare_scratch_inputs`, `forward_scratch_with_hidden`, +21 more - [`src/qwen35/load.rs`](src/qwen35/load.rs): `hipfire_runtime`, `qwen35_tensor_name_candidates`, `load_weight_tensor`, `load_weight_tensor_host`, `report_cpu_exec_coverage`, `load_weights`, `load_weights_with_fault`, `HfqSource`, `new`, `ParoSource`, `preflight_weights_dense_tp`, `load_weights_dense_tp_rank`, +4 more - [`src/qwen35/oracle.rs`](src/qwen35/oracle.rs): `SCHEMA_VERSION`, `SCHEMA_FIELDS`, `ArrayRef`, `Snapshot`, `Observation`, `begin`, `set_sequence`, `set_decode_position`, `set_prefill_start`, `decode_position`, `prefill_start`, `decode_before`, +8 more -- [`src/qwen35/prefill.rs`](src/qwen35/prefill.rs): `PREFILL_MAX_BATCH`, `prefill_max_batch`, `prefill_max_batch_tp`, `prefill_max_batch_ep`, `vmm_kv_token_bytes`, `minimum_prefill_reservation_bytes`, `ordinary_prefill_chunk_limit`, `ordinary_serve_prefill_chunk_len`, `upload_prefill_batch_inputs`, `forward_prefill_batch_single_chunk_captured`, `forward_prefill_batch_single_chunk_captured_opts`, `forward_prefill_batch`, +6 more +- [`src/qwen35/prefill.rs`](src/qwen35/prefill.rs): `PREFILL_MAX_BATCH`, `prefill_max_batch`, `prefill_max_batch_tp`, `prefill_max_batch_ep`, `vmm_kv_token_bytes`, `minimum_prefill_reservation_bytes`, `ordinary_prefill_chunk_limit`, `ordinary_serve_prefill_chunk_len`, `upload_prefill_batch_inputs`, `forward_prefill_batch_single_chunk_captured`, `forward_prefill_batch_single_chunk_captured_opts`, `forward_prefill_batch`, +8 more - [`src/qwen35/weights.rs`](src/qwen35/weights.rs): `DeltaNetLayerWeights`, `FullAttnLayerWeights`, `ExpertWeights`, `mixed_expert_tag`, `SharedExpertWeights`, `MoeFfnWeights`, `MoeParoSidecars`, `DeltaNetMoeLayerWeights`, `FullAttnMoeLayerWeights`, `LayerWeights`, `free_gpu`, `Qwen35HfqSourceIdentity`, +31 more - [`src/qwen35.rs`](src/qwen35.rs): `batch`, `config`, `ep_batch`, `forward`, `load`, `oracle`, `prefill`, `weights` - [`src/serve_engine.rs`](src/serve_engine.rs): `EngineConfig`, `SlotEngine`, `submit`, `cancel_waiting`, `close`, `reset`, `stats`, `spawn`, `shutdown_engine`, `MTP_RETIRE_WINDOW`, `MTP_RETIRE_MIN_ADVANCE`, `MTP_RETIRE_WINDOWS` @@ -102,6 +102,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 31 modules · 87,577 lines · 557 public items · 283 tests · 18 examples +- 31 modules · 87,762 lines · 560 public items · 286 tests · 18 examples diff --git a/crates/hipfire-arch-qwen35/src/dflash_spec.rs b/crates/hipfire-arch-qwen35/src/dflash_spec.rs index cdec47c388..50878fc2df 100644 --- a/crates/hipfire-arch-qwen35/src/dflash_spec.rs +++ b/crates/hipfire-arch-qwen35/src/dflash_spec.rs @@ -164,6 +164,7 @@ pub fn load_dflash_state( // Retained-PM4 admission facts owned by the loader/target, not the draft. target_weights: &Qwen35Weights, kv_is_q8: bool, + target_flash_partials_numel: usize, single_gpu: bool, // True when adaptive KV is engaged for this load (tier-switching cache). // Must be false for retained-PM4 admission. @@ -466,11 +467,18 @@ pub fn load_dflash_state( .as_ref() .map(|p| p.max_batch) .unwrap_or(0); - let gfx1100_split_verify = gpu.arch == "gfx1100" - && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false) - && target_config.n_heads == 24 - && target_config.n_kv_heads == 4 - && target_config.head_dim == 256; + let gfx1100_split_verify_min_ctx = + qwen35::prefill::gfx1100_split_verify_min_ctx(gpu, target_config).filter(|&threshold| { + qwen35::prefill::gfx1100_split_verify_admitted( + gpu, + target_config, + kv_is_q8, + DFLASH_VERIFY_PM4_BLOCK, + threshold.saturating_add(1), + target_flash_partials_numel, + ) + }); + let gfx1100_split_verify = gfx1100_split_verify_min_ctx.is_some(); // The draft's selector/dynamic-conv shape is deliberately NOT a gate: the // draft forward is outside the tape, so DFlash2 and legacy DFlash yield an // identical target verify body. @@ -502,11 +510,11 @@ pub fn load_dflash_state( " DFlash verify PM4: armed (B={}, exact {})", DFLASH_VERIFY_PM4_BLOCK, gpu.arch ); - if gfx1100_split_verify { - // Keep the faster incumbent below the split verifier's 4K - // crossover, and do not capture a short-context tape that - // would silently omit the new partial + merge launches. - DflashVerifyPm4::armed_after_context(4096) + if let Some(min_ctx) = gfx1100_split_verify_min_ctx { + // Keep the faster incumbent below the resolved crossover, and + // do not capture a short-context tape that would silently omit + // the new partial + merge launches. + DflashVerifyPm4::armed_after_context(min_ctx) } else { DflashVerifyPm4::armed() } diff --git a/crates/hipfire-arch-qwen35/src/qwen35/config.rs b/crates/hipfire-arch-qwen35/src/qwen35/config.rs index f639801c92..a3b7636b20 100644 --- a/crates/hipfire-arch-qwen35/src/qwen35/config.rs +++ b/crates/hipfire-arch-qwen35/src/qwen35/config.rs @@ -103,6 +103,22 @@ pub struct MaskEmbedOverride<'a> { pub enum DflashFusionCtx { Off, ChainVerify, + /// Linear-chain verify whose exact model/KV/context envelope admits the + /// gfx1100 FA2 split-KV route. This only widens capture-safe auxiliary + /// kernels; ordinary eager ChainVerify behavior is otherwise unchanged. + ChainVerifySplit, +} + +impl DflashFusionCtx { + #[inline] + pub const fn is_chain_verify(self) -> bool { + matches!(self, Self::ChainVerify | Self::ChainVerifySplit) + } + + #[inline] + pub const fn split_verify_active(self) -> bool { + matches!(self, Self::ChainVerifySplit) + } } #[derive(Clone, Copy)] diff --git a/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs b/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs index 2307d47431..5650f55d28 100644 --- a/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs +++ b/crates/hipfire-arch-qwen35/src/qwen35/prefill.rs @@ -5542,9 +5542,9 @@ pub(crate) fn batch_chunk_upload_positions( #[inline] fn mq_f16_projection_fast_route(gpu: &Gpu, fusion: DflashFusionCtx, n: usize, dim: usize) -> bool { let recording = gpu.graphs.capture_mode || gpu.replay.is_recording(); - let recording_supported = - !recording || hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false); - matches!(fusion, DflashFusionCtx::ChainVerify) + let recording_supported = !recording + || (fusion.split_verify_active() && gpu.gfx1100_q8_fa2_split_capture_assets_ready(n)); + fusion.is_chain_verify() && gpu.arch_caps.is_gfx1100() && !gpu.flags.mq_f16_projection_off && n >= 1 @@ -6209,7 +6209,7 @@ fn batch_chunk_delta_net_pre_gdn<'a>( // (kill switch, non-sequential batch, non-GQA route, tape absence or // overflow, ineligible shapes/arch) runs the pre-change sequence below // launch-for-launch. - if fusion == DflashFusionCtx::ChainVerify + if fusion.is_chain_verify() && !gpu.flags.gdn_pre_fuse_off && matches!(batch_semantics, BatchSemantics::Sequential) && config.linear_num_key_heads < n_v_heads @@ -6463,7 +6463,7 @@ fn s4_residual_fast( epilogue: &BatchEpilogue<'_>, n: usize, ) -> bool { - fusion == DflashFusionCtx::ChainVerify + fusion.is_chain_verify() && !gpu.flags.mq_f16_residual_off && gpu.arch_caps.supports_dflash_f16_residual_fusions() && w_dtype == DType::MQ4G256V2 @@ -7329,6 +7329,7 @@ fn batch_chunk_delta_net_ffn_gate_up( layer.w_up.m, layer.w_gate.k, n, + fusion.split_verify_active() && gpu.gfx1100_q8_fa2_split_capture_assets_ready(n), ) .map(|()| FfnGateOutput::Separate); } @@ -8530,7 +8531,7 @@ fn batch_chunk_full_attn_prepare( let fa_prep_shape_ok = matches!((config.n_heads, config.n_kv_heads), (16, 2) | (24, 4)) && config.head_dim == 256 && fa_prep_n_rot == 64; - let fa_prep_fused_ok = (fusion == DflashFusionCtx::ChainVerify + let fa_prep_fused_ok = (fusion.is_chain_verify() || (fusion == DflashFusionCtx::Off && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA_PREP", true))) && gpu.arch_caps.is_gfx1100() @@ -8953,6 +8954,87 @@ fn fa_pertoken_min_ctx_for(arch: &str, explicit: Option) -> Option } } +/// Resolved product envelope for the exact-gfx1100 FA2 split-KV verifier. +/// +/// This is deliberately narrower than the generic R4/R8 multi-row route: +/// dense Qwen H24/KV4/HD256, both the stable VerifyAttn parent switch and the +/// split-specific kill switch enabled, and a live-context crossover. Keeping +/// the threshold here also makes retained-PM4 defer capture at the same point +/// as eager/HipGraph dispatch, including an explicit +/// `HIPFIRE_FA_PERTOKEN_MIN_CTX=0` opt-out. +fn gfx1100_split_verify_model_admitted( + arch: &str, + verify_attn: bool, + split_verify: bool, + num_experts: usize, + n_heads: usize, + n_kv_heads: usize, + head_dim: usize, +) -> bool { + arch == "gfx1100" + && verify_attn + && split_verify + && num_experts == 0 + && n_heads == 24 + && n_kv_heads == 4 + && head_dim == 256 +} + +#[inline] +fn gfx1100_split_verify_threshold(resolved: Option) -> Option { + // The lower launcher independently enforces the measured >4K envelope. + // Clamp smaller developer overrides so HipGraph/PM4 never capture a + // pre-crossover batched tape while auxiliary split-only kernels are armed. + resolved.map(|threshold| threshold.max(4_096)) +} + +#[doc(hidden)] +pub fn gfx1100_split_verify_min_ctx(gpu: &Gpu, config: &Qwen35Config) -> Option { + gfx1100_split_verify_model_admitted( + gpu.arch_caps.arch(), + gpu.flags.verify_attn, + gpu.flags.gfx1100_fa2_split_verify, + config.num_experts, + config.n_heads, + config.n_kv_heads, + config.head_dim, + ) + .then(|| gfx1100_split_verify_threshold(fa_pertoken_min_ctx(gpu.arch_caps.arch()))) + .flatten() +} + +#[inline] +fn gfx1100_split_verify_partials_admitted( + batch_size: usize, + n_heads: usize, + head_dim: usize, + flash_partials_numel: usize, +) -> bool { + rdna_compute::attention::gfx1100_q8_fa2_split_partials_len(batch_size, n_heads, head_dim) + .is_some_and(|need| flash_partials_numel >= need) +} + +#[doc(hidden)] +pub fn gfx1100_split_verify_admitted( + gpu: &Gpu, + config: &Qwen35Config, + quant_q8: bool, + batch_size: usize, + logical_ctx_len: usize, + flash_partials_numel: usize, +) -> bool { + quant_q8 + && (4..=32).contains(&batch_size) + && gfx1100_split_verify_partials_admitted( + batch_size, + config.n_heads, + config.head_dim, + flash_partials_numel, + ) + && gfx1100_split_verify_min_ctx(gpu, config) + .is_some_and(|threshold| logical_ctx_len > threshold) +} + /// Opt-in gfx1201 split-KV packed-Q8 verifier: `HIPFIRE_GFX12_FA2_SPLIT_VERIFY` /// (`1`/`on`/`true`) enables it, `HIPFIRE_GFX12_FA2_SPLIT_COUNT` picks the /// split count (`2`, `4` or `8`; default `8`; anything else fails closed to @@ -9281,6 +9363,14 @@ fn batch_chunk_fa_attend( start_pos + n, n, &s.flash_partials, + gfx1100_split_verify_admitted( + gpu, + config, + kv_cache.quant_q8, + n, + start_pos + n, + s.flash_partials.numel(), + ), )? { return Ok(()); } @@ -9472,7 +9562,7 @@ pub(crate) fn batch_chunk_full_attn_attn( let q_dim = config.n_heads * config.head_dim; let gfx12_fa_prep = gpu.arch == "gfx1201" && gpu.flags.gfx12_fa_prep_fused - && fusion != DflashFusionCtx::ChainVerify + && !fusion.is_chain_verify() && !gpu.flags.rope_interleaved_legacy && !hipfire_runtime::triattn::tap_enabled() && config.head_dim == 256 @@ -9665,6 +9755,7 @@ fn batch_chunk_full_attn_ffn_gate_up( layer.w_up.m, layer.w_gate.k, n, + fusion.split_verify_active() && gpu.gfx1100_q8_fa2_split_capture_assets_ready(n), ) .map(|()| FfnGateOutput::Separate); } @@ -12443,8 +12534,6 @@ fn forward_prefill_chunk_pair( fa_pertoken_min_ctx(gpu.arch_caps.arch()), gpu.graphs.capture_mode, gpu.replay.is_recording(), - gpu.arch_caps.is_gfx1100() - && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false), ); let multirow_admitted = |max_ctx: usize| { q8_multirow_attn_admitted( @@ -12458,7 +12547,14 @@ fn forward_prefill_chunk_pair( false, multirow_common.4, multirow_common.5, - multirow_common.6, + gfx1100_split_verify_admitted( + gpu, + config, + kv_cache.quant_q8, + n, + max_ctx, + s.flash_partials.numel(), + ), ) }; let multirow_c = multirow_admitted(max_ctx_c); @@ -13131,7 +13227,8 @@ pub(crate) fn forward_batch_chunk_impl( // GEMMs stay batched either way — only the attend step switches to the // multi-row tile. The incumbent tile grid is sized from the live logical // context and therefore stays out of capture. The gfx1100 split-KV route - // uses fixed S=8 geometry and is admitted only behind its explicit opt-in; + // uses fixed S=8 geometry and is admitted only inside its exact default-on + // gfx1100 envelope (with parent and route-specific opt-outs); // the launcher independently rejects every other captured shape. let fa_attn_multirow = q8_multirow_attn_admitted( gpu.arch_caps.arch(), @@ -13144,8 +13241,14 @@ pub(crate) fn forward_batch_chunk_impl( batch_semantics.is_independent(), gpu.graphs.capture_mode, gpu.replay.is_recording(), - gpu.arch_caps.is_gfx1100() - && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false), + gfx1100_split_verify_admitted( + gpu, + config, + kv_cache.quant_q8, + n, + start_pos + n, + s.flash_partials.numel(), + ), ); let logical_max_ctx = match batch_semantics { BatchSemantics::Sequential => start_pos + n, @@ -14341,6 +14444,50 @@ mod tests { )); } + #[test] + fn gfx1100_split_verify_model_gate_is_dense_exact_shape_with_two_switches() { + let admitted = |arch, verify_attn, split_verify, num_experts, n_heads, n_kv, hd| { + gfx1100_split_verify_model_admitted( + arch, + verify_attn, + split_verify, + num_experts, + n_heads, + n_kv, + hd, + ) + }; + assert!(admitted("gfx1100", true, true, 0, 24, 4, 256)); + assert!(!admitted("gfx1100", false, true, 0, 24, 4, 256)); + assert!(!admitted("gfx1100", true, false, 0, 24, 4, 256)); + assert!(!admitted("gfx1100", true, true, 8, 24, 4, 256)); + assert!(!admitted("gfx1201", true, true, 0, 24, 4, 256)); + assert!(!admitted("gfx1100", true, true, 0, 32, 4, 256)); + assert!(!admitted("gfx1100", true, true, 0, 24, 8, 256)); + assert!(!admitted("gfx1100", true, true, 0, 24, 4, 128)); + } + + #[test] + fn gfx1100_split_verify_threshold_never_precedes_measured_crossover() { + assert_eq!(gfx1100_split_verify_threshold(None), None); + assert_eq!(gfx1100_split_verify_threshold(Some(1)), Some(4_096)); + assert_eq!(gfx1100_split_verify_threshold(Some(4_096)), Some(4_096)); + assert_eq!(gfx1100_split_verify_threshold(Some(8_192)), Some(8_192)); + } + + #[test] + fn gfx1100_split_verify_rejects_undersized_partials_workspace() { + let need = rdna_compute::attention::gfx1100_q8_fa2_split_partials_len(16, 24, 256) + .expect("fixed production shape"); + assert!(gfx1100_split_verify_partials_admitted(16, 24, 256, need,)); + assert!(!gfx1100_split_verify_partials_admitted( + 16, + 24, + 256, + need - 1, + )); + } + #[test] fn fa_pertoken_min_ctx_is_opt_in_on_gfx1151_only() { for arch in ["gfx1100", "gfx1201"] { diff --git a/crates/hipfire-arch-qwen35/src/speculative.rs b/crates/hipfire-arch-qwen35/src/speculative.rs index 2117772fd4..d58a6a4e8c 100644 --- a/crates/hipfire-arch-qwen35/src/speculative.rs +++ b/crates/hipfire-arch-qwen35/src/speculative.rs @@ -3626,10 +3626,24 @@ fn verify_dflash_block_inner( // shapes. sub_offset returns a non-owning view; do NOT free these. let final_hidden = verify_scratch.final_hidden.sub_offset(0, b * dim); let tree_verify_present = tree_verify.is_some(); - // Launch-fusion prescaffold: frozen AR/verify discriminator. Linear chain - // verify (`tree_verify` is `None`) arms `ChainVerify`; tree verify stays `Off`. + // Launch-fusion prescaffold: frozen AR/verify discriminator. The split + // variant is selected only when this exact live window also satisfies the + // dense gfx1100 Q8 attention route. That prevents the split flag from + // widening capture-time F16/LDS auxiliary kernels below the crossover or + // on unsupported model shapes. let fusion = if tree_verify.is_none() { - qwen35::DflashFusionCtx::ChainVerify + if qwen35::prefill::gfx1100_split_verify_admitted( + gpu, + &target.config, + target.kv_cache.quant_q8, + b, + required_tokens, + target.scratch.flash_partials.numel(), + ) { + qwen35::DflashFusionCtx::ChainVerifySplit + } else { + qwen35::DflashFusionCtx::ChainVerify + } } else { qwen35::DflashFusionCtx::Off }; diff --git a/crates/hipfire-config/map.md b/crates/hipfire-config/map.md index 0e07445ca3..9cf3a190d0 100644 --- a/crates/hipfire-config/map.md +++ b/crates/hipfire-config/map.md @@ -24,7 +24,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside |---|---:|---:|---:| | [`src/bin/hipfire-rocm-resolve.rs`](src/bin/hipfire-rocm-resolve.rs) | 105 | 0 | 0 | | [`src/devices.rs`](src/devices.rs) | 1,439 | 46 | 13 | -| [`src/lib.rs`](src/lib.rs) | 6,694 | 101 | 57 | +| [`src/lib.rs`](src/lib.rs) | 6,727 | 101 | 57 | | [`src/rocm.rs`](src/rocm.rs) | 2,460 | 39 | 37 | ### Public API surface @@ -47,6 +47,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 4 modules · 10,698 lines · 186 public items · 107 tests · 0 examples +- 4 modules · 10,731 lines · 186 public items · 107 tests · 0 examples diff --git a/crates/hipfire-config/src/lib.rs b/crates/hipfire-config/src/lib.rs index 69e6ea56a7..f97402c526 100644 --- a/crates/hipfire-config/src/lib.rs +++ b/crates/hipfire-config/src/lib.rs @@ -524,8 +524,23 @@ const KV_V_NAMES: &[&str] = &["", "q8", "lloyd2", "lloyd3", "lloyd4"]; // (on Qwen it names the same legacy K as `--kv-k legacy-asym3`). const KV_MODES: &[&str] = &[ // lifecycle: deprecated since 0.4.0, removal 0.5.0 — Givens asym KV and the asymN/turboN aliases are superseded by fwht3 (asymN / turbo*) - "auto", "f32", "f16", "bf16", "q8", "asym4", "asym3", "asym2", "fwht4", "fwht3", "fwht2", - "turbo", "turbo4", "turbo3", "turbo2", "fp8", "legacy-asym3", + "auto", + "f32", + "f16", + "bf16", + "q8", + "asym4", + "asym3", + "asym2", + "fwht4", + "fwht3", + "fwht2", + "turbo", + "turbo4", + "turbo3", + "turbo2", + "fp8", + "legacy-asym3", ]; const AUTO_ON_OFF: &[&str] = &["auto", "on", "off"]; /// VL image decode path: `cpu` (default) / `vcn` / `auto` (VCN when probed). @@ -2557,6 +2572,15 @@ pub static FIELDS: &[ConfigField] = &[ "HIPFIRE_VERIFY_ATTN", "Run speculative-verify attention (1..=32 rows, non-tree) through VerifyAttn: on gfx1201 and gfx1100 the GQA-shared split-K twin of the batched flash tile + reduce (on gfx1100 also of the multi-row R4/R8 Q8 tile), on gfx1151 the context-parallel twin of the single-slot WMMA flash prefill (default on exact gfx1201, gfx1100 and gfx1151; byte-identical output; set to false or HIPFIRE_VERIFY_ATTN=0 to opt out to attention_flash_*_tile_batched / attention_flash_q8_0_rows{4,8}_d8 / attention_q8_0_flash_prefill_wmma)." ), + process_bool_field!( + "kernel.gfx1100_fa2_split_verify", + "gfx1100_fa2_split_verify", + Kernel, + true, + false, + "HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", + "Enable the exact-gfx1100 Q8 FA2 split-KV VerifyAttn route (default on exact gfx1100 inside the dense H24/KV4/HD256 long-context envelope; false restores the established VerifyAttn/R4-R8/batched fallback). The parent kernel.verify_attn switch takes precedence." + ), process_bool_field!( "kernel.gfx12_fa_prep_fused", "gfx12_fa_prep_fused", @@ -3793,7 +3817,10 @@ fn legacy_table(config: &ProcessConfig) -> HashMap { let Some(name) = developer_env_for_key(key) else { continue; }; - if FIELDS.iter().any(|schema| schema.env_compat == Some(name.as_str())) { + if FIELDS + .iter() + .any(|schema| schema.env_compat == Some(name.as_str())) + { continue; } if let Some(value) = render_compat_value(value) { @@ -5334,7 +5361,10 @@ mod tests { values, }; let table = legacy_table(&config); - let mut names: Vec<&str> = FIELDS.iter().filter_map(|schema| schema.env_compat).collect(); + let mut names: Vec<&str> = FIELDS + .iter() + .filter_map(|schema| schema.env_compat) + .collect(); names.extend([ "HIPFIRE_DSPARK_Q8_WMMA", "HIPFIRE_NGRAM_WINDOW", @@ -5909,7 +5939,10 @@ mod tests { .unwrap(); let process = ProcessConfig::from_resolved(&resolved).unwrap(); - assert_eq!(process.legacy_value("HIPFIRE_DETERMINISTIC").as_deref(), Some("1")); + assert_eq!( + process.legacy_value("HIPFIRE_DETERMINISTIC").as_deref(), + Some("1") + ); assert_eq!( process.legacy_value("HIPFIRE_FLASH_ATTN_CK_LIB").as_deref(), Some("/opt/hipfire/ck.so") diff --git a/crates/hipfire-generate/map.md b/crates/hipfire-generate/map.md index 426558c258..da08b33f67 100644 --- a/crates/hipfire-generate/map.md +++ b/crates/hipfire-generate/map.md @@ -29,7 +29,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/img.rs`](src/img.rs) | 346 | 1 | 0 | | [`src/lib.rs`](src/lib.rs) | 61 | 8 | 0 | | [`src/qwen.rs`](src/qwen.rs) | 8,196 | 70 | 6 | -| [`src/redline.rs`](src/redline.rs) | 6,182 | 52 | 2 | +| [`src/redline.rs`](src/redline.rs) | 6,193 | 52 | 2 | | [`src/vision.rs`](src/vision.rs) | 3,660 | 9 | 8 | ### Public API surface @@ -57,6 +57,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 9 modules · 41,870 lines · 396 public items · 330 tests · 0 examples +- 9 modules · 41,881 lines · 396 public items · 330 tests · 0 examples diff --git a/crates/hipfire-generate/src/redline.rs b/crates/hipfire-generate/src/redline.rs index b3e1528c26..8811dd8160 100644 --- a/crates/hipfire-generate/src/redline.rs +++ b/crates/hipfire-generate/src/redline.rs @@ -3085,12 +3085,27 @@ pub fn redline_shadow_dflash_verify_pm4( return Err("DFlash shadow iterations must be non-zero".into()); } let count = iterations.max(12); - // The gfx1100 split-KV verifier is admitted only beyond its 4K crossover. - // Start the shadow at the 8K boundary so capture/record/PM4 evidence names - // the new partial + merge route rather than a short-context fallback tape. - let split_gfx1100 = gpu.arch == "gfx1100" - && hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false); - let base: usize = if split_gfx1100 { 8176 } else { 112 }; + let frame_checkpoint = rdna_compute::norm::gdn_requant_frame_checkpoint(); + let mut guard = Qwen35SlotGuard::take(&mut loaded.state, &loaded.model_path)?; + let slot = guard.model_slot()?; + // The gfx1100 split-KV verifier is admitted only beyond its resolved + // crossover. Start the shadow at no less than the 8K boundary so + // capture/record/PM4 evidence names the new partial + merge route rather + // than a short-context fallback tape. + let split_min_ctx = + qwen35::prefill::gfx1100_split_verify_min_ctx(gpu, &slot.config).filter(|&threshold| { + qwen35::prefill::gfx1100_split_verify_admitted( + gpu, + &slot.config, + slot.kv_cache.quant_q8, + batch, + threshold.saturating_add(1), + slot.scratch.flash_partials.numel(), + ) + }); + let base: usize = split_min_ctx + .map(|threshold| threshold.saturating_add(1).max(8192).saturating_sub(batch)) + .unwrap_or(112); let mut positions: Vec = Vec::with_capacity(count); for i in 0..count { positions.push(base + i * batch); @@ -3105,10 +3120,6 @@ pub fn redline_shadow_dflash_verify_pm4( last + batch )); } - - let frame_checkpoint = rdna_compute::norm::gdn_requant_frame_checkpoint(); - let mut guard = Qwen35SlotGuard::take(&mut loaded.state, &loaded.model_path)?; - let slot = guard.model_slot()?; let hidden_k = slot.config.dim.next_power_of_two(); let max_n = batch + 1; let mut fixtures = RedlineDflashFixtures { diff --git a/crates/hipfire-loader/map.md b/crates/hipfire-loader/map.md index 12bd5fd2da..fdbf3142cf 100644 --- a/crates/hipfire-loader/map.md +++ b/crates/hipfire-loader/map.md @@ -26,7 +26,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/admission.rs`](src/admission.rs) | 2,302 | 17 | 31 | | [`src/batch_staging.rs`](src/batch_staging.rs) | 466 | 3 | 0 | | [`src/carriers.rs`](src/carriers.rs) | 3,892 | 14 | 7 | -| [`src/lib.rs`](src/lib.rs) | 5,835 | 108 | 29 | +| [`src/lib.rs`](src/lib.rs) | 5,846 | 108 | 29 | | [`src/spec_build.rs`](src/spec_build.rs) | 326 | 6 | 0 | ### Public API surface @@ -50,6 +50,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 5 modules · 12,821 lines · 148 public items · 67 tests · 1 examples +- 5 modules · 12,832 lines · 148 public items · 67 tests · 1 examples diff --git a/crates/hipfire-loader/src/lib.rs b/crates/hipfire-loader/src/lib.rs index ffd8a50dd0..c808353179 100644 --- a/crates/hipfire-loader/src/lib.rs +++ b/crates/hipfire-loader/src/lib.rs @@ -1968,7 +1968,10 @@ fn resolve_qwen35_mtp_head( gpu: &mut rdna_compute::Gpu, physical_cap: usize, device: Option<&str>, -) -> (Option, Vec) { +) -> ( + Option, + Vec, +) { use hipfire_arch_qwen35::mtp_head; let tag = device.map(|d| format!(", {d}")).unwrap_or_default(); let sidecar = sidecar.unwrap_or_else(|| trunk_path.with_extension("mtp")); @@ -2191,6 +2194,7 @@ fn finish_qwen35_load( && !kv.quant_asym4 && matches!(kv.v_mode, llama::VMode::Q8) }, + bundle.scratch.flash_partials.numel(), // finish_qwen35_load is the single-GPU carrier path. true, // Fail-closed: adaptive KV starts FWHT4 and tier-switches at runtime. @@ -4183,7 +4187,9 @@ fn load_model_tp_qwen35_dense( } else { errors.join("; ") }; - return Err(format!("MTP head required (mtp=on) but not loaded: {reason}")); + return Err(format!( + "MTP head required (mtp=on) but not loaded: {reason}" + )); } head } else { @@ -4776,9 +4782,14 @@ mod ep_admission_tests { let before = active.request(); let mut effects = LoadEffects::default(); let refusal = attempt_candidate_swap( - &candidate, 1, admission::KvBackendRequest::Explicit(hipfire_runtime::kv_backend::KvBackend::Vmm), - "gfx1100", &mut active, &mut effects, - ).unwrap_err(); + &candidate, + 1, + admission::KvBackendRequest::Explicit(hipfire_runtime::kv_backend::KvBackend::Vmm), + "gfx1100", + &mut active, + &mut effects, + ) + .unwrap_err(); assert!(refusal.contains("vmm") && refusal.contains("unsupported")); assert_eq!(effects, LoadEffects::default()); assert_eq!(active.request(), before); diff --git a/crates/rdna-compute/map.md b/crates/rdna-compute/map.md index 9b68b4262a..89708c9188 100644 --- a/crates/rdna-compute/map.md +++ b/crates/rdna-compute/map.md @@ -24,7 +24,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | File | Lines | Public items | Tests | |---|---:|---:|---:| | [`src/arch_caps.rs`](src/arch_caps.rs) | 780 | 60 | 21 | -| [`src/attention.rs`](src/attention.rs) | 23,471 | 288 | 20 | +| [`src/attention.rs`](src/attention.rs) | 23,598 | 291 | 23 | | [`src/bin/hipfire-kernel-hash.rs`](src/bin/hipfire-kernel-hash.rs) | 142 | 0 | 0 | | [`src/bin/hipfire-kernel-manifest.rs`](src/bin/hipfire-kernel-manifest.rs) | 119 | 0 | 0 | | [`src/bin/hipfire-kernel-pack.rs`](src/bin/hipfire-kernel-pack.rs) | 101 | 0 | 0 | @@ -39,7 +39,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/dflash_state_copy.rs`](src/dflash_state_copy.rs) | 171 | 10 | 0 | | [`src/dispatch.rs`](src/dispatch.rs) | 7,739 | 158 | 31 | | [`src/embedding.rs`](src/embedding.rs) | 600 | 13 | 0 | -| [`src/feature_flags.rs`](src/feature_flags.rs) | 2,067 | 33 | 21 | +| [`src/feature_flags.rs`](src/feature_flags.rs) | 2,188 | 33 | 22 | | [`src/flash_attn_ck.rs`](src/flash_attn_ck.rs) | 1,775 | 26 | 15 | | [`src/flux_fused.rs`](src/flux_fused.rs) | 581 | 3 | 3 | | [`src/gap_timing.rs`](src/gap_timing.rs) | 147 | 7 | 0 | @@ -51,7 +51,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside | [`src/grouped_ops.rs`](src/grouped_ops.rs) | 1,541 | 15 | 6 | | [`src/hc_row_fold.rs`](src/hc_row_fold.rs) | 170 | 4 | 0 | | [`src/kernel_pack.rs`](src/kernel_pack.rs) | 332 | 12 | 1 | -| [`src/kernel_registry.rs`](src/kernel_registry.rs) | 1,321 | 30 | 7 | +| [`src/kernel_registry.rs`](src/kernel_registry.rs) | 3,761 | 30 | 8 | | [`src/kernels.rs`](src/kernels.rs) | 10,877 | 1575 | 38 | | [`src/kv_slots.rs`](src/kv_slots.rs) | 526 | 10 | 12 | | [`src/lib.rs`](src/lib.rs) | 120 | 47 | 1 | @@ -83,7 +83,7 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Public API surface - [`src/arch_caps.rs`](src/arch_caps.rs): `ArchCaps`, `note_process_gpu_arch`, `process_gpu_arch`, `new`, `should_use_mmq`, `is_gfx906`, `is_gfx908`, `is_gfx1010`, `is_gfx1011`, `is_gfx1012`, `is_gfx1030`, `is_gfx1031`, +48 more -- [`src/attention.rs`](src/attention.rs): `attention_q8_0_kv_independent_lds_bytes`, `attention_q8_0_kv_independent_max_lane_capacity`, `q8_flash_tile_size`, `GFX12_Q8_FA2_MAX_CTX`, `GFX12_QUERY16_MAX_CTX`, `fp8_e4m3_row_bytes`, `bf16_row_bytes`, `VerifyKv`, `VerifyWmmaPv`, `ALL`, `VerifyWmmaGeometry`, `dspark_stage_kv`, +276 more +- [`src/attention.rs`](src/attention.rs): `attention_q8_0_kv_independent_lds_bytes`, `attention_q8_0_kv_independent_max_lane_capacity`, `q8_flash_tile_size`, `GFX12_Q8_FA2_MAX_CTX`, `GFX12_QUERY16_MAX_CTX`, `fp8_e4m3_row_bytes`, `bf16_row_bytes`, `VerifyKv`, `VerifyWmmaPv`, `ALL`, `VerifyWmmaGeometry`, `gfx1100_q8_fa2_split_partials_len`, +279 more - [`src/bin/hipfire-kernel-hash.rs`](src/bin/hipfire-kernel-hash.rs): — - [`src/bin/hipfire-kernel-manifest.rs`](src/bin/hipfire-kernel-manifest.rs): — - [`src/bin/hipfire-kernel-pack.rs`](src/bin/hipfire-kernel-pack.rs): — @@ -152,6 +152,6 @@ _Generated by `scripts/check-crate-maps.py` from the tree — do not edit inside ### Totals -- 56 modules · 175,796 lines · 4001 public items · 479 tests · 224 examples +- 56 modules · 178,484 lines · 4004 public items · 484 tests · 224 examples diff --git a/crates/rdna-compute/src/attention.rs b/crates/rdna-compute/src/attention.rs index 749cbe326f..dc8b18a583 100644 --- a/crates/rdna-compute/src/attention.rs +++ b/crates/rdna-compute/src/attention.rs @@ -615,6 +615,52 @@ fn gfx1100_q8_fa2_split_admitted( && (4..=32).contains(&batch_size) } +/// F32 elements required by the production gfx1100 S=8 split-KV verifier. +/// +/// Kept shared with the product admission so a reduced +/// `HIPFIRE_FLASH_PARTIALS_BATCH` cannot arm split-only capture fusions while +/// the attention launcher itself must fall back for lack of workspace. +#[doc(hidden)] +pub fn gfx1100_q8_fa2_split_partials_len( + batch_size: usize, + n_heads: usize, + head_dim: usize, +) -> Option { + batch_size + .checked_mul(n_heads)? + .checked_mul(8)? + .checked_mul(head_dim.checked_add(1)?) +} + +#[inline] +fn q8_multirow_fail_closed_after_split_miss( + arch: &str, + split_admitted: bool, + capture_mode: bool, + replay_recording: bool, +) -> bool { + replay_recording || (arch == "gfx1100" && split_admitted && capture_mode) +} + +#[inline] +fn gfx1100_split_capture_assets_ready( + partial_loaded: bool, + merge_loaded: bool, + preconvert_loaded: bool, + q16_scratch_bytes: usize, + q16_scratch_present: bool, + need_q16_bytes: usize, +) -> bool { + partial_loaded + && merge_loaded + && preconvert_loaded + && !crate::scratch::scratch_will_grow( + q16_scratch_bytes, + q16_scratch_present, + need_q16_bytes, + ) +} + /// `head_dim` envelope of a batched flash-tile kernel, inclusive. /// /// The launcher-level check (`positive multiple of 32, <= 512`) describes the @@ -7038,7 +7084,7 @@ impl Gpu { /// Experimental exact-gfx1100 split-KV twin of the tuned Q16/R3 FA2 /// kernel. The partial grid partitions the live KT32 range across Z and /// the merge performs a stable LSE-weighted reduction. Product dispatch - /// reaches this only through the fail-closed, opt-in admission in + /// reaches this only through the fail-closed exact-shape admission in /// [`Self::attention_flash_q8_0_rows_masked`]. #[doc(hidden)] #[allow(clippy::too_many_arguments)] @@ -7139,6 +7185,7 @@ impl Gpu { const PARTIAL: &str = "attention_q8_0_fa2_gqa_partial_gfx1100"; const MERGE: &str = "attention_q8_0_fa2_gqa_merge_gfx1100"; const PRECONVERT: &str = "attention_fa2_q_preconvert_gfx1100"; + const MODULE: &str = "attention_q8_0_fa2_gqa_gfx1100"; let src = format!( "// HIPFIRE_COMPILER_FLAGS: -mcumode\n\ #define HIPFIRE_FA2_KT 32\n\ @@ -7148,13 +7195,13 @@ impl Gpu { kernels::ATTENTION_Q8_0_FA2_GQA_GFX11_SRC ); if !self.functions.contains_key(PARTIAL) { - self.ensure_kernel(PARTIAL, &src, PARTIAL)?; + self.ensure_kernel(MODULE, &src, PARTIAL)?; } if !self.functions.contains_key(PRECONVERT) { - self.ensure_kernel(PARTIAL, &src, PRECONVERT)?; + self.ensure_kernel(MODULE, &src, PRECONVERT)?; } if !self.functions.contains_key(MERGE) { - self.ensure_kernel(MERGE, &src, MERGE)?; + self.ensure_kernel(MODULE, &src, MERGE)?; } let need_q16_bytes = need_qo.checked_mul(2).ok_or_else(|| { @@ -7270,19 +7317,35 @@ impl Gpu { /// retained-PM4 capture is active. Capture must never trigger JIT or move /// the persistent Q16 scratch pointer; the direct warmup/prime window /// materializes both before either recorder reaches this predicate. - fn gfx1100_q8_fa2_split_capture_ready(&self, batch_size: usize) -> bool { + #[doc(hidden)] + pub fn gfx1100_q8_fa2_split_capture_assets_ready(&self, batch_size: usize) -> bool { const PARTIAL: &str = "attention_q8_0_fa2_gqa_partial_gfx1100"; const MERGE: &str = "attention_q8_0_fa2_gqa_merge_gfx1100"; const PRECONVERT: &str = "attention_fa2_q_preconvert_gfx1100"; let need_q16_bytes = batch_size * 24 * 256 * 2; - self.functions.contains_key(PARTIAL) - && self.functions.contains_key(MERGE) - && self.functions.contains_key(PRECONVERT) - && !crate::scratch::scratch_will_grow( - self.scratch.fa2_q16_scratch_bytes, - self.scratch.fa2_q16_scratch.is_some(), - need_q16_bytes, - ) + gfx1100_split_capture_assets_ready( + self.functions.contains_key(PARTIAL), + self.functions.contains_key(MERGE), + self.functions.contains_key(PRECONVERT), + self.scratch.fa2_q16_scratch_bytes, + self.scratch.fa2_q16_scratch.is_some(), + need_q16_bytes, + ) + } + + /// Full capture readiness, including the caller-owned partial/merge + /// workspace. The auxiliary F16/LDS producers use the asset-only twin + /// after `ChainVerifySplit` has already proved this capacity at the model + /// boundary. + #[doc(hidden)] + pub fn gfx1100_q8_fa2_split_capture_ready( + &self, + batch_size: usize, + partials_numel: usize, + ) -> bool { + self.gfx1100_q8_fa2_split_capture_assets_ready(batch_size) + && gfx1100_q8_fa2_split_partials_len(batch_size, 24, 256) + .is_some_and(|need| partials_numel >= need) } #[allow(clippy::too_many_arguments)] @@ -7299,13 +7362,15 @@ impl Gpu { head_dim: usize, logical_ctx_len: usize, batch_size: usize, + route_admitted: bool, ) -> HipResult { const SPLITS: usize = 8; let capturing = self.graphs.capture_mode || self.replay.is_recording(); - let capture_ready = !capturing || self.gfx1100_q8_fa2_split_capture_ready(batch_size); + let capture_ready = + !capturing || self.gfx1100_q8_fa2_split_capture_ready(batch_size, partials.numel()); if !capture_ready || !gfx1100_q8_fa2_split_admitted( - hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false), + route_admitted && self.flags.verify_attn && self.flags.gfx1100_fa2_split_verify, &self.arch, logical_ctx_len, n_heads, @@ -7322,10 +7387,7 @@ impl Gpu { else { return Ok(false); }; - let Some(need_partials) = batch_size - .checked_mul(n_heads) - .and_then(|v| v.checked_mul(SPLITS)) - .and_then(|v| v.checked_mul(head_dim + 1)) + let Some(need_partials) = gfx1100_q8_fa2_split_partials_len(batch_size, n_heads, head_dim) else { return Ok(false); }; @@ -8875,6 +8937,7 @@ impl Gpu { max_ctx_len, batch_size, partials, + true, ) } @@ -8899,6 +8962,7 @@ impl Gpu { logical_ctx_len: usize, batch_size: usize, partials: &GpuTensor, + allow_gfx1100_split: bool, ) -> HipResult { if !q8_multirow_arch_supported(self.arch_caps.arch()) { return Ok(false); @@ -8924,13 +8988,19 @@ impl Gpu { head_dim, logical_ctx_len, batch_size, + allow_gfx1100_split, )? { return Ok(true); } - // The incumbent multi-row and VerifyGQA routes are not certified for - // retained recording. The split-KV route above is: it has fixed launch - // geometry, recorder-owned kernargs, and a warmup-stable scratch base. - if self.replay.is_recording() { + // The incumbent multi-row/VerifyGQA routes are outside this captured + // split contract. If split admission or readiness misses during either + // recorder, fail closed to the caller's established batched route. + if q8_multirow_fail_closed_after_split_miss( + self.arch_caps.arch(), + allow_gfx1100_split, + self.graphs.capture_mode, + self.replay.is_recording(), + ) { return Ok(false); } if self.try_attention_verify_gqa_rows( @@ -23070,9 +23140,10 @@ fn kv_slot_paged_symbol(func: &str) -> Option<&'static str> { mod tests { use super::{ flux_attn_dtype_error, flux_attn_dtype_suffix, flux_attn_route_dtypes, - flux_attn_route_name, gfx1100_q8_fa2_split_admitted, - pack_attention_q8_0_fa2_gqa_gfx11_kernarg, q8_flash_default_tile_size, - q8_flash_reduce_safe_tile_size, q8_multirow_arch_supported, replay_stable_tile_count, + flux_attn_route_name, gfx1100_q8_fa2_split_admitted, gfx1100_q8_fa2_split_partials_len, + gfx1100_split_capture_assets_ready, pack_attention_q8_0_fa2_gqa_gfx11_kernarg, + q8_flash_default_tile_size, q8_flash_reduce_safe_tile_size, q8_multirow_arch_supported, + q8_multirow_fail_closed_after_split_miss, replay_stable_tile_count, }; use crate::DType; use std::ffi::c_void; @@ -23155,6 +23226,62 @@ mod tests { )); } + #[test] + fn split_miss_capture_fallback_is_gfx1100_only() { + assert!(q8_multirow_fail_closed_after_split_miss( + "gfx1100", true, true, false, + )); + for arch in ["gfx1151", "gfx1201"] { + assert!( + !q8_multirow_fail_closed_after_split_miss(arch, true, true, false), + "arch={arch} must retain its incumbent HipGraph VerifyAttn path" + ); + } + assert!(!q8_multirow_fail_closed_after_split_miss( + "gfx1100", false, true, false, + )); + for arch in ["gfx1100", "gfx1151", "gfx1201"] { + assert!(q8_multirow_fail_closed_after_split_miss( + arch, false, false, true, + )); + } + } + + #[test] + fn split_capture_assets_require_all_symbols_and_fixed_q16_scratch() { + assert!(gfx1100_split_capture_assets_ready( + true, true, true, 196_608, true, 196_608, + )); + assert!(!gfx1100_split_capture_assets_ready( + false, true, true, 196_608, true, 196_608, + )); + assert!(!gfx1100_split_capture_assets_ready( + true, false, true, 196_608, true, 196_608, + )); + assert!(!gfx1100_split_capture_assets_ready( + true, true, false, 196_608, true, 196_608, + )); + assert!(!gfx1100_split_capture_assets_ready( + true, true, true, 196_607, true, 196_608, + )); + assert!(!gfx1100_split_capture_assets_ready( + true, true, true, 196_608, false, 196_608, + )); + } + + #[test] + fn gfx1100_split_partials_len_is_checked_and_matches_s8_layout() { + assert_eq!( + gfx1100_q8_fa2_split_partials_len(16, 24, 256), + Some(16 * 24 * 8 * 257) + ); + assert_eq!( + gfx1100_q8_fa2_split_partials_len(32, 24, 256), + Some(32 * 24 * 8 * 257) + ); + assert_eq!(gfx1100_q8_fa2_split_partials_len(usize::MAX, 24, 256), None); + } + /// The suffix table is the mapping from tensor dtypes to a kernel symbol. /// Get an arm wrong and the launcher asks for an entry that either does /// not exist (a load failure) or exists but reads the buffer at the wrong diff --git a/crates/rdna-compute/src/feature_flags.rs b/crates/rdna-compute/src/feature_flags.rs index 2e3d1c332e..cfd75c8bf8 100644 --- a/crates/rdna-compute/src/feature_flags.rs +++ b/crates/rdna-compute/src/feature_flags.rs @@ -389,6 +389,11 @@ pub struct FeatureFlags { /// and `attention_q8_0_flash_prefill_wmma`. Byte-identical output; any /// byte difference kills it. pub verify_attn: bool, + /// Exact-gfx1100 FA2 split-KV implementation of VerifyAttn + /// (`HIPFIRE_GFX1100_FA2_SPLIT_VERIFY`, + /// `kernel.gfx1100_fa2_split_verify`). Default ON only on exact gfx1100; + /// the stable [`Self::verify_attn`] switch remains the parent opt-out. + pub gfx1100_fa2_split_verify: bool, /// Exact gfx1201 FA deinterleave, Q/K norm and RoPE fusion. /// `HIPFIRE_GFX12_FA_PREP_FUSED=0` restores the original chain. pub gfx12_fa_prep_fused: bool, @@ -696,13 +701,21 @@ impl FeatureFlags { qwen4_hc_fuse: value("HIPFIRE_QWEN4_HC_FUSE") .ok() .and_then(|v| v.trim().parse::().ok()) - .unwrap_or(if halo_hyper_default { 3 } else if is_gfx1201 { 1 } else { 0 }), + .unwrap_or(if halo_hyper_default { + 3 + } else if is_gfx1201 { + 1 + } else { + 0 + }), qwen4_hc_up_tile: parse_bool("HIPFIRE_QWEN4_HC_UP_TILE") .unwrap_or(halo_hyper_default || is_gfx1201), qwen4_moe_combine_zinit: parse_bool("HIPFIRE_QWEN4_MOE_COMBINE_ZINIT") .unwrap_or(exact_unit_default), - qwen4_hc_row_fold: parse_bool("HIPFIRE_QWEN4_HC_ROW_FOLD").unwrap_or(halo_hyper_default), - qwen4_router_fast: parse_bool("HIPFIRE_QWEN4_ROUTER_FAST").unwrap_or(exact_unit_default), + qwen4_hc_row_fold: parse_bool("HIPFIRE_QWEN4_HC_ROW_FOLD") + .unwrap_or(halo_hyper_default), + qwen4_router_fast: parse_bool("HIPFIRE_QWEN4_ROUTER_FAST") + .unwrap_or(exact_unit_default), qwen4_ple_fuse: parse_bool("HIPFIRE_QWEN4_PLE_FUSE").unwrap_or(exact_unit_default), gfx11_iu4_gridspec: parse_bool("HIPFIRE_GFX11_IU4_GRIDSPEC").unwrap_or(true), gfx11_iu4_shape: parse_bool("HIPFIRE_GFX11_IU4_SHAPE").unwrap_or(true), @@ -879,18 +892,15 @@ impl FeatureFlags { gfx1100_dec_norm: parse_bool("HIPFIRE_GFX1100_DEC_NORM").unwrap_or(arch == "gfx1100"), gfx11_producer_quant_fused: parse_bool("HIPFIRE_GFX11_PRODUCER_QUANT_FUSED") .unwrap_or(matches!(arch, "gfx1100" | "gfx1151")), - gfx12_fp8_stream: parse_bool("HIPFIRE_GFX12_FP8_STREAM") - .unwrap_or(arch == "gfx1201"), - gfx12_fa2_prefill: parse_bool("HIPFIRE_GFX12_FA2_PREFILL") - .unwrap_or(arch == "gfx1201"), - gfx12_fa_packet: parse_bool("HIPFIRE_GFX12_FA_PACKET") - .unwrap_or(arch == "gfx1201"), - attn_qresident: parse_bool("HIPFIRE_ATTN_QRESIDENT") - .unwrap_or(arch == "gfx1201"), - attn_qresident_v2: parse_bool("HIPFIRE_ATTN_QRESIDENT_V2") - .unwrap_or(arch == "gfx1201"), + gfx12_fp8_stream: parse_bool("HIPFIRE_GFX12_FP8_STREAM").unwrap_or(arch == "gfx1201"), + gfx12_fa2_prefill: parse_bool("HIPFIRE_GFX12_FA2_PREFILL").unwrap_or(arch == "gfx1201"), + gfx12_fa_packet: parse_bool("HIPFIRE_GFX12_FA_PACKET").unwrap_or(arch == "gfx1201"), + attn_qresident: parse_bool("HIPFIRE_ATTN_QRESIDENT").unwrap_or(arch == "gfx1201"), + attn_qresident_v2: parse_bool("HIPFIRE_ATTN_QRESIDENT_V2").unwrap_or(arch == "gfx1201"), verify_attn: parse_bool("HIPFIRE_VERIFY_ATTN") .unwrap_or(matches!(arch, "gfx1201" | "gfx1100" | "gfx1151")), + gfx1100_fa2_split_verify: parse_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY") + .unwrap_or(arch == "gfx1100"), gfx12_fa_prep_fused: parse_bool("HIPFIRE_GFX12_FA_PREP_FUSED") .unwrap_or(arch == "gfx1201"), gfx12_fa_prep_fp8q: parse_bool("HIPFIRE_GFX12_FA_PREP_FP8Q") @@ -1113,8 +1123,7 @@ impl FeatureFlags { /// SwiGLU/FWHT emit `block_i4_128` in-register; otherwise consumers /// keep standalone `quantize_int4_mmq_ds128`. pub fn iu4_producer_sidecar_enabled(&self) -> bool { - self.iu4_prefill.unwrap_or(true) - && matches!(self.arch.as_str(), "gfx1100" | "gfx1151") + self.iu4_prefill.unwrap_or(true) && matches!(self.arch.as_str(), "gfx1100" | "gfx1151") } /// True only on exact gfx1201 with the opt-in set. The producer emits /// the shared `block_i4_128` recipe, so output is bit-identical to the @@ -1348,6 +1357,7 @@ impl FeatureFlags { attn_qresident: false, attn_qresident_v2: false, verify_attn: false, + gfx1100_fa2_split_verify: false, gfx12_fa_prep_fused: false, gfx12_fa_prep_fp8q: false, gfx11_q8_fa2_wide: false, @@ -1490,12 +1500,16 @@ mod tests { move |name: &str| -> std::result::Result { match name { "HIPFIRE_QWEN4_HC_FUSE" => Ok(fuse.into()), - "HIPFIRE_QWEN4_HC_UP_TILE" | "HIPFIRE_QWEN4_MOE_COMBINE_ZINIT" => Ok(rest.into()), + "HIPFIRE_QWEN4_HC_UP_TILE" | "HIPFIRE_QWEN4_MOE_COMBINE_ZINIT" => { + Ok(rest.into()) + } _ => Err(()), } } }; - for arch in ["gfx906", "gfx1030", "gfx1100", "gfx1150", "gfx1151", "gfx1200", "gfx1201"] { + for arch in [ + "gfx906", "gfx1030", "gfx1100", "gfx1150", "gfx1151", "gfx1200", "gfx1201", + ] { let halo = arch == "gfx1151"; let exact = matches!(arch, "gfx1151" | "gfx1201"); let default_level = match arch { @@ -1512,7 +1526,11 @@ mod tests { assert_eq!(off.qwen4_hc_fuse_level(), 0, "{arch}"); assert!(!off.qwen4_hc_up_tile_enabled() && !off.qwen4_moe_combine_zinit_enabled()); let on = FeatureFlags::from_lookup(arch, with("2", "1")); - assert_eq!(on.qwen4_hc_fuse_level(), if exact { 2 } else { 0 }, "{arch}"); + assert_eq!( + on.qwen4_hc_fuse_level(), + if exact { 2 } else { 0 }, + "{arch}" + ); assert_eq!(on.qwen4_hc_up_tile_enabled(), exact, "{arch}"); assert_eq!(on.qwen4_moe_combine_zinit_enabled(), exact, "{arch}"); assert_eq!( @@ -1531,10 +1549,16 @@ mod tests { fn qwen4_ple_fuse_flag_defaults_on_exact_gfx1151_gfx1201() { let with = |value: &'static str| { move |name: &str| -> std::result::Result { - if name == "HIPFIRE_QWEN4_PLE_FUSE" { Ok(value.into()) } else { Err(()) } + if name == "HIPFIRE_QWEN4_PLE_FUSE" { + Ok(value.into()) + } else { + Err(()) + } } }; - for arch in ["gfx906", "gfx1100", "gfx1150", "gfx1151", "gfx1200", "gfx1201"] { + for arch in [ + "gfx906", "gfx1100", "gfx1150", "gfx1151", "gfx1200", "gfx1201", + ] { let exact = matches!(arch, "gfx1151" | "gfx1201"); assert_eq!( FeatureFlags::from_lookup(arch, |_| Err(())).qwen4_ple_fuse_enabled(), @@ -1553,16 +1577,26 @@ mod tests { /// The wave-per-token router defaults on for exact gfx1151 / gfx1201; `0` turns it off. #[test] fn qwen4_router_fast_defaults_on_exact_gfx1151_gfx1201() { - for arch in ["gfx906", "gfx1100", "gfx1150", "gfx1151", "gfx1200", "gfx1201"] { + for arch in [ + "gfx906", "gfx1100", "gfx1150", "gfx1151", "gfx1200", "gfx1201", + ] { let exact = matches!(arch, "gfx1151" | "gfx1201"); let unset = FeatureFlags::from_lookup(arch, |_| Err(())); assert_eq!(unset.qwen4_router_fast_enabled(), exact, "{arch}"); let off = FeatureFlags::from_lookup(arch, |n| { - if n == "HIPFIRE_QWEN4_ROUTER_FAST" { Ok("0".into()) } else { Err(()) } + if n == "HIPFIRE_QWEN4_ROUTER_FAST" { + Ok("0".into()) + } else { + Err(()) + } }); assert!(!off.qwen4_router_fast_enabled(), "{arch}"); let on = FeatureFlags::from_lookup(arch, |n| { - if n == "HIPFIRE_QWEN4_ROUTER_FAST" { Ok("1".into()) } else { Err(()) } + if n == "HIPFIRE_QWEN4_ROUTER_FAST" { + Ok("1".into()) + } else { + Err(()) + } }); assert_eq!(on.qwen4_router_fast_enabled(), exact, "{arch}"); } @@ -1597,14 +1631,26 @@ mod tests { "gfx1200", "gfx1201", ] { let unset = FeatureFlags::from_lookup(arch, |_| Err(())); - assert_eq!(unset.qwen4_gdn_conv_qknorm_enabled(), arch == "gfx1151", "{arch}: unset"); + assert_eq!( + unset.qwen4_gdn_conv_qknorm_enabled(), + arch == "gfx1151", + "{arch}: unset" + ); assert!(!unset.qwen4_gdn_q8_inline_enabled(), "{arch}: unset"); let zero = FeatureFlags::from_lookup(arch, off); assert!(!zero.qwen4_gdn_conv_qknorm_enabled(), "{arch}: 0"); assert!(!zero.qwen4_gdn_q8_inline_enabled(), "{arch}: 0"); let one = FeatureFlags::from_lookup(arch, on); - assert_eq!(one.qwen4_gdn_conv_qknorm_enabled(), arch == "gfx1151", "{arch}: 1"); - assert_eq!(one.qwen4_gdn_q8_inline_enabled(), arch == "gfx1151", "{arch}: 1"); + assert_eq!( + one.qwen4_gdn_conv_qknorm_enabled(), + arch == "gfx1151", + "{arch}: 1" + ); + assert_eq!( + one.qwen4_gdn_q8_inline_enabled(), + arch == "gfx1151", + "{arch}: 1" + ); } assert!(!FeatureFlags::for_test("gfx1151").qwen4_gdn_conv_qknorm_enabled()); assert!(!FeatureFlags::for_test("gfx1151").qwen4_gdn_q8_inline_enabled()); @@ -1733,6 +1779,56 @@ mod tests { assert!(!test_flags.gfx11_fa2_prefill, "arch={arch}"); } } + + #[test] + fn gfx1100_split_verify_defaults_only_on_exact_arch_and_keeps_two_opt_outs() { + let resolved = resolve([]).unwrap(); + let process = ProcessConfig::from_resolved(&resolved).unwrap(); + let gfx1100 = FeatureFlags::from_process_config("gfx1100", &process); + assert!(gfx1100.verify_attn); + assert!(gfx1100.gfx1100_fa2_split_verify); + for arch in ["gfx1101", "gfx1151", "gfx1200", "gfx1201", "gfx942"] { + assert!( + !FeatureFlags::from_process_config(arch, &process).gfx1100_fa2_split_verify, + "arch={arch}" + ); + } + + let mut parent_off = ConfigLayer::default(); + parent_off.set_cli("kernel.verify_attn", "false").unwrap(); + let resolved = resolve([NamedLayer { + source: ConfigSource::GlobalUser { + path: "config.toml".into(), + }, + layer: parent_off, + }]) + .unwrap(); + let flags = FeatureFlags::from_process_config( + "gfx1100", + &ProcessConfig::from_resolved(&resolved).unwrap(), + ); + assert!(!flags.verify_attn); + assert!(flags.gfx1100_fa2_split_verify); + + let mut split_off = ConfigLayer::default(); + split_off + .set_cli("kernel.gfx1100_fa2_split_verify", "false") + .unwrap(); + let resolved = resolve([NamedLayer { + source: ConfigSource::GlobalUser { + path: "config.toml".into(), + }, + layer: split_off, + }]) + .unwrap(); + let flags = FeatureFlags::from_process_config( + "gfx1100", + &ProcessConfig::from_resolved(&resolved).unwrap(), + ); + assert!(flags.verify_attn); + assert!(!flags.gfx1100_fa2_split_verify); + } + #[test] fn gfx12_silu_quant_fused_default_on_gfx1201_with_opt_out() { // Default process policy: the exact-gfx1201 silu+quant fusion admits @@ -1743,14 +1839,18 @@ mod tests { let gfx1201 = FeatureFlags::from_process_config("gfx1201", &process); assert!(gfx1201.gfx12_silu_quant_fused); assert!(gfx1201.gfx12_silu_quant_fused_enabled()); - for arch in ["gfx1100", "gfx1151", "gfx1101", "gfx1102", "gfx1150", "gfx1200", "gfx942"] { + for arch in [ + "gfx1100", "gfx1151", "gfx1101", "gfx1102", "gfx1150", "gfx1200", "gfx942", + ] { let flags = FeatureFlags::from_process_config(arch, &process); assert!(!flags.gfx12_silu_quant_fused, "arch={arch}"); assert!(!flags.gfx12_silu_quant_fused_enabled(), "arch={arch}"); } let mut layer = ConfigLayer::default(); - layer.set_cli("kernel.gfx12_silu_quant_fused", "false").unwrap(); + layer + .set_cli("kernel.gfx12_silu_quant_fused", "false") + .unwrap(); let resolved = resolve([NamedLayer { source: ConfigSource::GlobalUser { path: "config.toml".into(), @@ -1780,14 +1880,18 @@ mod tests { let gfx1201 = FeatureFlags::from_process_config("gfx1201", &process); assert!(gfx1201.gfx12_producer_quant_fused); assert!(gfx1201.gfx12_producer_quant_fused_enabled()); - for arch in ["gfx1100", "gfx1151", "gfx1101", "gfx1102", "gfx1150", "gfx1200", "gfx942"] { + for arch in [ + "gfx1100", "gfx1151", "gfx1101", "gfx1102", "gfx1150", "gfx1200", "gfx942", + ] { let flags = FeatureFlags::from_process_config(arch, &process); assert!(!flags.gfx12_producer_quant_fused, "arch={arch}"); assert!(!flags.gfx12_producer_quant_fused_enabled(), "arch={arch}"); } let mut layer = ConfigLayer::default(); - layer.set_cli("kernel.gfx12_producer_quant_fused", "false").unwrap(); + layer + .set_cli("kernel.gfx12_producer_quant_fused", "false") + .unwrap(); let resolved = resolve([NamedLayer { source: ConfigSource::GlobalUser { path: "config.toml".into(), @@ -1819,14 +1923,18 @@ mod tests { assert!(flags.gfx11_producer_quant_fused, "arch={arch}"); assert!(flags.gfx11_producer_quant_fused_enabled(), "arch={arch}"); } - for arch in ["gfx1201", "gfx1101", "gfx1102", "gfx1150", "gfx1200", "gfx942"] { + for arch in [ + "gfx1201", "gfx1101", "gfx1102", "gfx1150", "gfx1200", "gfx942", + ] { let flags = FeatureFlags::from_process_config(arch, &process); assert!(!flags.gfx11_producer_quant_fused, "arch={arch}"); assert!(!flags.gfx11_producer_quant_fused_enabled(), "arch={arch}"); } let mut layer = ConfigLayer::default(); - layer.set_cli("kernel.gfx11_producer_quant_fused", "false").unwrap(); + layer + .set_cli("kernel.gfx11_producer_quant_fused", "false") + .unwrap(); let resolved = resolve([NamedLayer { source: ConfigSource::GlobalUser { path: "config.toml".into(), @@ -1863,7 +1971,9 @@ mod tests { let mut off = ConfigLayer::default(); off.set_cli("kernel.gfx12_fp8_stream", "false").unwrap(); let resolved = resolve([NamedLayer { - source: ConfigSource::GlobalUser { path: "config.toml".into() }, + source: ConfigSource::GlobalUser { + path: "config.toml".into(), + }, layer: off, }]) .unwrap(); @@ -2049,18 +2159,29 @@ mod tests { ("gfx1151", true), ("gfx1201", false), ] { - assert_eq!(FeatureFlags::from_process_config(arch, &process).gfx11_q8_fa2_wide, expected, "arch={arch}"); + assert_eq!( + FeatureFlags::from_process_config(arch, &process).gfx11_q8_fa2_wide, + expected, + "arch={arch}" + ); } for (value, expected) in [("false", false), ("true", true)] { let mut layer = ConfigLayer::default(); layer.set_cli("kernel.gfx11_q8_fa2_wide", value).unwrap(); let resolved = resolve([NamedLayer { - source: ConfigSource::GlobalUser { path: "config.toml".into() }, + source: ConfigSource::GlobalUser { + path: "config.toml".into(), + }, layer, - }]).unwrap(); + }]) + .unwrap(); let process = ProcessConfig::from_resolved(&resolved).unwrap(); for arch in ["gfx1100", "gfx1151"] { - assert_eq!(FeatureFlags::from_process_config(arch, &process).gfx11_q8_fa2_wide, expected, "arch={arch} value={value}"); + assert_eq!( + FeatureFlags::from_process_config(arch, &process).gfx11_q8_fa2_wide, + expected, + "arch={arch} value={value}" + ); } } } diff --git a/crates/rdna-compute/src/kernel_registry.rs b/crates/rdna-compute/src/kernel_registry.rs index 05e938114f..dc9eaeec0e 100644 --- a/crates/rdna-compute/src/kernel_registry.rs +++ b/crates/rdna-compute/src/kernel_registry.rs @@ -29,31 +29,42 @@ mod kernels { #[cfg(not(feature = "deltanet"))] pub const L2_NORM_SRC: &str = include_str!("../../../kernels/src/l2_norm.hip"); #[cfg(not(feature = "deltanet"))] - pub const FUSED_QK_L2_NORM_SCALE_SRC: &str = include_str!("../../../kernels/src/fused_qk_l2_norm_scale.hip"); + pub const FUSED_QK_L2_NORM_SCALE_SRC: &str = + include_str!("../../../kernels/src/fused_qk_l2_norm_scale.hip"); #[cfg(not(feature = "deltanet"))] - pub const FUSED_SIGMOID_ALPHA_GATE_SRC: &str = include_str!("../../../kernels/src/fused_sigmoid_alpha_gate.hip"); + pub const FUSED_SIGMOID_ALPHA_GATE_SRC: &str = + include_str!("../../../kernels/src/fused_sigmoid_alpha_gate.hip"); #[cfg(not(feature = "deltanet"))] - pub const CONV1D_SILU_SPLIT_SRC: &str = include_str!("../../../kernels/src/conv1d_silu_split.hip"); + pub const CONV1D_SILU_SPLIT_SRC: &str = + include_str!("../../../kernels/src/conv1d_silu_split.hip"); #[cfg(not(feature = "deltanet"))] - pub const CONV1D_SILU_SPLIT_TREE_SRC: &str = include_str!("../../../kernels/src/conv1d_silu_split_tree.hip"); + pub const CONV1D_SILU_SPLIT_TREE_SRC: &str = + include_str!("../../../kernels/src/conv1d_silu_split_tree.hip"); #[cfg(not(feature = "deltanet"))] - pub const GATED_DELTA_NET_Q8_TREE_SRC: &str = include_str!("../../../kernels/src/gated_delta_net_q8_tree.hip"); + pub const GATED_DELTA_NET_Q8_TREE_SRC: &str = + include_str!("../../../kernels/src/gated_delta_net_q8_tree.hip"); #[cfg(not(feature = "deltanet"))] pub const SCALE_F32_SRC: &str = include_str!("../../../kernels/src/scale_f32.hip"); #[cfg(not(feature = "deltanet"))] pub const GATED_NORM_SRC: &str = include_str!("../../../kernels/src/gated_norm.hip"); #[cfg(not(feature = "deltanet"))] - pub const ROPE_PARTIAL_INTERLEAVED_SRC: &str = include_str!("../../../kernels/src/rope_partial_interleaved.hip"); + pub const ROPE_PARTIAL_INTERLEAVED_SRC: &str = + include_str!("../../../kernels/src/rope_partial_interleaved.hip"); #[cfg(not(feature = "deltanet"))] - pub const GATED_DELTA_NET_Q8_SRC: &str = include_str!("../../../kernels/src/gated_delta_net_q8.hip"); + pub const GATED_DELTA_NET_Q8_SRC: &str = + include_str!("../../../kernels/src/gated_delta_net_q8.hip"); #[cfg(not(feature = "deltanet"))] - pub const ROPE_PARTIAL_HALFSPLIT_SRC: &str = include_str!("../../../kernels/src/rope_partial_halfsplit.hip"); + pub const ROPE_PARTIAL_HALFSPLIT_SRC: &str = + include_str!("../../../kernels/src/rope_partial_halfsplit.hip"); #[cfg(not(feature = "deltanet"))] - pub const ROPE_PARTIAL_HALFSPLIT_HEADGRID_SRC: &str = include_str!("../../../kernels/src/rope_partial_halfsplit_headgrid.hip"); + pub const ROPE_PARTIAL_HALFSPLIT_HEADGRID_SRC: &str = + include_str!("../../../kernels/src/rope_partial_halfsplit_headgrid.hip"); #[cfg(not(feature = "deltanet"))] - pub const ROPE_PARTIAL_HALFSPLIT_BATCHED_SRC: &str = include_str!("../../../kernels/src/rope_partial_halfsplit_batched.hip"); + pub const ROPE_PARTIAL_HALFSPLIT_BATCHED_SRC: &str = + include_str!("../../../kernels/src/rope_partial_halfsplit_batched.hip"); #[cfg(not(feature = "deltanet"))] - pub const QWEN35_FA_PREP_GFX1100_SRC: &str = include_str!("../../../kernels/src/qwen35_fa_prep.gfx1100.hip"); + pub const QWEN35_FA_PREP_GFX1100_SRC: &str = + include_str!("../../../kernels/src/qwen35_fa_prep.gfx1100.hip"); #[cfg(not(feature = "deltanet"))] pub const QWEN35_FA_PREP_GFX1151_SRC: &str = concat!( "#define HIPFIRE_QWEN35_FA_PREP_KERNEL qwen35_fa_prep_gfx1151\n", @@ -190,276 +201,1408 @@ pub fn entries(arch: &str, extra_flags: &str) -> Result, Regist add!("alpha_gate", kernels::ALPHA_GATE_SRC, ["alpha_gate_f32"]); add!("conv1d_silu", kernels::CONV1D_SILU_SRC, ["conv1d_silu_f32"]); add!("l2_norm", kernels::L2_NORM_SRC, ["l2_norm_f32"]); - add!("fused_qk_l2_norm_scale", kernels::FUSED_QK_L2_NORM_SCALE_SRC, ["fused_qk_l2_norm_scale_f32"]); - add!("fused_sigmoid_alpha_gate", kernels::FUSED_SIGMOID_ALPHA_GATE_SRC, ["fused_sigmoid_alpha_gate_f32"]); - add!("conv1d_silu_split", kernels::CONV1D_SILU_SPLIT_SRC, ["conv1d_silu_split_f32"]); - add!("conv1d_silu_split_tree", kernels::CONV1D_SILU_SPLIT_TREE_SRC, ["conv1d_silu_split_tree_f32"]); - add!("gated_delta_net_q8_tree", kernels::GATED_DELTA_NET_Q8_TREE_SRC, ["gated_delta_net_q8_tree"]); + add!( + "fused_qk_l2_norm_scale", + kernels::FUSED_QK_L2_NORM_SCALE_SRC, + ["fused_qk_l2_norm_scale_f32"] + ); + add!( + "fused_sigmoid_alpha_gate", + kernels::FUSED_SIGMOID_ALPHA_GATE_SRC, + ["fused_sigmoid_alpha_gate_f32"] + ); + add!( + "conv1d_silu_split", + kernels::CONV1D_SILU_SPLIT_SRC, + ["conv1d_silu_split_f32"] + ); + add!( + "conv1d_silu_split_tree", + kernels::CONV1D_SILU_SPLIT_TREE_SRC, + ["conv1d_silu_split_tree_f32"] + ); + add!( + "gated_delta_net_q8_tree", + kernels::GATED_DELTA_NET_Q8_TREE_SRC, + ["gated_delta_net_q8_tree"] + ); add!("sigmoid_mul", kernels::SIGMOID_MUL_SRC, ["sigmoid_mul_f32"]); add!("topk_logits", kernels::TOPK_LOGITS_SRC, ["topk_logits_f32"]); add!("scale_f32", kernels::SCALE_F32_SRC, ["scale_f32"]); add!("gated_norm", kernels::GATED_NORM_SRC, ["gated_norm_f32"]); - add!("rope_partial_interleaved", kernels::ROPE_PARTIAL_INTERLEAVED_SRC, ["rope_partial_interleaved_f32"]); - add!("deinterleave", kernels::DEINTERLEAVE_SRC, ["deinterleave_f32"]); - add!("repeat_interleave_qk", kernels::REPEAT_INTERLEAVE_QK_SRC, ["repeat_interleave_qk_f32"]); + add!( + "rope_partial_interleaved", + kernels::ROPE_PARTIAL_INTERLEAVED_SRC, + ["rope_partial_interleaved_f32"] + ); + add!( + "deinterleave", + kernels::DEINTERLEAVE_SRC, + ["deinterleave_f32"] + ); + add!( + "repeat_interleave_qk", + kernels::REPEAT_INTERLEAVE_QK_SRC, + ["repeat_interleave_qk_f32"] + ); add!("embedding_q8", kernels::EMBEDDING_Q8_SRC, ["embedding_q8"]); - add!("embedding_q8_batched", kernels::EMBEDDING_Q8_BATCHED_SRC, ["embedding_q8_batched"]); - add!("embedding_hfq4g128", kernels::EMBEDDING_HFQ4G128_SRC, ["embedding_hfq4g128"]); - add!("embedding_hfq4g128_batched", kernels::EMBEDDING_HFQ4G128_BATCHED_SRC, ["embedding_hfq4g128_batched"]); - add!("embedding_hfq4g256", kernels::EMBEDDING_HFQ4G256_SRC, ["embedding_hfq4g256"]); - add!("embedding_hfq4g256_batched", kernels::EMBEDDING_HFQ4G256_BATCHED_SRC, ["embedding_hfq4g256_batched"]); - add!("gated_delta_net_q8", kernels::GATED_DELTA_NET_Q8_SRC, ["gated_delta_net_q8"]); + add!( + "embedding_q8_batched", + kernels::EMBEDDING_Q8_BATCHED_SRC, + ["embedding_q8_batched"] + ); + add!( + "embedding_hfq4g128", + kernels::EMBEDDING_HFQ4G128_SRC, + ["embedding_hfq4g128"] + ); + add!( + "embedding_hfq4g128_batched", + kernels::EMBEDDING_HFQ4G128_BATCHED_SRC, + ["embedding_hfq4g128_batched"] + ); + add!( + "embedding_hfq4g256", + kernels::EMBEDDING_HFQ4G256_SRC, + ["embedding_hfq4g256"] + ); + add!( + "embedding_hfq4g256_batched", + kernels::EMBEDDING_HFQ4G256_BATCHED_SRC, + ["embedding_hfq4g256_batched"] + ); + add!( + "gated_delta_net_q8", + kernels::GATED_DELTA_NET_Q8_SRC, + ["gated_delta_net_q8"] + ); // Matrix-dependent default precompile modules; each name and source is // taken from the same constants as dispatch.rs:4900-5174. - add!("gemv_hfq6g256", kernels::GEMV_HFQ6G256_SRC, ["gemv_hfq6g256"]); + add!( + "gemv_hfq6g256", + kernels::GEMV_HFQ6G256_SRC, + ["gemv_hfq6g256"] + ); add!("gemv_mq6g256", kernels::GEMV_MQ6G256_SRC, ["gemv_mq6g256"]); add!("gemm_mq6g256", kernels::GEMM_MQ6G256_SRC, ["gemm_mq6g256"]); add!("gemv_q8_0", kernels::GEMV_Q8_0_SRC, ["gemv_q8_0"]); - add!("gemv_mq4g256", kernels::GEMV_MQ4G256_SRC, ["gemv_mq4g256", "mq_rotate_x"]); - add!("gemv_hfq4g256_wide", kernels::GEMV_HFQ4G256_WIDE_SRC, ["gemv_hfq4g256_wide"]); - add!("fused_qkvza_hfq4g256", kernels::FUSED_QKVZA_HFQ4G256_SRC, ["fused_qkvza_hfq4g256"]); - add!("fused_qkv_hfq4g256", kernels::FUSED_QKV_HFQ4G256_SRC, ["fused_qkv_hfq4g256"]); - add!("fused_gate_up_hfq4g256", kernels::FUSED_GATE_UP_HFQ4G256_SRC, ["fused_gate_up_hfq4g256"]); - add!("fused_rmsnorm_mq_rotate", kernels::FUSED_RMSNORM_MQ_ROTATE_SRC, ["fused_rmsnorm_mq_rotate"]); - add!("fused_silu_mul_mq_rotate", kernels::FUSED_SILU_MUL_MQ_ROTATE_SRC, ["fused_silu_mul_mq_rotate"]); + add!( + "gemv_mq4g256", + kernels::GEMV_MQ4G256_SRC, + ["gemv_mq4g256", "mq_rotate_x"] + ); + add!( + "gemv_hfq4g256_wide", + kernels::GEMV_HFQ4G256_WIDE_SRC, + ["gemv_hfq4g256_wide"] + ); + add!( + "fused_qkvza_hfq4g256", + kernels::FUSED_QKVZA_HFQ4G256_SRC, + ["fused_qkvza_hfq4g256"] + ); + add!( + "fused_qkv_hfq4g256", + kernels::FUSED_QKV_HFQ4G256_SRC, + ["fused_qkv_hfq4g256"] + ); + add!( + "fused_gate_up_hfq4g256", + kernels::FUSED_GATE_UP_HFQ4G256_SRC, + ["fused_gate_up_hfq4g256"] + ); + add!( + "fused_rmsnorm_mq_rotate", + kernels::FUSED_RMSNORM_MQ_ROTATE_SRC, + ["fused_rmsnorm_mq_rotate"] + ); + add!( + "fused_silu_mul_mq_rotate", + kernels::FUSED_SILU_MUL_MQ_ROTATE_SRC, + ["fused_silu_mul_mq_rotate"] + ); match arch { - "gfx1100" => add!("gemv_hfq4g256_rdna3", kernels::GEMV_HFQ4G256_GFX1100_SRC, ["gemv_hfq4g256"]), - _ => add!("gemv_hfq4g256", kernels::GEMV_HFQ4G256_SRC, ["gemv_hfq4g256"]), + "gfx1100" => add!( + "gemv_hfq4g256_rdna3", + kernels::GEMV_HFQ4G256_GFX1100_SRC, + ["gemv_hfq4g256"] + ), + _ => add!( + "gemv_hfq4g256", + kernels::GEMV_HFQ4G256_SRC, + ["gemv_hfq4g256"] + ), } if arch == "gfx1100" { - add!("fused_gate_up_mq4g256v2_k5120_gfx1100", kernels::fused_gate_up_mq4g256v2_k5120_gfx1100_src(), ["fused_gate_up_mq4g256v2_k5120_gfx1100"]); + add!( + "fused_gate_up_mq4g256v2_k5120_gfx1100", + kernels::fused_gate_up_mq4g256v2_k5120_gfx1100_src(), + ["fused_gate_up_mq4g256v2_k5120_gfx1100"] + ); // Default gfx1100 precompile branches (dispatch.rs:4944-4962,5304-5308). - add!("fused_qkvza_hfq4g256_k2048_gfx1100", kernels::FUSED_QKVZA_HFQ4G256_K2048_GFX1100_SRC, ["fused_qkvza_hfq4g256_k2048"]); - add!("fused_gate_up_hfq4g256_stage_x32_gfx1100", kernels::FUSED_GATE_UP_HFQ4G256_STAGE_X32_GFX1100_SRC, ["fused_gate_up_hfq4g256_stage_x32_gfx1100"]); - add!("kv_cache_write_asym3_q8_pair_gfx1100", assemble_asym(kernels::KV_CACHE_WRITE_ASYM3_Q8_PAIR_GFX1100_SRC), ["kv_cache_write_asym3_q8_pair_gfx1100"]); + add!( + "fused_qkvza_hfq4g256_k2048_gfx1100", + kernels::FUSED_QKVZA_HFQ4G256_K2048_GFX1100_SRC, + ["fused_qkvza_hfq4g256_k2048"] + ); + add!( + "fused_gate_up_hfq4g256_stage_x32_gfx1100", + kernels::FUSED_GATE_UP_HFQ4G256_STAGE_X32_GFX1100_SRC, + ["fused_gate_up_hfq4g256_stage_x32_gfx1100"] + ); + add!( + "kv_cache_write_asym3_q8_pair_gfx1100", + assemble_asym(kernels::KV_CACHE_WRITE_ASYM3_Q8_PAIR_GFX1100_SRC), + ["kv_cache_write_asym3_q8_pair_gfx1100"] + ); } if matches!(arch, "gfx906" | "gfx942") { - add!("fused_qkvza_hfq4g256_wave64", kernels::FUSED_QKVZA_HFQ4G256_WAVE64_SRC, ["fused_qkvza_hfq4g256_wave64"]); - add!("fused_qkv_hfq4g256_wave64", kernels::FUSED_QKV_HFQ4G256_WAVE64_SRC, ["fused_qkv_hfq4g256_wave64"]); - add!("fused_gate_up_hfq4g256_wave64", kernels::FUSED_GATE_UP_HFQ4G256_WAVE64_SRC, ["fused_gate_up_hfq4g256_wave64"]); - add!("gemv_hfq4g256_moe_gate_up_indexed_wave64", kernels::GEMV_HFQ4G256_MOE_GATE_UP_INDEXED_WAVE64_SRC, ["gemv_hfq4g256_moe_gate_up_k8_indexed_wave64"]); - add!("gemv_hfq4g256_moe_down_indexed_wave64", kernels::GEMV_HFQ4G256_MOE_DOWN_INDEXED_WAVE64_SRC, ["gemv_hfq4g256_moe_down_residual_scaled_k8_indexed_wave64"]); - add!("gemm_qkvza_hfq4g256_wave64", kernels::GEMM_QKVZA_HFQ4G256_WAVE64_SRC, ["gemm_qkvza_hfq4g256_wave64"]); - add!("gemm_qkv_hfq4g256_wave64", kernels::GEMM_QKV_HFQ4G256_WAVE64_SRC, ["gemm_qkv_hfq4g256_wave64"]); - add!("gemm_hfq4g256_wave64", kernels::GEMM_HFQ4G256_WAVE64_SRC, ["gemm_hfq4g256_wave64"]); - add!("gemm_hfq4g256_residual_wave64", kernels::GEMM_HFQ4G256_RESIDUAL_WAVE64_SRC, ["gemm_hfq4g256_residual_wave64"]); - add!("gemv_hfq4g256_moe_gate_up_indexed_batched_wave64", kernels::GEMV_HFQ4G256_MOE_GATE_UP_INDEXED_BATCHED_WAVE64_SRC, ["gemv_hfq4g256_moe_gate_up_k8_indexed_batched_wave64"]); - add!("gemv_hfq4g256_moe_down_indexed_batched_wave64", kernels::GEMV_HFQ4G256_MOE_DOWN_INDEXED_BATCHED_WAVE64_SRC, ["gemv_hfq4g256_moe_down_residual_scaled_k8_indexed_batched_wave64"]); + add!( + "fused_qkvza_hfq4g256_wave64", + kernels::FUSED_QKVZA_HFQ4G256_WAVE64_SRC, + ["fused_qkvza_hfq4g256_wave64"] + ); + add!( + "fused_qkv_hfq4g256_wave64", + kernels::FUSED_QKV_HFQ4G256_WAVE64_SRC, + ["fused_qkv_hfq4g256_wave64"] + ); + add!( + "fused_gate_up_hfq4g256_wave64", + kernels::FUSED_GATE_UP_HFQ4G256_WAVE64_SRC, + ["fused_gate_up_hfq4g256_wave64"] + ); + add!( + "gemv_hfq4g256_moe_gate_up_indexed_wave64", + kernels::GEMV_HFQ4G256_MOE_GATE_UP_INDEXED_WAVE64_SRC, + ["gemv_hfq4g256_moe_gate_up_k8_indexed_wave64"] + ); + add!( + "gemv_hfq4g256_moe_down_indexed_wave64", + kernels::GEMV_HFQ4G256_MOE_DOWN_INDEXED_WAVE64_SRC, + ["gemv_hfq4g256_moe_down_residual_scaled_k8_indexed_wave64"] + ); + add!( + "gemm_qkvza_hfq4g256_wave64", + kernels::GEMM_QKVZA_HFQ4G256_WAVE64_SRC, + ["gemm_qkvza_hfq4g256_wave64"] + ); + add!( + "gemm_qkv_hfq4g256_wave64", + kernels::GEMM_QKV_HFQ4G256_WAVE64_SRC, + ["gemm_qkv_hfq4g256_wave64"] + ); + add!( + "gemm_hfq4g256_wave64", + kernels::GEMM_HFQ4G256_WAVE64_SRC, + ["gemm_hfq4g256_wave64"] + ); + add!( + "gemm_hfq4g256_residual_wave64", + kernels::GEMM_HFQ4G256_RESIDUAL_WAVE64_SRC, + ["gemm_hfq4g256_residual_wave64"] + ); + add!( + "gemv_hfq4g256_moe_gate_up_indexed_batched_wave64", + kernels::GEMV_HFQ4G256_MOE_GATE_UP_INDEXED_BATCHED_WAVE64_SRC, + ["gemv_hfq4g256_moe_gate_up_k8_indexed_batched_wave64"] + ); + add!( + "gemv_hfq4g256_moe_down_indexed_batched_wave64", + kernels::GEMV_HFQ4G256_MOE_DOWN_INDEXED_BATCHED_WAVE64_SRC, + ["gemv_hfq4g256_moe_down_residual_scaled_k8_indexed_batched_wave64"] + ); } // KV paths: precompile's asym3/q8 assembly mirrors the runtime header // stripping; FP8 substitutions are only admitted for observed gfx1201. - add!("kv_cache_write_asym_k_givens3", assemble_asym(kernels::KV_CACHE_WRITE_ASYM_K_GIVENS3_SRC), ["kv_cache_write_asym_k_givens3"]); - add!("kv_cache_write_asym_k_givens3_batched", assemble_asym(kernels::KV_CACHE_WRITE_ASYM_K_GIVENS3_BATCHED_SRC), ["kv_cache_write_asym_k_givens3_batched"]); - add!("attention_flash_asym3_tile", assemble_asym(kernels::ATTENTION_FLASH_ASYM3_TILE_SRC), ["attention_flash_asym3_tile"]); - add!("attention_flash_asym3_tile_batched", assemble_asym(kernels::ATTENTION_FLASH_ASYM3_TILE_BATCHED_SRC), ["attention_flash_asym3_tile_batched"]); - add!("attention_flash_asym_reduce_batched", kernels::ATTENTION_FLASH_ASYM_REDUCE_BATCHED_SRC, ["attention_flash_asym_reduce_batched"]); - add!("kv_cache_write_q8_0", kernels::KV_CACHE_WRITE_Q8_0_SRC, ["kv_cache_write_q8_0"]); - add!("attention_q8_0_kv", kernels::ATTENTION_Q8_0_KV_SRC, ["attention_q8_0_kv"]); - add!("attention_q8_0_kv_batched", prepend_kv_slot_desc(kernels::ATTENTION_Q8_0_KV_BATCHED_SRC), ["attention_q8_0_kv_batched"]); - add!("attention_q8_0_kv_independent_masked_windowed", prepend_kv_slot_desc(kernels::ATTENTION_Q8_0_KV_BATCHED_SRC), ["attention_q8_0_kv_independent_masked_windowed"]); - add!("attention_q8_0_flash_prefill", prepend_kv_slot_desc(kernels::ATTENTION_Q8_0_FLASH_PREFILL_SRC), ["attention_q8_0_flash_prefill"]); + add!( + "kv_cache_write_asym_k_givens3", + assemble_asym(kernels::KV_CACHE_WRITE_ASYM_K_GIVENS3_SRC), + ["kv_cache_write_asym_k_givens3"] + ); + add!( + "kv_cache_write_asym_k_givens3_batched", + assemble_asym(kernels::KV_CACHE_WRITE_ASYM_K_GIVENS3_BATCHED_SRC), + ["kv_cache_write_asym_k_givens3_batched"] + ); + add!( + "attention_flash_asym3_tile", + assemble_asym(kernels::ATTENTION_FLASH_ASYM3_TILE_SRC), + ["attention_flash_asym3_tile"] + ); + add!( + "attention_flash_asym3_tile_batched", + assemble_asym(kernels::ATTENTION_FLASH_ASYM3_TILE_BATCHED_SRC), + ["attention_flash_asym3_tile_batched"] + ); + add!( + "attention_flash_asym_reduce_batched", + kernels::ATTENTION_FLASH_ASYM_REDUCE_BATCHED_SRC, + ["attention_flash_asym_reduce_batched"] + ); + add!( + "kv_cache_write_q8_0", + kernels::KV_CACHE_WRITE_Q8_0_SRC, + ["kv_cache_write_q8_0"] + ); + add!( + "attention_q8_0_kv", + kernels::ATTENTION_Q8_0_KV_SRC, + ["attention_q8_0_kv"] + ); + add!( + "attention_q8_0_kv_batched", + prepend_kv_slot_desc(kernels::ATTENTION_Q8_0_KV_BATCHED_SRC), + ["attention_q8_0_kv_batched"] + ); + add!( + "attention_q8_0_kv_independent_masked_windowed", + prepend_kv_slot_desc(kernels::ATTENTION_Q8_0_KV_BATCHED_SRC), + ["attention_q8_0_kv_independent_masked_windowed"] + ); + add!( + "attention_q8_0_flash_prefill", + prepend_kv_slot_desc(kernels::ATTENTION_Q8_0_FLASH_PREFILL_SRC), + ["attention_q8_0_flash_prefill"] + ); // The installer-only unspecialized module above cannot satisfy this // runtime's default scalar-prefill module or its BR/BC-specialized source. - add!("attention_q8_0_flash_prefill_br8_bc16", q8_flash_prefill_default_source(), ["attention_q8_0_flash_prefill"]); - add!("kv_cache_write_q8_0_batched", prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC), ["kv_cache_write_q8_0_batched"]); - add!("kv_cache_write_q8_0_independent", prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC), ["kv_cache_write_q8_0_independent"]); - add!("kv_cache_write_q8_0_independent_masked", prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC), ["kv_cache_write_q8_0_independent_masked"]); - add!("attention_flash_q8_0_tile", kernels::ATTENTION_FLASH_Q8_0_TILE_SRC, ["attention_flash_q8_0_tile"]); - add!("attention_flash_q8_0_reduce", kernels::ATTENTION_FLASH_Q8_0_REDUCE_SRC, ["attention_flash_q8_0_reduce"]); + add!( + "attention_q8_0_flash_prefill_br8_bc16", + q8_flash_prefill_default_source(), + ["attention_q8_0_flash_prefill"] + ); + add!( + "kv_cache_write_q8_0_batched", + prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC), + ["kv_cache_write_q8_0_batched"] + ); + add!( + "kv_cache_write_q8_0_independent", + prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC), + ["kv_cache_write_q8_0_independent"] + ); + add!( + "kv_cache_write_q8_0_independent_masked", + prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC), + ["kv_cache_write_q8_0_independent_masked"] + ); + add!( + "attention_flash_q8_0_tile", + kernels::ATTENTION_FLASH_Q8_0_TILE_SRC, + ["attention_flash_q8_0_tile"] + ); + add!( + "attention_flash_q8_0_reduce", + kernels::ATTENTION_FLASH_Q8_0_REDUCE_SRC, + ["attention_flash_q8_0_reduce"] + ); // Multi-slot (slot-descriptor) launches of the q8 and asym3 KV routes run // separately named `*_paged` modules (attention.rs `ensure_kv_slot_kernel`, // `ensure_givens4_kv_slot_kernel`, `attention_q8_0_flash_prefill_wmma_slots`). - add!("kv_cache_write_q8_0_batched_paged", kernels::kv_slot_desc_paged_source(kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC, "kv_cache_write_q8_0_batched", "kv_cache_write_q8_0_batched_paged"), ["kv_cache_write_q8_0_batched_paged"]); - add!("attention_q8_0_kv_batched_paged", kernels::kv_slot_desc_paged_source(kernels::ATTENTION_Q8_0_KV_BATCHED_SRC, "attention_q8_0_kv_batched", "attention_q8_0_kv_batched_paged"), ["attention_q8_0_kv_batched_paged"]); - add!("attention_flash_q8_0_tile_batched_paged", kernels::kv_slot_givens4_paged_source(kernels::ATTENTION_FLASH_Q8_0_TILE_BATCHED_SRC, "attention_flash_q8_0_tile_batched", "attention_flash_q8_0_tile_batched_paged"), ["attention_flash_q8_0_tile_batched_paged"]); - add!("kv_cache_write_asym_k_givens3_batched_paged", kernels::kv_slot_givens4_paged_source(kernels::KV_CACHE_WRITE_ASYM_K_GIVENS3_BATCHED_SRC, "kv_cache_write_asym_k_givens3_batched", "kv_cache_write_asym_k_givens3_batched_paged"), ["kv_cache_write_asym_k_givens3_batched_paged"]); - add!("attention_flash_asym3_tile_batched_paged", kernels::kv_slot_givens4_paged_source(kernels::ATTENTION_FLASH_ASYM3_TILE_BATCHED_SRC, "attention_flash_asym3_tile_batched", "attention_flash_asym3_tile_batched_paged"), ["attention_flash_asym3_tile_batched_paged"]); + add!( + "kv_cache_write_q8_0_batched_paged", + kernels::kv_slot_desc_paged_source( + kernels::KV_CACHE_WRITE_Q8_0_BATCHED_SRC, + "kv_cache_write_q8_0_batched", + "kv_cache_write_q8_0_batched_paged" + ), + ["kv_cache_write_q8_0_batched_paged"] + ); + add!( + "attention_q8_0_kv_batched_paged", + kernels::kv_slot_desc_paged_source( + kernels::ATTENTION_Q8_0_KV_BATCHED_SRC, + "attention_q8_0_kv_batched", + "attention_q8_0_kv_batched_paged" + ), + ["attention_q8_0_kv_batched_paged"] + ); + add!( + "attention_flash_q8_0_tile_batched_paged", + kernels::kv_slot_givens4_paged_source( + kernels::ATTENTION_FLASH_Q8_0_TILE_BATCHED_SRC, + "attention_flash_q8_0_tile_batched", + "attention_flash_q8_0_tile_batched_paged" + ), + ["attention_flash_q8_0_tile_batched_paged"] + ); + add!( + "kv_cache_write_asym_k_givens3_batched_paged", + kernels::kv_slot_givens4_paged_source( + kernels::KV_CACHE_WRITE_ASYM_K_GIVENS3_BATCHED_SRC, + "kv_cache_write_asym_k_givens3_batched", + "kv_cache_write_asym_k_givens3_batched_paged" + ), + ["kv_cache_write_asym_k_givens3_batched_paged"] + ); + add!( + "attention_flash_asym3_tile_batched_paged", + kernels::kv_slot_givens4_paged_source( + kernels::ATTENTION_FLASH_ASYM3_TILE_BATCHED_SRC, + "attention_flash_asym3_tile_batched", + "attention_flash_asym3_tile_batched_paged" + ), + ["attention_flash_asym3_tile_batched_paged"] + ); if matches!(arch, "gfx1100" | "gfx1151") { // gfx11 admits the multi-slot WMMA prefill tile by default // (forward_slots `q8_flash_prefill_wmma_eligible`); Qwen3.5 head_dim // 256 with the default SPLIT_Q=0, fixed head dim and V prefetch. - add!("attention_q8_0_flash_prefill_wmma_gfx11_hd256_paged", format!("#define SPLIT_Q 0\n#define FIXED_HEAD_DIM 256\n#define PREFETCH_V 1\n{}", kernels::kv_slot_desc_paged_source(kernels::ATTENTION_Q8_0_FLASH_PREFILL_WMMA_SRC, "attention_q8_0_flash_prefill_wmma", "attention_q8_0_flash_prefill_wmma_paged")), ["attention_q8_0_flash_prefill_wmma_paged"]); + add!( + "attention_q8_0_flash_prefill_wmma_gfx11_hd256_paged", + format!( + "#define SPLIT_Q 0\n#define FIXED_HEAD_DIM 256\n#define PREFETCH_V 1\n{}", + kernels::kv_slot_desc_paged_source( + kernels::ATTENTION_Q8_0_FLASH_PREFILL_WMMA_SRC, + "attention_q8_0_flash_prefill_wmma", + "attention_q8_0_flash_prefill_wmma_paged" + ) + ), + ["attention_q8_0_flash_prefill_wmma_paged"] + ); } - add!("sample_top_p_parallel", sampling::sample_top_p_parallel_src(), ["sample_apply_repeat_penalty", "sample_topk_partial", "sample_topk_finalize"]); - add!("sample_top_p_parallel_w64", sampling::sample_top_p_parallel_w64_src(), ["sample_apply_repeat_penalty_w64", "sample_topk_partial_w64", "sample_topk_finalize_w64"]); - add!("sample_top_p_parallel_fast21", sampling::sample_top_p_parallel_fast_src(21, "fast21"), ["sample_apply_repeat_penalty_fast21", "sample_topk_partial_fast21", "sample_topk_finalize_fast21"]); - add!("sample_top_p_parallel_fast65", sampling::sample_top_p_parallel_fast_src(65, "fast65"), ["sample_apply_repeat_penalty_fast65", "sample_topk_partial_fast65", "sample_topk_finalize_fast65"]); + add!( + "sample_top_p_parallel", + sampling::sample_top_p_parallel_src(), + [ + "sample_apply_repeat_penalty", + "sample_topk_partial", + "sample_topk_finalize" + ] + ); + add!( + "sample_top_p_parallel_w64", + sampling::sample_top_p_parallel_w64_src(), + [ + "sample_apply_repeat_penalty_w64", + "sample_topk_partial_w64", + "sample_topk_finalize_w64" + ] + ); + add!( + "sample_top_p_parallel_fast21", + sampling::sample_top_p_parallel_fast_src(21, "fast21"), + [ + "sample_apply_repeat_penalty_fast21", + "sample_topk_partial_fast21", + "sample_topk_finalize_fast21" + ] + ); + add!( + "sample_top_p_parallel_fast65", + sampling::sample_top_p_parallel_fast_src(65, "fast65"), + [ + "sample_apply_repeat_penalty_fast65", + "sample_topk_partial_fast65", + "sample_topk_finalize_fast65" + ] + ); if arch == "gfx1201" { // H2 and small-model first-token JIT: use the actual module key, even // when it differs from the kernel symbol or the source filename. - add!("attention_flash_fp8_e4m3_tile", kernels::ATTENTION_FLASH_FP8_E4M3_TILE_SRC, ["attention_flash_fp8_e4m3_tile"]); - add!("attention_fp8_e4m3_kv_batched", prepend_kv_slot_desc(kernels::ATTENTION_FP8_E4M3_KV_BATCHED_SRC), ["attention_fp8_e4m3_kv_batched"]); - add!("conv1d_silu_split_qknorm_b256", kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SRC, ["conv1d_silu_split_qknorm_b256"]); - add!("convert_f32_to_f16", kernels::GEMM_HFQ4G256_RESIDUAL_FP16_SRC, ["convert_f32_to_f16"]); - add!("fused_gate_up_hfq4g256_mq4v2", kernels::FUSED_GATE_UP_MQ4G256V2_SRC, ["fused_gate_up_mq4g256v2"]); - add!("fused_qkv_hfq4g256_mq4v2", kernels::FUSED_QKV_MQ4G256V2_SRC, ["fused_qkv_mq4g256v2"]); - add!("fused_qkvza_hfq4g256_mq4v2", kernels::FUSED_QKVZA_MQ4G256V2_SRC, ["fused_qkvza_mq4g256v2"]); - add!("fused_rmsnorm_mq_rotate_awq", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_SRC, ["fused_rmsnorm_mq_rotate_awq"]); - add!("fused_rmsnorm_mq_rotate_awq_g12dec", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_G12DEC_SRC, ["fused_rmsnorm_mq_rotate_awq_g12dec"]); - add!("fused_silu_mul_mq_rotate_awq", kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_SRC, ["fused_silu_mul_mq_rotate_awq"]); - add!("gated_delta_net_q8_fast", kernels::GATED_DELTA_NET_Q8_FAST_SRC, ["gated_delta_net_q8_fast"]); - add!("gdn_pre_batched_gfx1201", include_str!("../../../kernels/src/gdn_pre_batched.gfx1201.hip"), ["gdn_pre_batched_gfx1201"]); - add!("gemm_gate_up_hfq4g256_wmma_gfx12_mq4v2", kernels::GEMM_GATE_UP_MQ4G256V2_WMMA_GFX12_SRC, ["gemm_gate_up_mq4g256v2_wmma_gfx12"]); - add!("gemm_hfq4g256_residual_wmma_gfx12_mq4v2", kernels::GEMM_MQ4G256V2_RESIDUAL_WMMA_GFX12_SRC, ["gemm_mq4g256v2_residual_wmma_gfx12"]); - add!("gemm_qkv_hfq4g256_wmma_gfx12_mq4v2", kernels::GEMM_QKV_MQ4G256V2_WMMA_GFX12_SRC, ["gemm_qkv_mq4g256v2_wmma_gfx12"]); - add!("gemm_qkvza_hfq4g256_wmma_gfx12_mq4v2", kernels::GEMM_QKVZA_MQ4G256V2_WMMA_GFX12_SRC, ["gemm_qkvza_mq4g256v2_wmma_gfx12"]); - add!("gemv_hfq4g256_multirow_default_mq4v2", kernels::GEMV_MQ4G256V2_MULTIROW_SRC, ["gemv_mq4g256v2_multirow_r2", "gemv_mq4g256v2_multirow_r4", "gemv_mq4g256v2_multirow_r8"]); - add!("gemv_hfq4g256_residual_mq4v2", kernels::GEMV_MQ4G256V2_RESIDUAL_SRC, ["gemv_mq4g256v2_residual"]); - add!("gemv_mq4g256v2_mq4v2", kernels::GEMV_MQ4G256V2_SRC, ["gemv_mq4g256v2"]); - add!("kv_cache_write_fp8_e4m3", kernels::KV_CACHE_WRITE_FP8_E4M3_SRC, ["kv_cache_write_fp8_e4m3"]); - add!("kv_cache_write_fp8_e4m3_batched", prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_FP8_E4M3_BATCHED_SRC), ["kv_cache_write_fp8_e4m3_batched"]); + add!( + "attention_flash_fp8_e4m3_tile", + kernels::ATTENTION_FLASH_FP8_E4M3_TILE_SRC, + ["attention_flash_fp8_e4m3_tile"] + ); + add!( + "attention_fp8_e4m3_kv_batched", + prepend_kv_slot_desc(kernels::ATTENTION_FP8_E4M3_KV_BATCHED_SRC), + ["attention_fp8_e4m3_kv_batched"] + ); + add!( + "conv1d_silu_split_qknorm_b256", + kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SRC, + ["conv1d_silu_split_qknorm_b256"] + ); + add!( + "convert_f32_to_f16", + kernels::GEMM_HFQ4G256_RESIDUAL_FP16_SRC, + ["convert_f32_to_f16"] + ); + add!( + "fused_gate_up_hfq4g256_mq4v2", + kernels::FUSED_GATE_UP_MQ4G256V2_SRC, + ["fused_gate_up_mq4g256v2"] + ); + add!( + "fused_qkv_hfq4g256_mq4v2", + kernels::FUSED_QKV_MQ4G256V2_SRC, + ["fused_qkv_mq4g256v2"] + ); + add!( + "fused_qkvza_hfq4g256_mq4v2", + kernels::FUSED_QKVZA_MQ4G256V2_SRC, + ["fused_qkvza_mq4g256v2"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_SRC, + ["fused_rmsnorm_mq_rotate_awq"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_g12dec", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_G12DEC_SRC, + ["fused_rmsnorm_mq_rotate_awq_g12dec"] + ); + add!( + "fused_silu_mul_mq_rotate_awq", + kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_SRC, + ["fused_silu_mul_mq_rotate_awq"] + ); + add!( + "gated_delta_net_q8_fast", + kernels::GATED_DELTA_NET_Q8_FAST_SRC, + ["gated_delta_net_q8_fast"] + ); + add!( + "gdn_pre_batched_gfx1201", + include_str!("../../../kernels/src/gdn_pre_batched.gfx1201.hip"), + ["gdn_pre_batched_gfx1201"] + ); + add!( + "gemm_gate_up_hfq4g256_wmma_gfx12_mq4v2", + kernels::GEMM_GATE_UP_MQ4G256V2_WMMA_GFX12_SRC, + ["gemm_gate_up_mq4g256v2_wmma_gfx12"] + ); + add!( + "gemm_hfq4g256_residual_wmma_gfx12_mq4v2", + kernels::GEMM_MQ4G256V2_RESIDUAL_WMMA_GFX12_SRC, + ["gemm_mq4g256v2_residual_wmma_gfx12"] + ); + add!( + "gemm_qkv_hfq4g256_wmma_gfx12_mq4v2", + kernels::GEMM_QKV_MQ4G256V2_WMMA_GFX12_SRC, + ["gemm_qkv_mq4g256v2_wmma_gfx12"] + ); + add!( + "gemm_qkvza_hfq4g256_wmma_gfx12_mq4v2", + kernels::GEMM_QKVZA_MQ4G256V2_WMMA_GFX12_SRC, + ["gemm_qkvza_mq4g256v2_wmma_gfx12"] + ); + add!( + "gemv_hfq4g256_multirow_default_mq4v2", + kernels::GEMV_MQ4G256V2_MULTIROW_SRC, + [ + "gemv_mq4g256v2_multirow_r2", + "gemv_mq4g256v2_multirow_r4", + "gemv_mq4g256v2_multirow_r8" + ] + ); + add!( + "gemv_hfq4g256_residual_mq4v2", + kernels::GEMV_MQ4G256V2_RESIDUAL_SRC, + ["gemv_mq4g256v2_residual"] + ); + add!( + "gemv_mq4g256v2_mq4v2", + kernels::GEMV_MQ4G256V2_SRC, + ["gemv_mq4g256v2"] + ); + add!( + "kv_cache_write_fp8_e4m3", + kernels::KV_CACHE_WRITE_FP8_E4M3_SRC, + ["kv_cache_write_fp8_e4m3"] + ); + add!( + "kv_cache_write_fp8_e4m3_batched", + prepend_kv_slot_desc(kernels::KV_CACHE_WRITE_FP8_E4M3_BATCHED_SRC), + ["kv_cache_write_fp8_e4m3_batched"] + ); add!("mq_rotate_x", kernels::GEMV_MQ4G256_SRC, ["mq_rotate_x"]); - add!("qwen35_fa_prep_batched_gfx1201", include_str!("../../../kernels/src/qwen35_fa_prep_batched.gfx1201.hip"), ["qwen35_fa_prep_batched_gfx1201"]); + add!( + "qwen35_fa_prep_batched_gfx1201", + include_str!("../../../kernels/src/qwen35_fa_prep_batched.gfx1201.hip"), + ["qwen35_fa_prep_batched_gfx1201"] + ); add!("rmsnorm_f32", kernels::RMSNORM_SRC, ["rmsnorm_f32"]); - add!("rmsnorm_f32_rowsplit", kernels::RMSNORM_ROWSPLIT_SRC, ["rmsnorm_f32_rowsplit"]); - add!("rope_partial_halfsplit", kernels::ROPE_PARTIAL_HALFSPLIT_SRC, ["rope_partial_halfsplit_f32"]); - add!("rope_partial_halfsplit_f32_headgrid", kernels::ROPE_PARTIAL_HALFSPLIT_HEADGRID_SRC, ["rope_partial_halfsplit_f32_headgrid"]); - add!("rotate_x_mq_awq", kernels::ROTATE_X_MQ_AWQ_SRC, ["rotate_x_mq_awq"]); - add!("deinterleave_q_rmsnorm_f32_batched", kernels::DEINTERLEAVE_Q_RMSNORM_BATCHED_SRC, ["deinterleave_q_rmsnorm_f32_batched"]); - add!("fused_gate_up_hfq4g256_k1024_gfx1201", kernels::FUSED_GATE_UP_HFQ4G256_K1024_GFX1201_SRC, ["fused_gate_up_hfq4g256_k1024_gfx1201"]); - add!("gemm_gate_up_hfq4g256_wmma_gfx12", kernels::GEMM_GATE_UP_HFQ4G256_WMMA_GFX12_SRC, ["gemm_gate_up_hfq4g256_wmma_gfx12"]); - add!("gemm_hfq4g256_residual_wmma_gfx12", kernels::GEMM_HFQ4G256_RESIDUAL_WMMA_GFX12_SRC, ["gemm_hfq4g256_residual_wmma_gfx12"]); - add!("gemm_qkv_hfq4g256_wmma_gfx12", kernels::GEMM_QKV_HFQ4G256_WMMA_GFX12_SRC, ["gemm_qkv_hfq4g256_wmma_gfx12"]); - add!("gemm_qkvza_hfq4g256_wmma_gfx12", kernels::GEMM_QKVZA_HFQ4G256_WMMA_GFX12_SRC, ["gemm_qkvza_hfq4g256_wmma_gfx12"]); - add!("gemv_hfq4g256_residual", kernels::GEMV_HFQ4G256_RESIDUAL_SRC, ["gemv_hfq4g256_residual"]); - add!("gemv_q8_0_wide", kernels::GEMV_Q8_0_WIDE_SRC, ["gemv_q8_0_wide"]); - add!("rope_partial_halfsplit_batched", kernels::ROPE_PARTIAL_HALFSPLIT_BATCHED_SRC, ["rope_partial_halfsplit_batched_f32"]); + add!( + "rmsnorm_f32_rowsplit", + kernels::RMSNORM_ROWSPLIT_SRC, + ["rmsnorm_f32_rowsplit"] + ); + add!( + "rope_partial_halfsplit", + kernels::ROPE_PARTIAL_HALFSPLIT_SRC, + ["rope_partial_halfsplit_f32"] + ); + add!( + "rope_partial_halfsplit_f32_headgrid", + kernels::ROPE_PARTIAL_HALFSPLIT_HEADGRID_SRC, + ["rope_partial_halfsplit_f32_headgrid"] + ); + add!( + "rotate_x_mq_awq", + kernels::ROTATE_X_MQ_AWQ_SRC, + ["rotate_x_mq_awq"] + ); + add!( + "deinterleave_q_rmsnorm_f32_batched", + kernels::DEINTERLEAVE_Q_RMSNORM_BATCHED_SRC, + ["deinterleave_q_rmsnorm_f32_batched"] + ); + add!( + "fused_gate_up_hfq4g256_k1024_gfx1201", + kernels::FUSED_GATE_UP_HFQ4G256_K1024_GFX1201_SRC, + ["fused_gate_up_hfq4g256_k1024_gfx1201"] + ); + add!( + "gemm_gate_up_hfq4g256_wmma_gfx12", + kernels::GEMM_GATE_UP_HFQ4G256_WMMA_GFX12_SRC, + ["gemm_gate_up_hfq4g256_wmma_gfx12"] + ); + add!( + "gemm_hfq4g256_residual_wmma_gfx12", + kernels::GEMM_HFQ4G256_RESIDUAL_WMMA_GFX12_SRC, + ["gemm_hfq4g256_residual_wmma_gfx12"] + ); + add!( + "gemm_qkv_hfq4g256_wmma_gfx12", + kernels::GEMM_QKV_HFQ4G256_WMMA_GFX12_SRC, + ["gemm_qkv_hfq4g256_wmma_gfx12"] + ); + add!( + "gemm_qkvza_hfq4g256_wmma_gfx12", + kernels::GEMM_QKVZA_HFQ4G256_WMMA_GFX12_SRC, + ["gemm_qkvza_hfq4g256_wmma_gfx12"] + ); + add!( + "gemv_hfq4g256_residual", + kernels::GEMV_HFQ4G256_RESIDUAL_SRC, + ["gemv_hfq4g256_residual"] + ); + add!( + "gemv_q8_0_wide", + kernels::GEMV_Q8_0_WIDE_SRC, + ["gemv_q8_0_wide"] + ); + add!( + "rope_partial_halfsplit_batched", + kernels::ROPE_PARTIAL_HALFSPLIT_BATCHED_SRC, + ["rope_partial_halfsplit_batched_f32"] + ); // H2 (Qwen3.8-27B MQ4v2, fp8 KV) decode, observed JIT-compiled by // `hipfire run` greedy AR/MTP/DFlash on the installer inventory above. - add!("attention_flash_fp8_e4m3_tile_gqa_gfx1201", kernels::ATTENTION_FLASH_FP8_E4M3_TILE_GQA_GFX1201_SRC, ["attention_flash_fp8_e4m3_tile_gqa_gfx1201"]); - add!("attention_flash_reduce_dsplit_gfx1201", kernels::ATTENTION_FLASH_REDUCE_DSPLIT_GFX1201_SRC, ["attention_flash_reduce_dsplit_gfx1201"]); - add!("dflash_state_bulk_copy_gfx1100", crate::dflash_state_copy::DFLASH_STATE_BULK_COPY_GFX1100_SRC, ["dflash_state_bulk_copy_gfx1100"]); - add!("gated_delta_net_q8_compact3_b2", kernels::GATED_DELTA_NET_Q8_COMPACT3_B2_SRC, ["gated_delta_net_q8_compact3_b2"]); - add!("gated_norm_mq_rotate_awq_k6144_gfx1201", kernels::gated_norm_mq_rotate_awq_k6144_gfx1201_src(), ["gated_norm_mq_rotate_awq_k6144_gfx1201"]); - add!("qwen36_27b_fa_prep_gfx1201", kernels::qwen36_27b_fa_prep_gfx1201_src(), ["qwen36_27b_fa_prep_gfx1201"]); + add!( + "attention_flash_fp8_e4m3_tile_gqa_gfx1201", + kernels::ATTENTION_FLASH_FP8_E4M3_TILE_GQA_GFX1201_SRC, + ["attention_flash_fp8_e4m3_tile_gqa_gfx1201"] + ); + add!( + "attention_flash_reduce_dsplit_gfx1201", + kernels::ATTENTION_FLASH_REDUCE_DSPLIT_GFX1201_SRC, + ["attention_flash_reduce_dsplit_gfx1201"] + ); + add!( + "dflash_state_bulk_copy_gfx1100", + crate::dflash_state_copy::DFLASH_STATE_BULK_COPY_GFX1100_SRC, + ["dflash_state_bulk_copy_gfx1100"] + ); + add!( + "gated_delta_net_q8_compact3_b2", + kernels::GATED_DELTA_NET_Q8_COMPACT3_B2_SRC, + ["gated_delta_net_q8_compact3_b2"] + ); + add!( + "gated_norm_mq_rotate_awq_k6144_gfx1201", + kernels::gated_norm_mq_rotate_awq_k6144_gfx1201_src(), + ["gated_norm_mq_rotate_awq_k6144_gfx1201"] + ); + add!( + "qwen36_27b_fa_prep_gfx1201", + kernels::qwen36_27b_fa_prep_gfx1201_src(), + ["qwen36_27b_fa_prep_gfx1201"] + ); } // H2 (Qwen3.8-27B MQ4v2) prefill, MTP and DFlash modules JIT-compiled by // `hipfire serve`/`run` beyond the inventory above: gfx1201 fp8 KV, and // gfx1100/gfx1151 Q8 KV. Symbols are every kernel each object defines. if matches!(arch, "gfx1201" | "gfx1100" | "gfx1151") { - add!("attention_dflash_sliding_f32", kernels::ATTENTION_DFLASH_SLIDING_SRC, ["attention_dflash_sliding_f32"]); - add!("deinterleave_batched", kernels::DEINTERLEAVE_BATCHED_SRC, ["deinterleave_f32_batched"]); - add!("dynamic_conv_f32", kernels::DYNAMIC_CONV_F32_SRC, ["dynamic_causal_conv_f32", "dynamic_conv_f32"]); - add!("fused_qk_l2_norm_scale_interleave_f32_batched", kernels::FUSED_QK_L2_NORM_SCALE_INTERLEAVE_F32_BATCHED_SRC, ["fused_qk_l2_norm_scale_interleave_f32_batched"]); - add!("gemm_hfq4g256", kernels::GEMM_HFQ4G256_SRC, ["gemm_hfq4g256"]); - add!("rope_batched", kernels::ROPE_BATCHED_SRC, ["rope_batched_f32"]); + add!( + "attention_dflash_sliding_f32", + kernels::ATTENTION_DFLASH_SLIDING_SRC, + ["attention_dflash_sliding_f32"] + ); + add!( + "deinterleave_batched", + kernels::DEINTERLEAVE_BATCHED_SRC, + ["deinterleave_f32_batched"] + ); + add!( + "dynamic_conv_f32", + kernels::DYNAMIC_CONV_F32_SRC, + ["dynamic_causal_conv_f32", "dynamic_conv_f32"] + ); + add!( + "fused_qk_l2_norm_scale_interleave_f32_batched", + kernels::FUSED_QK_L2_NORM_SCALE_INTERLEAVE_F32_BATCHED_SRC, + ["fused_qk_l2_norm_scale_interleave_f32_batched"] + ); + add!( + "gemm_hfq4g256", + kernels::GEMM_HFQ4G256_SRC, + ["gemm_hfq4g256"] + ); + add!( + "rope_batched", + kernels::ROPE_BATCHED_SRC, + ["rope_batched_f32"] + ); } if matches!(arch, "gfx1201" | "gfx1151") { - add!("add", kernels::ADD_SRC, ["add_f32", "broadcast_add_rows_f32"]); - add!("argmax_token_chain", kernels::ARGMAX_TOKEN_CHAIN_SRC, ["argmax_token_chain_f32"]); - add!("gdn_chunk_kkt_solve", kernels::GDN_CHUNK_KKT_SOLVE_SRC, ["gdn_chunk_kkt_solve"]); - add!("gemv_hfq4g256_multirow_default", kernels::GEMV_HFQ4G256_MULTIROW_SRC, ["gemv_hfq4g256_multirow_r2", "gemv_hfq4g256_multirow_r4", "gemv_hfq4g256_multirow_r8"]); - add!("greedy_accept", kernels::GREEDY_ACCEPT_SRC, ["greedy_accept_from_argmax_i32"]); + add!( + "add", + kernels::ADD_SRC, + ["add_f32", "broadcast_add_rows_f32"] + ); + add!( + "argmax_token_chain", + kernels::ARGMAX_TOKEN_CHAIN_SRC, + ["argmax_token_chain_f32"] + ); + add!( + "gdn_chunk_kkt_solve", + kernels::GDN_CHUNK_KKT_SOLVE_SRC, + ["gdn_chunk_kkt_solve"] + ); + add!( + "gemv_hfq4g256_multirow_default", + kernels::GEMV_HFQ4G256_MULTIROW_SRC, + [ + "gemv_hfq4g256_multirow_r2", + "gemv_hfq4g256_multirow_r4", + "gemv_hfq4g256_multirow_r8" + ] + ); + add!( + "greedy_accept", + kernels::GREEDY_ACCEPT_SRC, + ["greedy_accept_from_argmax_i32"] + ); } if matches!(arch, "gfx1100" | "gfx1151") { - add!("argmax_f32_batched", kernels::ARGMAX_BATCHED_SRC, ["argmax_f32_batched"]); - add!("attention_q8_0_flash_prefill_wmma_gfx11_hd256", q8_flash_prefill_wmma_gfx11_hd256_source(), ["attention_q8_0_flash_prefill_wmma"]); - add!("convert_f32_to_f16", kernels::GEMM_HFQ4G256_RESIDUAL_FP16_SRC, ["convert_f32_to_f16", "gemm_hfq4g256_residual_fp16"]); - add!("fused_gate_up_hfq4g256_mq4v2", kernels::FUSED_GATE_UP_MQ4G256V2_SRC, ["fused_gate_up_mq4g256v2"]); - add!("fused_qkv_hfq4g256_mq4v2", kernels::FUSED_QKV_MQ4G256V2_SRC, ["fused_qkv_mq4g256v2"]); - add!("fused_qkvza_hfq4g256_mq4v2", kernels::FUSED_QKVZA_MQ4G256V2_SRC, ["fused_qkvza_mq4g256v2"]); - add!("fused_rmsnorm_mq_rotate_awq", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_SRC, ["fused_rmsnorm_mq_rotate_awq"]); - add!("fused_rmsnorm_mq_rotate_awq_g12dec", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_G12DEC_SRC, ["fused_rmsnorm_mq_rotate_awq_g12dec"]); - add!("fused_silu_mul_mq_rotate_awq", kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_SRC, ["fused_silu_mul_mq_rotate_awq"]); - add!("fused_silu_mul_mq_rotate_awq_i4", kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_I4_SRC, ["fused_silu_mul_mq_rotate_awq_i4"]); - add!("fused_silu_mul_mq_rotate_awq_i4_hin", kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_I4_HIN_SRC, ["fused_silu_mul_mq_rotate_awq_i4_hin"]); - add!("gated_delta_net_q8_fast", kernels::GATED_DELTA_NET_Q8_FAST_SRC, ["gated_delta_net_q8_fast", "gated_delta_net_q8_fast_independent_masked"]); - add!("gdn_chunk_prep_gfx11", kernels::GDN_CHUNK_PREP_GFX11_SRC, ["gdn_chunk_prep_gfx11"]); - add!("gemm_gate_up_mq4g256v2_wmma", kernels::GEMM_GATE_UP_MQ4G256V2_WMMA_SRC, ["gemm_gate_up_mq4g256v2_wmma"]); - add!("gemm_mq4g256v2_residual_wmma", kernels::GEMM_MQ4G256V2_RESIDUAL_WMMA_SRC, ["gemm_mq4g256v2_residual_wmma"]); - add!("gemm_qkv_mq4g256v2_wmma", kernels::GEMM_QKV_MQ4G256V2_WMMA_SRC, ["gemm_qkv_mq4g256v2_wmma"]); - add!("gemm_qkvza_mq4g256v2_wmma", kernels::GEMM_QKVZA_MQ4G256V2_WMMA_SRC, ["gemm_qkvza_mq4g256v2_wmma"]); - add!("mq_rotate_x", kernels::GEMV_MQ4G256_SRC, ["gemv_mq4g256", "mq_rotate_x"]); - add!("rmsnorm_f32_rowsplit", kernels::RMSNORM_ROWSPLIT_SRC, ["rmsnorm_f32_rowsplit"]); - add!("rope_partial_halfsplit_batched", kernels::ROPE_PARTIAL_HALFSPLIT_BATCHED_SRC, ["rope_partial_halfsplit_batched_f32"]); - add!("rotate_x_mq_awq", kernels::ROTATE_X_MQ_AWQ_SRC, ["rotate_x_mq_awq"]); - add!("sigmoid_mul_rotate_x_mq_awq_i4_gfx11", kernels::SIGMOID_MUL_MQ_ROTATE_X_AWQ_I4_GFX11_SRC, ["sigmoid_mul_rotate_x_mq_awq_i4_gfx11"]); - add!("softmax_temp_topp_batched", kernels::SOFTMAX_TEMP_BATCHED_SRC, ["softmax_temp_batched_f32", "softmax_temp_topp_batched_f32"]); - add!("split_mq4v2_z_betaalpha", kernels::SPLIT_MQ4V2_Z_BETAALPHA_SRC, ["split_mq4v2_z_betaalpha"]); - add!("topk_logsumexp_batched", kernels::TOPK_LOGSUMEXP_BATCHED_SRC, ["topk_logsumexp_batched_f32"]); + add!( + "argmax_f32_batched", + kernels::ARGMAX_BATCHED_SRC, + ["argmax_f32_batched"] + ); + add!( + "attention_q8_0_flash_prefill_wmma_gfx11_hd256", + q8_flash_prefill_wmma_gfx11_hd256_source(), + ["attention_q8_0_flash_prefill_wmma"] + ); + add!( + "convert_f32_to_f16", + kernels::GEMM_HFQ4G256_RESIDUAL_FP16_SRC, + ["convert_f32_to_f16", "gemm_hfq4g256_residual_fp16"] + ); + add!( + "fused_gate_up_hfq4g256_mq4v2", + kernels::FUSED_GATE_UP_MQ4G256V2_SRC, + ["fused_gate_up_mq4g256v2"] + ); + add!( + "fused_qkv_hfq4g256_mq4v2", + kernels::FUSED_QKV_MQ4G256V2_SRC, + ["fused_qkv_mq4g256v2"] + ); + add!( + "fused_qkvza_hfq4g256_mq4v2", + kernels::FUSED_QKVZA_MQ4G256V2_SRC, + ["fused_qkvza_mq4g256v2"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_SRC, + ["fused_rmsnorm_mq_rotate_awq"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_g12dec", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_G12DEC_SRC, + ["fused_rmsnorm_mq_rotate_awq_g12dec"] + ); + add!( + "fused_silu_mul_mq_rotate_awq", + kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_SRC, + ["fused_silu_mul_mq_rotate_awq"] + ); + add!( + "fused_silu_mul_mq_rotate_awq_i4", + kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_I4_SRC, + ["fused_silu_mul_mq_rotate_awq_i4"] + ); + add!( + "fused_silu_mul_mq_rotate_awq_i4_hin", + kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_I4_HIN_SRC, + ["fused_silu_mul_mq_rotate_awq_i4_hin"] + ); + add!( + "gated_delta_net_q8_fast", + kernels::GATED_DELTA_NET_Q8_FAST_SRC, + [ + "gated_delta_net_q8_fast", + "gated_delta_net_q8_fast_independent_masked" + ] + ); + add!( + "gdn_chunk_prep_gfx11", + kernels::GDN_CHUNK_PREP_GFX11_SRC, + ["gdn_chunk_prep_gfx11"] + ); + add!( + "gemm_gate_up_mq4g256v2_wmma", + kernels::GEMM_GATE_UP_MQ4G256V2_WMMA_SRC, + ["gemm_gate_up_mq4g256v2_wmma"] + ); + add!( + "gemm_mq4g256v2_residual_wmma", + kernels::GEMM_MQ4G256V2_RESIDUAL_WMMA_SRC, + ["gemm_mq4g256v2_residual_wmma"] + ); + add!( + "gemm_qkv_mq4g256v2_wmma", + kernels::GEMM_QKV_MQ4G256V2_WMMA_SRC, + ["gemm_qkv_mq4g256v2_wmma"] + ); + add!( + "gemm_qkvza_mq4g256v2_wmma", + kernels::GEMM_QKVZA_MQ4G256V2_WMMA_SRC, + ["gemm_qkvza_mq4g256v2_wmma"] + ); + add!( + "mq_rotate_x", + kernels::GEMV_MQ4G256_SRC, + ["gemv_mq4g256", "mq_rotate_x"] + ); + add!( + "rmsnorm_f32_rowsplit", + kernels::RMSNORM_ROWSPLIT_SRC, + ["rmsnorm_f32_rowsplit"] + ); + add!( + "rope_partial_halfsplit_batched", + kernels::ROPE_PARTIAL_HALFSPLIT_BATCHED_SRC, + ["rope_partial_halfsplit_batched_f32"] + ); + add!( + "rotate_x_mq_awq", + kernels::ROTATE_X_MQ_AWQ_SRC, + ["rotate_x_mq_awq"] + ); + add!( + "sigmoid_mul_rotate_x_mq_awq_i4_gfx11", + kernels::SIGMOID_MUL_MQ_ROTATE_X_AWQ_I4_GFX11_SRC, + ["sigmoid_mul_rotate_x_mq_awq_i4_gfx11"] + ); + add!( + "softmax_temp_topp_batched", + kernels::SOFTMAX_TEMP_BATCHED_SRC, + ["softmax_temp_batched_f32", "softmax_temp_topp_batched_f32"] + ); + add!( + "split_mq4v2_z_betaalpha", + kernels::SPLIT_MQ4V2_Z_BETAALPHA_SRC, + ["split_mq4v2_z_betaalpha"] + ); + add!( + "topk_logsumexp_batched", + kernels::TOPK_LOGSUMEXP_BATCHED_SRC, + ["topk_logsumexp_batched_f32"] + ); } if arch == "gfx1201" { - add!("attention_flash_fp8_e4m3_tile_batched", assemble_asym(kernels::ATTENTION_FLASH_FP8_E4M3_TILE_BATCHED_SRC), ["attention_flash_fp8_e4m3_tile_batched", "attention_flash_q8_0_tile_batched"]); - add!("attention_fp8_e4m3_fa2_gqa_qresident_v2_gfx1201", kernels::ATTENTION_FP8_E4M3_FA2_GQA_QRESIDENT_V2_GFX1201_SRC, ["attention_fp8_e4m3_fa2_gqa_gfx1201", "attention_fp8_e4m3_fa2_gqa_merge_gfx1201", "attention_fp8_e4m3_fa2_gqa_partial_gfx1201", "attention_fp8_e4m3_fa2_gqa_qresident_v2_gfx1201", "attention_fp8_e4m3_fa2_q_preconvert_f16_gfx1201", "attention_fp8_e4m3_fa2_q_preconvert_fp8_gfx1201"]); - add!("attention_fp8_e4m3_fa2_gqa_qresident_v2_q8_a4epi_gfx1201", kernels::ATTENTION_FP8_E4M3_FA2_GQA_QRESIDENT_V2_Q8_A4EPI_GFX1201_SRC, ["attention_fp8_e4m3_fa2_gqa_gfx1201", "attention_fp8_e4m3_fa2_gqa_merge_gfx1201", "attention_fp8_e4m3_fa2_gqa_partial_gfx1201", "attention_fp8_e4m3_fa2_gqa_qresident_v2_q8_a4epi_gfx1201", "attention_fp8_e4m3_fa2_q_preconvert_f16_gfx1201", "attention_fp8_e4m3_fa2_q_preconvert_fp8_gfx1201"]); - add!("dflash_gdn_replay_pre_ml", crate::dflash_gdn_replay::DFLASH_GDN_REPLAY_PRE_ML_SRC, ["dflash_gdn_replay_pre_ml"]); - add!("fused_rmsnorm_mq_rotate_awq_i4_gfx12_v2_slab_fdiv", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_GFX12_V2_SLAB_FDIV_SRC, ["fused_rmsnorm_mq_rotate_awq_i4_gfx12_v2_slab_fdiv"]); - add!("fused_silu_mul_mq_rotate_awq_i4_hin_gfx12_slab_tokfast", kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_I4_HIN_GFX12_SLAB_TOKFAST_SRC, ["fused_silu_mul_mq_rotate_awq_i4_hin_gfx12_slab_tokfast"]); - add!("gated_delta_net_q8_fast_ml", crate::dflash_gdn_replay::GATED_DELTA_NET_Q8_FAST_ML_SRC, ["gated_delta_net_q8_fast_ml"]); - add!("gated_norm_mq_rotate_awq_i4_gfx12_v2_slab", kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX12_V2_SLAB_SRC, ["gated_norm_mq_rotate_awq_i4_gfx12_v2_slab"]); - add!("gated_norm_mq_rotate_awq_i4_gfx12_v2_xbf16_slab", kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX12_V2_XBF16_SLAB_SRC, ["gated_norm_mq_rotate_awq_i4_gfx12_v2_xbf16_slab"]); - add!("gdn_chunk_kkt_solve_batched", kernels::GDN_CHUNK_KKT_SOLVE_BATCHED_SRC, ["gdn_chunk_kkt_solve_batched"]); - add!("gdn_chunk_prep", kernels::GDN_CHUNK_PREP_SRC, ["gdn_chunk_prep"]); - add!("gdn_chunk_prep_fixup", kernels::GDN_CHUNK_PREP_FIXUP_SRC, ["gdn_chunk_prep_fixup"]); - add!("gdn_chunk_scan_bf16", kernels::GDN_CHUNK_SCAN_BF16_SRC, ["gdn_chunk_scan_bf16"]); - add!("gdn_chunk_scan_bf16_mseg", kernels::GDN_CHUNK_SCAN_BF16_MSEG_SRC, ["gdn_chunk_scan_bf16_mseg"]); - add!("gemm_mq4g256v2_residual_mmq_iu4", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_SRC, ["gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_full_add", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "quantize_int4_mmq_ds128"]); - add!("qwen35_fa_prep_fp8q_nogate_batched_gfx1201", crate::qwen35_fa_batch::FA_PREP_BATCHED_GFX1201_SRC, ["qwen35_fa_prep_batched_gfx1201", "qwen35_fa_prep_fp8q_batched_gfx1201", "qwen35_fa_prep_fp8q_nogate_batched_gfx1201"]); - add!("select_regrid", crate::select_regrid::SELECT_REGRID_SRC, ["argmax_f32_batched_regrid", "topk_values_regrid_f32"]); - add!("sigmoid_mul_rotate_x_mq_awq_i4_gfx12_slab", kernels::SIGMOID_MUL_MQ_ROTATE_X_AWQ_I4_GFX12_SLAB_SRC, ["sigmoid_mul_rotate_x_mq_awq_i4_gfx12_slab"]); - add!("topk_values_fixup_f32", crate::select_regrid::TOPK_VALUES_FIXUP_SRC, ["topk_values_fixup_f32"]); + add!( + "attention_flash_fp8_e4m3_tile_batched", + assemble_asym(kernels::ATTENTION_FLASH_FP8_E4M3_TILE_BATCHED_SRC), + [ + "attention_flash_fp8_e4m3_tile_batched", + "attention_flash_q8_0_tile_batched" + ] + ); + add!( + "attention_fp8_e4m3_fa2_gqa_qresident_v2_gfx1201", + kernels::ATTENTION_FP8_E4M3_FA2_GQA_QRESIDENT_V2_GFX1201_SRC, + [ + "attention_fp8_e4m3_fa2_gqa_gfx1201", + "attention_fp8_e4m3_fa2_gqa_merge_gfx1201", + "attention_fp8_e4m3_fa2_gqa_partial_gfx1201", + "attention_fp8_e4m3_fa2_gqa_qresident_v2_gfx1201", + "attention_fp8_e4m3_fa2_q_preconvert_f16_gfx1201", + "attention_fp8_e4m3_fa2_q_preconvert_fp8_gfx1201" + ] + ); + add!( + "attention_fp8_e4m3_fa2_gqa_qresident_v2_q8_a4epi_gfx1201", + kernels::ATTENTION_FP8_E4M3_FA2_GQA_QRESIDENT_V2_Q8_A4EPI_GFX1201_SRC, + [ + "attention_fp8_e4m3_fa2_gqa_gfx1201", + "attention_fp8_e4m3_fa2_gqa_merge_gfx1201", + "attention_fp8_e4m3_fa2_gqa_partial_gfx1201", + "attention_fp8_e4m3_fa2_gqa_qresident_v2_q8_a4epi_gfx1201", + "attention_fp8_e4m3_fa2_q_preconvert_f16_gfx1201", + "attention_fp8_e4m3_fa2_q_preconvert_fp8_gfx1201" + ] + ); + add!( + "dflash_gdn_replay_pre_ml", + crate::dflash_gdn_replay::DFLASH_GDN_REPLAY_PRE_ML_SRC, + ["dflash_gdn_replay_pre_ml"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_i4_gfx12_v2_slab_fdiv", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_GFX12_V2_SLAB_FDIV_SRC, + ["fused_rmsnorm_mq_rotate_awq_i4_gfx12_v2_slab_fdiv"] + ); + add!( + "fused_silu_mul_mq_rotate_awq_i4_hin_gfx12_slab_tokfast", + kernels::FUSED_SILU_MUL_MQ_ROTATE_AWQ_I4_HIN_GFX12_SLAB_TOKFAST_SRC, + ["fused_silu_mul_mq_rotate_awq_i4_hin_gfx12_slab_tokfast"] + ); + add!( + "gated_delta_net_q8_fast_ml", + crate::dflash_gdn_replay::GATED_DELTA_NET_Q8_FAST_ML_SRC, + ["gated_delta_net_q8_fast_ml"] + ); + add!( + "gated_norm_mq_rotate_awq_i4_gfx12_v2_slab", + kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX12_V2_SLAB_SRC, + ["gated_norm_mq_rotate_awq_i4_gfx12_v2_slab"] + ); + add!( + "gated_norm_mq_rotate_awq_i4_gfx12_v2_xbf16_slab", + kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX12_V2_XBF16_SLAB_SRC, + ["gated_norm_mq_rotate_awq_i4_gfx12_v2_xbf16_slab"] + ); + add!( + "gdn_chunk_kkt_solve_batched", + kernels::GDN_CHUNK_KKT_SOLVE_BATCHED_SRC, + ["gdn_chunk_kkt_solve_batched"] + ); + add!( + "gdn_chunk_prep", + kernels::GDN_CHUNK_PREP_SRC, + ["gdn_chunk_prep"] + ); + add!( + "gdn_chunk_prep_fixup", + kernels::GDN_CHUNK_PREP_FIXUP_SRC, + ["gdn_chunk_prep_fixup"] + ); + add!( + "gdn_chunk_scan_bf16", + kernels::GDN_CHUNK_SCAN_BF16_SRC, + ["gdn_chunk_scan_bf16"] + ); + add!( + "gdn_chunk_scan_bf16_mseg", + kernels::GDN_CHUNK_SCAN_BF16_MSEG_SRC, + ["gdn_chunk_scan_bf16_mseg"] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "quantize_int4_mmq_ds128" + ] + ); + add!( + "qwen35_fa_prep_fp8q_nogate_batched_gfx1201", + crate::qwen35_fa_batch::FA_PREP_BATCHED_GFX1201_SRC, + [ + "qwen35_fa_prep_batched_gfx1201", + "qwen35_fa_prep_fp8q_batched_gfx1201", + "qwen35_fa_prep_fp8q_nogate_batched_gfx1201" + ] + ); + add!( + "select_regrid", + crate::select_regrid::SELECT_REGRID_SRC, + ["argmax_f32_batched_regrid", "topk_values_regrid_f32"] + ); + add!( + "sigmoid_mul_rotate_x_mq_awq_i4_gfx12_slab", + kernels::SIGMOID_MUL_MQ_ROTATE_X_AWQ_I4_GFX12_SLAB_SRC, + ["sigmoid_mul_rotate_x_mq_awq_i4_gfx12_slab"] + ); + add!( + "topk_values_fixup_f32", + crate::select_regrid::TOPK_VALUES_FIXUP_SRC, + ["topk_values_fixup_f32"] + ); } if arch == "gfx1100" { - add!("attention_flash_q8_0_reduce_gated_mq_rotate_awq_dec_gfx1100", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_AWQ_DEC_GFX1100_SRC, ["attention_flash_q8_0_reduce_gated_mq_rotate_awq_dec_gfx1100"]); - add!("attention_flash_q8_0_tile_batched", assemble_asym(kernels::ATTENTION_FLASH_Q8_0_TILE_BATCHED_SRC), ["attention_flash_q8_0_tile_batched"]); - add!("attention_flash_q8_0_tile_gqa_gfx1100", kernels::ATTENTION_FLASH_Q8_0_TILE_GQA_GFX1100_SRC, ["attention_flash_q8_0_tile_gqa_gfx1100"]); - add!("attention_q8_0_fa2_gqa_gfx1100", q8_fa2_gqa_gfx1100_source(), ["attention_fa2_q_preconvert_gfx1100", "attention_q8_0_fa2_gqa_gfx1100"]); - add!("conv1d_silu_split_qknorm_b256_scalar_prep", kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SCALAR_PREP_SRC, ["conv1d_silu_split_qknorm_b256_scalar_prep"]); - add!("dflash_gdn_pre_gfx1100", crate::dflash_gdn_pre::DFLASH_GDN_PRE_GFX1100_SRC, ["dflash_gdn_pre_capture_gfx1100", "dflash_gdn_pre_replay_gfx1100"]); - add!("dflash_hidden_commit5_gfx1100", crate::dflash_hidden_scatter::DFLASH_HIDDEN_SCATTER_SRC, ["dflash_hidden_commit5_gfx1100", "dflash_hidden_scatter5_gfx1100"]); - add!("dflash_hidden_scatter5_gfx1100", crate::dflash_hidden_scatter::DFLASH_HIDDEN_SCATTER_SRC, ["dflash_hidden_commit5_gfx1100", "dflash_hidden_scatter5_gfx1100"]); - add!("dflash_state_bulk_copy_gfx1100", crate::dflash_state_copy::DFLASH_STATE_BULK_COPY_GFX1100_SRC, ["dflash_state_bulk_copy_gfx1100"]); - add!("dynamic_conv_residual_gfx1100", crate::dflash_draft_fusion::COLLAPSE_SRC, ["dynamic_conv_residual_gfx1100", "gemm_hfq4g256_overwrite_wmma_k2_dflash_gfx1100", "gemm_ksplit_det_overwrite_finalize_dflash_gfx1100", "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks2", "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks4", "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks8", "mq_rotate_x_f16_dflash_gfx1100", "rmsnorm_residual_dual_gfx1100"]); - add!("fused_rmsnorm_mq_rotate_awq_f16", crate::mq_f16_producers::FUSED_RMSNORM_MQ_ROTATE_F16_SRC, ["fused_rmsnorm_mq_rotate_awq_f16", "fused_rmsnorm_mq_rotate_f16"]); - add!("fused_rmsnorm_mq_rotate_awq_i4_b8", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_B8_SRC, ["fused_rmsnorm_mq_rotate_awq_i4_b8"]); - add!("fused_silu_mul_mq_rotate_f16", crate::mq_f16_residual_producers::FUSED_SILU_F16_SRC, ["fused_silu_mul_mq_rotate_awq_f16_batched_gfx1100", "fused_silu_mul_mq_rotate_f16_batched_gfx1100"]); - add!("gated_delta_net_q8_compact3_b2", kernels::GATED_DELTA_NET_Q8_COMPACT3_B2_SRC, ["gated_delta_net_q8_compact3_b2", "gated_delta_net_q8_fast_independent_masked"]); - add!("gated_norm_mq_rotate_awq_i4_gfx1100_v2", kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX1100_V2_SRC, ["gated_norm_mq_rotate_awq_i4_gfx1100_v2"]); - add!("gated_norm_mq_rotate_awq_k6144_gfx1100", kernels::gated_norm_mq_rotate_awq_k6144_gfx1100_src(), ["gated_norm_mq_rotate_awq_k6144_gfx1100"]); - add!("gated_norm_mq_rotate_f16", crate::mq_f16_residual_producers::GATED_NORM_F16_SRC, ["gated_norm_mq_rotate_awq_f16_batched_gfx1100", "gated_norm_mq_rotate_f16_batched_gfx1100"]); - add!("gdn_chunk_kkt_solve_gfx1100", kernels::GDN_CHUNK_KKT_SOLVE_GFX1100_SRC, ["gdn_chunk_kkt_solve_gfx1100"]); - add!("gdn_chunk_scan", kernels::GDN_CHUNK_SCAN_SRC, ["gdn_chunk_scan"]); - add!("gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage", kernels::GEMM_GATE_UP_MQ4G256V2_WMMA_GFX1100_LDSSTAGE_SRC, ["gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage"]); - add!("gemm_mq4g256v2_residual_mmq_iu4", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_SRC, ["gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_full_add", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "quantize_int4_mmq_ds128"]); - add!("gemm_mq4g256v2_residual_mmq_iu4_gridspec_symfold", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SYMFOLD_SRC, ["gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", "gemm_mq4g256v2_residual_mmq_iu4_full_add", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_symfold", "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", "quantize_int4_mmq_ds128"]); - add!("gemm_mq4g256v2_residual_wmma_gfx1100_ldsstage", kernels::GEMM_MQ4G256V2_RESIDUAL_WMMA_GFX1100_LDSSTAGE_SRC, ["gemm_mq4g256v2_residual_wmma_gfx1100_ldsstage"]); - add!("gemm_mqv2_wmma_gfx1100_mw_lds", kernels::GEMM_MQV2_WMMA_GFX11_MW_LDS_SRC, ["gemm_gate_up_mq3g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq3g256v2_wmma_gfx11_mw8_lds", "gemm_gate_up_mq4g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq4g256v2_wmma_gfx11_mw8_lds", "gemm_gate_up_mq5g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq5g256v2_wmma_gfx11_mw8_lds", "gemm_gate_up_mq6g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq6g256v2_wmma_gfx11_mw8_lds", "gemm_mq3g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq3g256v2_residual_wmma_gfx11_mw8_lds", "gemm_mq4g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq4g256v2_residual_wmma_gfx11_mw8_lds", "gemm_mq5g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq5g256v2_residual_wmma_gfx11_mw8_lds", "gemm_mq6g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq6g256v2_residual_wmma_gfx11_mw8_lds"]); - add!("gemv_hfq4g256_residual_rdna3_mq4v2", kernels::GEMV_MQ4G256V2_RESIDUAL_SRC, ["gemv_mq4g256v2_residual"]); - add!("gemv_mq4g256v2_rdna3_mq4v2", kernels::GEMV_MQ4G256V2_SRC, ["gemv_mq4g256v2"]); - add!("kv_cache_write_q8_0_pair", kernels::KV_CACHE_WRITE_Q8_0_PAIR_GFX1100_SRC, ["kv_cache_write_q8_0_pair"]); - add!("kv_cache_write_q8_0_pair_batched_gfx1100", crate::qwen35_fa_batch::KV_PAIR_BATCHED_SRC, ["kv_cache_write_q8_0_pair_batched_gfx1100"]); - add!("qwen35_fa_prep_batched_gfx1100", crate::qwen35_fa_batch::FA_PREP_BATCHED_SRC, ["qwen35_fa_prep_batched_gfx1100"]); - add!("qwen36_27b_fa_prep_gfx1100", kernels::qwen36_27b_fa_prep_gfx1100_src(), ["qwen36_27b_fa_prep_gfx1100"]); - add!("rmsnorm_residual_dual_gfx1100", crate::dflash_draft_fusion::COLLAPSE_SRC, ["dynamic_conv_residual_gfx1100", "gemm_hfq4g256_overwrite_wmma_k2_dflash_gfx1100", "gemm_ksplit_det_overwrite_finalize_dflash_gfx1100", "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks2", "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks4", "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks8", "mq_rotate_x_f16_dflash_gfx1100", "rmsnorm_residual_dual_gfx1100"]); - add!("rope_partial_halfsplit", kernels::ROPE_PARTIAL_HALFSPLIT_SRC, ["rope_partial_halfsplit_f32"]); - add!("sigmoid_mul_mq_rotate_f16", crate::mq_f16_residual_producers::SIGMOID_MUL_F16_SRC, ["sigmoid_mul_mq_rotate_awq_f16_batched_gfx1100", "sigmoid_mul_mq_rotate_f16_batched_gfx1100"]); + add!( + "attention_flash_q8_0_reduce_gated_mq_rotate_awq_dec_gfx1100", + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_AWQ_DEC_GFX1100_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_awq_dec_gfx1100"] + ); + add!( + "attention_flash_q8_0_tile_batched", + assemble_asym(kernels::ATTENTION_FLASH_Q8_0_TILE_BATCHED_SRC), + ["attention_flash_q8_0_tile_batched"] + ); + add!( + "attention_flash_q8_0_tile_gqa_gfx1100", + kernels::ATTENTION_FLASH_Q8_0_TILE_GQA_GFX1100_SRC, + ["attention_flash_q8_0_tile_gqa_gfx1100"] + ); + add!( + "attention_q8_0_fa2_gqa_gfx1100", + q8_fa2_gqa_gfx1100_source(), + [ + "attention_fa2_q_preconvert_gfx1100", + "attention_q8_0_fa2_gqa_gfx1100", + "attention_q8_0_fa2_gqa_partial_gfx1100", + "attention_q8_0_fa2_gqa_merge_gfx1100" + ] + ); + add!( + "conv1d_silu_split_qknorm_b256_scalar_prep", + kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SCALAR_PREP_SRC, + ["conv1d_silu_split_qknorm_b256_scalar_prep"] + ); + add!( + "dflash_gdn_pre_gfx1100", + crate::dflash_gdn_pre::DFLASH_GDN_PRE_GFX1100_SRC, + [ + "dflash_gdn_pre_capture_gfx1100", + "dflash_gdn_pre_replay_gfx1100" + ] + ); + add!( + "dflash_hidden_commit5_gfx1100", + crate::dflash_hidden_scatter::DFLASH_HIDDEN_SCATTER_SRC, + [ + "dflash_hidden_commit5_gfx1100", + "dflash_hidden_scatter5_gfx1100" + ] + ); + add!( + "dflash_hidden_scatter5_gfx1100", + crate::dflash_hidden_scatter::DFLASH_HIDDEN_SCATTER_SRC, + [ + "dflash_hidden_commit5_gfx1100", + "dflash_hidden_scatter5_gfx1100" + ] + ); + add!( + "dflash_state_bulk_copy_gfx1100", + crate::dflash_state_copy::DFLASH_STATE_BULK_COPY_GFX1100_SRC, + ["dflash_state_bulk_copy_gfx1100"] + ); + add!( + "dynamic_conv_residual_gfx1100", + crate::dflash_draft_fusion::COLLAPSE_SRC, + [ + "dynamic_conv_residual_gfx1100", + "gemm_hfq4g256_overwrite_wmma_k2_dflash_gfx1100", + "gemm_ksplit_det_overwrite_finalize_dflash_gfx1100", + "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks2", + "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks4", + "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks8", + "mq_rotate_x_f16_dflash_gfx1100", + "rmsnorm_residual_dual_gfx1100" + ] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_f16", + crate::mq_f16_producers::FUSED_RMSNORM_MQ_ROTATE_F16_SRC, + [ + "fused_rmsnorm_mq_rotate_awq_f16", + "fused_rmsnorm_mq_rotate_f16" + ] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_i4_b8", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_B8_SRC, + ["fused_rmsnorm_mq_rotate_awq_i4_b8"] + ); + add!( + "fused_silu_mul_mq_rotate_f16", + crate::mq_f16_residual_producers::FUSED_SILU_F16_SRC, + [ + "fused_silu_mul_mq_rotate_awq_f16_batched_gfx1100", + "fused_silu_mul_mq_rotate_f16_batched_gfx1100" + ] + ); + add!( + "gated_delta_net_q8_compact3_b2", + kernels::GATED_DELTA_NET_Q8_COMPACT3_B2_SRC, + [ + "gated_delta_net_q8_compact3_b2", + "gated_delta_net_q8_fast_independent_masked" + ] + ); + add!( + "gated_norm_mq_rotate_awq_i4_gfx1100_v2", + kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX1100_V2_SRC, + ["gated_norm_mq_rotate_awq_i4_gfx1100_v2"] + ); + add!( + "gated_norm_mq_rotate_awq_k6144_gfx1100", + kernels::gated_norm_mq_rotate_awq_k6144_gfx1100_src(), + ["gated_norm_mq_rotate_awq_k6144_gfx1100"] + ); + add!( + "gated_norm_mq_rotate_f16", + crate::mq_f16_residual_producers::GATED_NORM_F16_SRC, + [ + "gated_norm_mq_rotate_awq_f16_batched_gfx1100", + "gated_norm_mq_rotate_f16_batched_gfx1100" + ] + ); + add!( + "gdn_chunk_kkt_solve_gfx1100", + kernels::GDN_CHUNK_KKT_SOLVE_GFX1100_SRC, + ["gdn_chunk_kkt_solve_gfx1100"] + ); + add!( + "gdn_chunk_scan", + kernels::GDN_CHUNK_SCAN_SRC, + ["gdn_chunk_scan"] + ); + add!( + "gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage", + kernels::GEMM_GATE_UP_MQ4G256V2_WMMA_GFX1100_LDSSTAGE_SRC, + ["gemm_gate_up_mq4g256v2_wmma_gfx1100_ldsstage"] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "quantize_int4_mmq_ds128" + ] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4_gridspec_symfold", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SYMFOLD_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", + "quantize_int4_mmq_ds128" + ] + ); + add!( + "gemm_mq4g256v2_residual_wmma_gfx1100_ldsstage", + kernels::GEMM_MQ4G256V2_RESIDUAL_WMMA_GFX1100_LDSSTAGE_SRC, + ["gemm_mq4g256v2_residual_wmma_gfx1100_ldsstage"] + ); + add!( + "gemm_mqv2_wmma_gfx1100_mw_lds", + kernels::GEMM_MQV2_WMMA_GFX11_MW_LDS_SRC, + [ + "gemm_gate_up_mq3g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq3g256v2_wmma_gfx11_mw8_lds", + "gemm_gate_up_mq4g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq4g256v2_wmma_gfx11_mw8_lds", + "gemm_gate_up_mq5g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq5g256v2_wmma_gfx11_mw8_lds", + "gemm_gate_up_mq6g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq6g256v2_wmma_gfx11_mw8_lds", + "gemm_mq3g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq3g256v2_residual_wmma_gfx11_mw8_lds", + "gemm_mq4g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq4g256v2_residual_wmma_gfx11_mw8_lds", + "gemm_mq5g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq5g256v2_residual_wmma_gfx11_mw8_lds", + "gemm_mq6g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq6g256v2_residual_wmma_gfx11_mw8_lds" + ] + ); + add!( + "gemv_hfq4g256_residual_rdna3_mq4v2", + kernels::GEMV_MQ4G256V2_RESIDUAL_SRC, + ["gemv_mq4g256v2_residual"] + ); + add!( + "gemv_mq4g256v2_rdna3_mq4v2", + kernels::GEMV_MQ4G256V2_SRC, + ["gemv_mq4g256v2"] + ); + add!( + "kv_cache_write_q8_0_pair", + kernels::KV_CACHE_WRITE_Q8_0_PAIR_GFX1100_SRC, + ["kv_cache_write_q8_0_pair"] + ); + add!( + "kv_cache_write_q8_0_pair_batched_gfx1100", + crate::qwen35_fa_batch::KV_PAIR_BATCHED_SRC, + ["kv_cache_write_q8_0_pair_batched_gfx1100"] + ); + add!( + "qwen35_fa_prep_batched_gfx1100", + crate::qwen35_fa_batch::FA_PREP_BATCHED_SRC, + ["qwen35_fa_prep_batched_gfx1100"] + ); + add!( + "qwen36_27b_fa_prep_gfx1100", + kernels::qwen36_27b_fa_prep_gfx1100_src(), + ["qwen36_27b_fa_prep_gfx1100"] + ); + add!( + "rmsnorm_residual_dual_gfx1100", + crate::dflash_draft_fusion::COLLAPSE_SRC, + [ + "dynamic_conv_residual_gfx1100", + "gemm_hfq4g256_overwrite_wmma_k2_dflash_gfx1100", + "gemm_ksplit_det_overwrite_finalize_dflash_gfx1100", + "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks2", + "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks4", + "gemm_mq4g256v2_overwrite_ksplit_lds_dflash_gfx1100_ks8", + "mq_rotate_x_f16_dflash_gfx1100", + "rmsnorm_residual_dual_gfx1100" + ] + ); + add!( + "rope_partial_halfsplit", + kernels::ROPE_PARTIAL_HALFSPLIT_SRC, + ["rope_partial_halfsplit_f32"] + ); + add!( + "sigmoid_mul_mq_rotate_f16", + crate::mq_f16_residual_producers::SIGMOID_MUL_F16_SRC, + [ + "sigmoid_mul_mq_rotate_awq_f16_batched_gfx1100", + "sigmoid_mul_mq_rotate_f16_batched_gfx1100" + ] + ); } if arch == "gfx1151" { - add!("attention_flash_q8_0_tile_gqa_gfx1151", kernels::ATTENTION_FLASH_Q8_0_TILE_GQA_GFX1151_SRC, ["attention_flash_q8_0_tile_gqa_gfx1151"]); - add!("attention_flash_reduce_dsplit_gfx1151", kernels::ATTENTION_FLASH_REDUCE_DSPLIT_GFX1151_SRC, ["attention_flash_reduce_dsplit_gfx1151"]); - add!("attention_q8_0_fa2_gqa_gfx1151", kernels::ATTENTION_Q8_0_FA2_GQA_GFX1151_SRC, ["attention_fa2_q_preconvert_gfx1151", "attention_q8_0_fa2_gqa_gfx1151"]); - add!("conv1d_silu_split_qknorm_b256", kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SRC, ["conv1d_silu_split_qknorm_b256"]); - add!("deinterleave_q_rmsnorm_f32_batched", kernels::DEINTERLEAVE_Q_RMSNORM_BATCHED_SRC, ["deinterleave_q_rmsnorm_f32_batched"]); - add!("fused_rmsnorm_mq_rotate_awq_i4_b8_gfx1151", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_B8_GFX1151_SRC, ["fused_rmsnorm_mq_rotate_awq_i4_b8_gfx1151"]); - add!("fused_rmsnorm_mq_rotate_awq_i4_fold_b8", kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_FOLD_B8_SRC, ["fused_rmsnorm_mq_rotate_awq_i4_fold_b8"]); + add!( + "attention_flash_q8_0_tile_gqa_gfx1151", + kernels::ATTENTION_FLASH_Q8_0_TILE_GQA_GFX1151_SRC, + ["attention_flash_q8_0_tile_gqa_gfx1151"] + ); + add!( + "attention_flash_reduce_dsplit_gfx1151", + kernels::ATTENTION_FLASH_REDUCE_DSPLIT_GFX1151_SRC, + ["attention_flash_reduce_dsplit_gfx1151"] + ); + add!( + "attention_q8_0_fa2_gqa_gfx1151", + kernels::ATTENTION_Q8_0_FA2_GQA_GFX1151_SRC, + [ + "attention_fa2_q_preconvert_gfx1151", + "attention_q8_0_fa2_gqa_gfx1151" + ] + ); + add!( + "conv1d_silu_split_qknorm_b256", + kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SRC, + ["conv1d_silu_split_qknorm_b256"] + ); + add!( + "deinterleave_q_rmsnorm_f32_batched", + kernels::DEINTERLEAVE_Q_RMSNORM_BATCHED_SRC, + ["deinterleave_q_rmsnorm_f32_batched"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_i4_b8_gfx1151", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_B8_GFX1151_SRC, + ["fused_rmsnorm_mq_rotate_awq_i4_b8_gfx1151"] + ); + add!( + "fused_rmsnorm_mq_rotate_awq_i4_fold_b8", + kernels::FUSED_RMSNORM_MQ_ROTATE_AWQ_I4_FOLD_B8_SRC, + ["fused_rmsnorm_mq_rotate_awq_i4_fold_b8"] + ); // Qwen4 chunked-GDN inline-Q8 recurrence (opt-in `HIPFIRE_QWEN4_GDN_Q8_INLINE`, exact gfx1151; tensor_ops.rs gated_delta_step_gate_wmma_arms). - add!("gated_delta_chunk_q8_wmma", crate::tensor_ops::GATED_DELTA_CHUNK_Q8_WMMA_SRC, ["gated_delta_chunk_gate_q8_wmma"]); - add!("gated_norm_mq_rotate_awq_i4_gfx1151_v2", kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX1151_V2_SRC, ["gated_norm_mq_rotate_awq_i4_gfx1151_v2"]); - add!("gdn_chunk_scan_gfx1151", kernels::GDN_CHUNK_SCAN_GFX1151_SRC, ["gdn_chunk_scan_gfx1151"]); - add!("gemm_mq4g256v2_residual_iu4_v2b_gfx11", kernels::GEMM_MQ4G256V2_RESIDUAL_IU4_V2B_GFX11_SRC, ["gemm_mq4g256v2_gate_up_silu_iu4_v2b_gfx11", "gemm_mq4g256v2_residual_iu4_v2b_add_gfx11", "gemm_mq4g256v2_residual_iu4_v2b_add_touch_gfx11", "gemm_mq4g256v2_residual_iu4_v2b_add_touch_swz_gfx11", "gemm_mq4g256v2_residual_iu4_v2b_set_gfx11", "gemm_mq4g256v2_residual_iu4_v2b_set_zba_gfx11"]); - add!("gemm_mq4g256v2_residual_mmq_iu4", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_SRC, ["gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_full_add", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "quantize_int4_mmq_ds128"]); - add!("gemm_mq4g256v2_residual_mmq_iu4_gfx11_x5_symfold", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GFX11_X5_SYMFOLD_SRC, ["gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_full_add", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_x5_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_x5_col_gfx1151_symfold", "quantize_int4_mmq_ds128"]); - add!("gemm_mq4g256v2_residual_mmq_iu4_gridspec_symfold", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SYMFOLD_SRC, ["gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", "gemm_mq4g256v2_residual_mmq_iu4_full_add", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_add_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_symfold", "gemm_mq4g256v2_residual_mmq_iu4_full_set_symfold", "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", "quantize_int4_mmq_ds128"]); - add!("gemm_mqv2_wmma_gfx1151_mw_lds", kernels::GEMM_MQV2_WMMA_GFX11_MW_LDS_SRC, ["gemm_gate_up_mq3g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq3g256v2_wmma_gfx11_mw8_lds", "gemm_gate_up_mq4g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq4g256v2_wmma_gfx11_mw8_lds", "gemm_gate_up_mq5g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq5g256v2_wmma_gfx11_mw8_lds", "gemm_gate_up_mq6g256v2_wmma_gfx11_mw4_lds", "gemm_gate_up_mq6g256v2_wmma_gfx11_mw8_lds", "gemm_mq3g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq3g256v2_residual_wmma_gfx11_mw8_lds", "gemm_mq4g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq4g256v2_residual_wmma_gfx11_mw8_lds", "gemm_mq5g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq5g256v2_residual_wmma_gfx11_mw8_lds", "gemm_mq6g256v2_residual_wmma_gfx11_mw4_lds", "gemm_mq6g256v2_residual_wmma_gfx11_mw8_lds"]); - add!("gemv_hfq4g256_multirow_default_mq4v2", kernels::GEMV_MQ4G256V2_MULTIROW_SRC, ["gemv_mq4g256v2_multirow_r2", "gemv_mq4g256v2_multirow_r4", "gemv_mq4g256v2_multirow_r8"]); - add!("gemv_mq4g256v2_mq4v2", kernels::GEMV_MQ4G256V2_SRC, ["gemv_mq4g256v2"]); - add!("gemv_mq4g256v2_residual_row_serial_gfx1151", kernels::GEMV_MQ4G256V2_RESIDUAL_ROW_SERIAL_GFX1151_SRC, ["gemv_mq4g256v2_residual_row_serial_gfx1151"]); - add!("qwen35_fa_prep_batched_gfx1151", crate::qwen35_fa_batch::FA_PREP_BATCHED_GFX1151_SRC, ["qwen35_fa_prep_batched_gfx1151"]); - add!("repeat_interleave_qk_batched", kernels::REPEAT_INTERLEAVE_QK_BATCHED_SRC, ["repeat_interleave_qk_f32_batched"]); - add!("rope_partial_halfsplit_f32_headgrid", kernels::ROPE_PARTIAL_HALFSPLIT_HEADGRID_SRC, ["rope_partial_halfsplit_f32_headgrid"]); + add!( + "gated_delta_chunk_q8_wmma", + crate::tensor_ops::GATED_DELTA_CHUNK_Q8_WMMA_SRC, + ["gated_delta_chunk_gate_q8_wmma"] + ); + add!( + "gated_norm_mq_rotate_awq_i4_gfx1151_v2", + kernels::GATED_NORM_MQ_ROTATE_AWQ_I4_GFX1151_V2_SRC, + ["gated_norm_mq_rotate_awq_i4_gfx1151_v2"] + ); + add!( + "gdn_chunk_scan_gfx1151", + kernels::GDN_CHUNK_SCAN_GFX1151_SRC, + ["gdn_chunk_scan_gfx1151"] + ); + add!( + "gemm_mq4g256v2_residual_iu4_v2b_gfx11", + kernels::GEMM_MQ4G256V2_RESIDUAL_IU4_V2B_GFX11_SRC, + [ + "gemm_mq4g256v2_gate_up_silu_iu4_v2b_gfx11", + "gemm_mq4g256v2_residual_iu4_v2b_add_gfx11", + "gemm_mq4g256v2_residual_iu4_v2b_add_touch_gfx11", + "gemm_mq4g256v2_residual_iu4_v2b_add_touch_swz_gfx11", + "gemm_mq4g256v2_residual_iu4_v2b_set_gfx11", + "gemm_mq4g256v2_residual_iu4_v2b_set_zba_gfx11" + ] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "quantize_int4_mmq_ds128" + ] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4_gfx11_x5_symfold", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GFX11_X5_SYMFOLD_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_x5_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_x5_col_gfx1151_symfold", + "quantize_int4_mmq_ds128" + ] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4_gridspec_symfold", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SYMFOLD_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_symfold", + "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", + "quantize_int4_mmq_ds128" + ] + ); + add!( + "gemm_mqv2_wmma_gfx1151_mw_lds", + kernels::GEMM_MQV2_WMMA_GFX11_MW_LDS_SRC, + [ + "gemm_gate_up_mq3g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq3g256v2_wmma_gfx11_mw8_lds", + "gemm_gate_up_mq4g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq4g256v2_wmma_gfx11_mw8_lds", + "gemm_gate_up_mq5g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq5g256v2_wmma_gfx11_mw8_lds", + "gemm_gate_up_mq6g256v2_wmma_gfx11_mw4_lds", + "gemm_gate_up_mq6g256v2_wmma_gfx11_mw8_lds", + "gemm_mq3g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq3g256v2_residual_wmma_gfx11_mw8_lds", + "gemm_mq4g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq4g256v2_residual_wmma_gfx11_mw8_lds", + "gemm_mq5g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq5g256v2_residual_wmma_gfx11_mw8_lds", + "gemm_mq6g256v2_residual_wmma_gfx11_mw4_lds", + "gemm_mq6g256v2_residual_wmma_gfx11_mw8_lds" + ] + ); + add!( + "gemv_hfq4g256_multirow_default_mq4v2", + kernels::GEMV_MQ4G256V2_MULTIROW_SRC, + [ + "gemv_mq4g256v2_multirow_r2", + "gemv_mq4g256v2_multirow_r4", + "gemv_mq4g256v2_multirow_r8" + ] + ); + add!( + "gemv_mq4g256v2_mq4v2", + kernels::GEMV_MQ4G256V2_SRC, + ["gemv_mq4g256v2"] + ); + add!( + "gemv_mq4g256v2_residual_row_serial_gfx1151", + kernels::GEMV_MQ4G256V2_RESIDUAL_ROW_SERIAL_GFX1151_SRC, + ["gemv_mq4g256v2_residual_row_serial_gfx1151"] + ); + add!( + "qwen35_fa_prep_batched_gfx1151", + crate::qwen35_fa_batch::FA_PREP_BATCHED_GFX1151_SRC, + ["qwen35_fa_prep_batched_gfx1151"] + ); + add!( + "repeat_interleave_qk_batched", + kernels::REPEAT_INTERLEAVE_QK_BATCHED_SRC, + ["repeat_interleave_qk_f32_batched"] + ); + add!( + "rope_partial_halfsplit_f32_headgrid", + kernels::ROPE_PARTIAL_HALFSPLIT_HEADGRID_SRC, + ["rope_partial_halfsplit_f32_headgrid"] + ); } // Qwen3.8-Flash-Next (Qwen4) load, AR prefill/decode and native MTP: // every module a kernel-load trace (tests/fixtures/kernel-trace-qwen4- @@ -471,340 +1614,1016 @@ pub fn entries(arch: &str, extra_flags: &str) -> Result, Regist // the exact expressions the callsites pass to `ensure_kernel`; symbols are // every kernel the arch's object defines. if matches!(arch, "gfx1201" | "gfx1100" | "gfx1151") { - add!("copy_f32_buffer", kernels::COPY_F32_BUFFER_SRC, ["copy_f32_buffer", "copy_f32_strided_slot_buffer"]); - add!("gemm_bf16_xf32_multirow", kernels::GEMM_BF16_XF32_MULTIROW_SRC, [ - "convert_bf16_to_f16", "gemm_bf16_xf32_multirow", "gemm_bf16_xf32_multirow_pto2", "gemm_bf16_xf32_multirow_rows3", - ]); - add!("gemm_bf16_xf32_multirow_r16w4_gfx1151", kernels::GEMM_BF16_XF32_MULTIROW_R16_GFX1151_SRC, [ - "gemm_bf16_xf32_multirow_r16_gfx1151", "gemm_bf16_xf32_multirow_r16w4_gfx1151", "gemm_bf16_xf32_multirow_r16w4t_gfx1151", - ]); - add!("gemm_bf16_xf32_multirow_r16w4t_gfx1151", kernels::GEMM_BF16_XF32_MULTIROW_R16_GFX1151_SRC, [ - "gemm_bf16_xf32_multirow_r16_gfx1151", "gemm_bf16_xf32_multirow_r16w4_gfx1151", "gemm_bf16_xf32_multirow_r16w4t_gfx1151", - ]); - add!("gemm_bf16_xf32_multirow_rows3", kernels::GEMM_BF16_XF32_MULTIROW_SRC, [ - "convert_bf16_to_f16", "gemm_bf16_xf32_multirow", "gemm_bf16_xf32_multirow_pto2", "gemm_bf16_xf32_multirow_rows3", - ]); - add!("gemm_mq4g128v2_moe_grouped_top10_o8_r16_gfx1151", kernels::GEMM_MQ4G128V2_MOE_GROUPED_TOP10_O8_R16_GFX1151_SRC, [ + add!( + "copy_f32_buffer", + kernels::COPY_F32_BUFFER_SRC, + ["copy_f32_buffer", "copy_f32_strided_slot_buffer"] + ); + add!( + "gemm_bf16_xf32_multirow", + kernels::GEMM_BF16_XF32_MULTIROW_SRC, + [ + "convert_bf16_to_f16", + "gemm_bf16_xf32_multirow", + "gemm_bf16_xf32_multirow_pto2", + "gemm_bf16_xf32_multirow_rows3", + ] + ); + add!( + "gemm_bf16_xf32_multirow_r16w4_gfx1151", + kernels::GEMM_BF16_XF32_MULTIROW_R16_GFX1151_SRC, + [ + "gemm_bf16_xf32_multirow_r16_gfx1151", + "gemm_bf16_xf32_multirow_r16w4_gfx1151", + "gemm_bf16_xf32_multirow_r16w4t_gfx1151", + ] + ); + add!( + "gemm_bf16_xf32_multirow_r16w4t_gfx1151", + kernels::GEMM_BF16_XF32_MULTIROW_R16_GFX1151_SRC, + [ + "gemm_bf16_xf32_multirow_r16_gfx1151", + "gemm_bf16_xf32_multirow_r16w4_gfx1151", + "gemm_bf16_xf32_multirow_r16w4t_gfx1151", + ] + ); + add!( + "gemm_bf16_xf32_multirow_rows3", + kernels::GEMM_BF16_XF32_MULTIROW_SRC, + [ + "convert_bf16_to_f16", + "gemm_bf16_xf32_multirow", + "gemm_bf16_xf32_multirow_pto2", + "gemm_bf16_xf32_multirow_rows3", + ] + ); + add!( "gemm_mq4g128v2_moe_grouped_top10_o8_r16_gfx1151", - ]); - add!("gemm_mq4g256v2_moe_grouped_top10_o4_r8_x4", kernels::GEMM_MQ4G256V2_MOE_GROUPED_TOP10_O4_R8_X4_GFX1151_SRC, [ - "gemm_mq4g256v2_moe_grouped_top10_o4_r2_x4", "gemm_mq4g256v2_moe_grouped_top10_o4_r8_x4", - ]); - add!("gemm_mq6g256v2_f32_rows", kernels::GEMM_MQ6G256V2_F32_ROWS_SRC, [ - "gemm_mq6g256v2_f32_rows_r2", "gemm_mq6g256v2_f32_rows_r3", "gemm_mq6g256v2_f32_rows_r4", "gemm_mq6g256v2_f32_rows_r5", - "gemm_mq6g256v2_f32_rows_r6", "gemm_mq6g256v2_f32_rows_r7", "gemm_mq6g256v2_f32_rows_r8", "gemm_mq6g256v2_f32_rows_x4_r2", - "gemm_mq6g256v2_f32_rows_x4_r3", "gemm_mq6g256v2_f32_rows_x4_r4", "gemm_mq6g256v2_f32_rows_x4_r5", "gemm_mq6g256v2_f32_rows_x4_r6", - "gemm_mq6g256v2_f32_rows_x4_r7", "gemm_mq6g256v2_f32_rows_x4_r8", - ]); - add!("gemv_mq4g128v2_moe_down_top10_indexed_batched_expanded", kernels::GEMV_MQ4G128V2_MOE_DOWN_TOP10_INDEXED_BATCHED_EXPANDED_SRC, [ + kernels::GEMM_MQ4G128V2_MOE_GROUPED_TOP10_O8_R16_GFX1151_SRC, + ["gemm_mq4g128v2_moe_grouped_top10_o8_r16_gfx1151",] + ); + add!( + "gemm_mq4g256v2_moe_grouped_top10_o4_r8_x4", + kernels::GEMM_MQ4G256V2_MOE_GROUPED_TOP10_O4_R8_X4_GFX1151_SRC, + [ + "gemm_mq4g256v2_moe_grouped_top10_o4_r2_x4", + "gemm_mq4g256v2_moe_grouped_top10_o4_r8_x4", + ] + ); + add!( + "gemm_mq6g256v2_f32_rows", + kernels::GEMM_MQ6G256V2_F32_ROWS_SRC, + [ + "gemm_mq6g256v2_f32_rows_r2", + "gemm_mq6g256v2_f32_rows_r3", + "gemm_mq6g256v2_f32_rows_r4", + "gemm_mq6g256v2_f32_rows_r5", + "gemm_mq6g256v2_f32_rows_r6", + "gemm_mq6g256v2_f32_rows_r7", + "gemm_mq6g256v2_f32_rows_r8", + "gemm_mq6g256v2_f32_rows_x4_r2", + "gemm_mq6g256v2_f32_rows_x4_r3", + "gemm_mq6g256v2_f32_rows_x4_r4", + "gemm_mq6g256v2_f32_rows_x4_r5", + "gemm_mq6g256v2_f32_rows_x4_r6", + "gemm_mq6g256v2_f32_rows_x4_r7", + "gemm_mq6g256v2_f32_rows_x4_r8", + ] + ); + add!( "gemv_mq4g128v2_moe_down_top10_indexed_batched_expanded", - ]); - add!("gemv_mq4g256v2_moe_gate_up_top10_indexed_batched", kernels::GEMV_MQ4G256V2_MOE_GATE_UP_TOP10_INDEXED_BATCHED_SRC, [ - "gemv_mq4g256v2_moe_gate_up_k2560_rows8_indexed_batched", "gemv_mq4g256v2_moe_gate_up_top10_indexed_batched", - ]); - add!("grouped_ops", crate::grouped_ops::GROUPED_KERNEL_SRC, [ - "grouped_depthwise_conv_silu_add_bf16", "grouped_depthwise_conv_silu_add_f32", "grouped_gate_bf16", "grouped_gate_f32", - "grouped_gather_convert_bf16", "grouped_gather_convert_bf16_f16", "grouped_linear_f32", "grouped_norm_bf16", "grouped_norm_f32", - "ple_conv_add_bf16s", "ple_conv_add_bf16s_k4d3", "ple_gate_rows_bf16s", "ple_norm_inv", - ]); - add!("hc_streams_init_from_embed_batched", kernels::HC_STREAMS_INIT_FROM_EMBED_BATCHED_SRC, ["hc_streams_init_from_embed_batched"]); - add!("moe_down_combine_grouped_top10", kernels::MOE_DOWN_COMBINE_GROUPED_TOP10_SRC, [ - "moe_combine_order_top10", "moe_down_combine_grouped_top10", "moe_down_combine_grouped_top10_bf16in", - "moe_down_combine_grouped_top10_bf16in_zinit", - ]); - add!("moe_down_combine_top10_batched", kernels::MOE_DOWN_COMBINE_TOP10_BATCHED_SRC, ["moe_down_combine_top10_batched"]); - add!("moe_gate_up_unscatter_silu_top10", kernels::MOE_GATE_UP_UNSCATTER_SILU_TOP10_SRC, [ - "moe_gate_up_unscatter_silu_rotate128_top10", "moe_gate_up_unscatter_silu_top10", "moe_gate_up_unscatter_silu_top10_bf16in", - "moe_unscatter_rotate128_f16", - ]); - add!("moe_router_softmax_top10_f32", kernels::MOE_ROUTER_SOFTMAX_TOP10_F32_SRC, ["moe_router_softmax_top10_f32"]); - add!("moe_scatter_fused_top10", kernels::MOE_SCATTER_FUSED_TOP10_SRC, ["moe_scatter_fused_top10"]); - add!("mq_rotate_x_128_v2", kernels::MQ_ROTATE_X_128_V2_SRC, ["mq_rotate_x_128_v2", "mq_rotate_x_128_v2_f16", "mq_rotate_x_128_v2_silu_bf16"]); - add!("qwen4_gemv_bf16_xf32", kernels::QWEN4_GEMV_BF16_XF32_SRC, [ - "gemv_bf16_xf32", "gemv_bf16_xf32_bf16_scaled_add", "gemv_bf16_xf32_k4", "gemv_bf16_xf32_k4_rows_r2", "gemv_bf16_xf32_k4_rows_r3", - "gemv_bf16_xf32_k4_rows_r4", "gemv_bf16_xf32_k4_rows_r5", "gemv_bf16_xf32_k4_rows_r6", "gemv_bf16_xf32_k4_rows_r7", - "gemv_bf16_xf32_k4_rows_r8", "gemv_bf16_xf32_k4_rows_tiled_r2", "gemv_bf16_xf32_k4_rows_tiled_r3", - "gemv_bf16_xf32_k4_rows_tiled_r4", "gemv_bf16_xf32_k4_rows_tiled_r5", "gemv_bf16_xf32_k4_rows_tiled_r6", - "gemv_bf16_xf32_k4_rows_tiled_r7", "gemv_bf16_xf32_k4_rows_tiled_r8", "gemv_bf16_xf32_x4", "gemv_bf16_xf32_x4_rows_r2", - "gemv_bf16_xf32_x4_rows_r3", - "gemv_bf16_xf32_x4_rows_r4", "gemv_bf16_xf32_x4_rows_r5", "gemv_bf16_xf32_x4_rows_r6", "gemv_bf16_xf32_x4_rows_r7", - "gemv_bf16_xf32_x4_rows_r8", "hyper_write_norm_f32", - ]); - add!("qwen4_gemv_mq4g256", kernels::QWEN4_GEMV_MQ4G256_SRC, ["hyper_read_projected_rotate_f32", "mq_rotate_x_bf16_f16", "mq_rotate_x_f16"]); - add!("qwen4_gemv_mq6g256v2", kernels::QWEN4_GEMV_MQ6G256V2_SRC, ["gemv_mq6g256v2", "gemv_mq6g256v2_x4"]); - add!("qwen4_gemv_q8_0", kernels::QWEN4_GEMV_Q8_0_SRC, [ - "gemv_q8_0_k2560_staged", "gemv_q8_0_k2560_staged_pair", "gemv_q8_0_k2560_staged_rows", "gemv_q8_0_k2560_staged_rows_tiled", - "gemv_q8_0_k320_staged", - "gemv_q8_0_k320_staged_rows", "gemv_q8_0_k8", "gemv_q8_0_k8_rows_r2", "gemv_q8_0_k8_rows_r3", "gemv_q8_0_k8_rows_r4", - "gemv_q8_0_k8_rows_r5", "gemv_q8_0_k8_rows_r6", "gemv_q8_0_k8_rows_r7", "gemv_q8_0_k8_rows_r8", "quantize_bf16_q8_0", - "topk8_partial_f32", "topk8_rescore_q8_0_k2560", - ]); - add!("qwen4_hc_streams_init_from_embed_batched", kernels::QWEN4_HC_STREAMS_INIT_FROM_EMBED_BATCHED_SRC, ["hc_streams_init_from_embed_batched_bf16"]); - add!("requant_g256", kernels::REQUANT_G256_SRC, ["requant_bf16_to_f32", "requant_mqg256v2_to_f32", "requant_pack_mqg256v2", "requant_q8_0_to_f32"]); - add!("topk8_rescore_mq6g256v2", kernels::TOPK8_RESCORE_MQ6G256V2_SRC, ["gemv_mq6g256v2", "gemv_mq6g256v2_x4", "topk8_rescore_mq6g256v2_k2560"]); + kernels::GEMV_MQ4G128V2_MOE_DOWN_TOP10_INDEXED_BATCHED_EXPANDED_SRC, + ["gemv_mq4g128v2_moe_down_top10_indexed_batched_expanded",] + ); + add!( + "gemv_mq4g256v2_moe_gate_up_top10_indexed_batched", + kernels::GEMV_MQ4G256V2_MOE_GATE_UP_TOP10_INDEXED_BATCHED_SRC, + [ + "gemv_mq4g256v2_moe_gate_up_k2560_rows8_indexed_batched", + "gemv_mq4g256v2_moe_gate_up_top10_indexed_batched", + ] + ); + add!( + "grouped_ops", + crate::grouped_ops::GROUPED_KERNEL_SRC, + [ + "grouped_depthwise_conv_silu_add_bf16", + "grouped_depthwise_conv_silu_add_f32", + "grouped_gate_bf16", + "grouped_gate_f32", + "grouped_gather_convert_bf16", + "grouped_gather_convert_bf16_f16", + "grouped_linear_f32", + "grouped_norm_bf16", + "grouped_norm_f32", + "ple_conv_add_bf16s", + "ple_conv_add_bf16s_k4d3", + "ple_gate_rows_bf16s", + "ple_norm_inv", + ] + ); + add!( + "hc_streams_init_from_embed_batched", + kernels::HC_STREAMS_INIT_FROM_EMBED_BATCHED_SRC, + ["hc_streams_init_from_embed_batched"] + ); + add!( + "moe_down_combine_grouped_top10", + kernels::MOE_DOWN_COMBINE_GROUPED_TOP10_SRC, + [ + "moe_combine_order_top10", + "moe_down_combine_grouped_top10", + "moe_down_combine_grouped_top10_bf16in", + "moe_down_combine_grouped_top10_bf16in_zinit", + ] + ); + add!( + "moe_down_combine_top10_batched", + kernels::MOE_DOWN_COMBINE_TOP10_BATCHED_SRC, + ["moe_down_combine_top10_batched"] + ); + add!( + "moe_gate_up_unscatter_silu_top10", + kernels::MOE_GATE_UP_UNSCATTER_SILU_TOP10_SRC, + [ + "moe_gate_up_unscatter_silu_rotate128_top10", + "moe_gate_up_unscatter_silu_top10", + "moe_gate_up_unscatter_silu_top10_bf16in", + "moe_unscatter_rotate128_f16", + ] + ); + add!( + "moe_router_softmax_top10_f32", + kernels::MOE_ROUTER_SOFTMAX_TOP10_F32_SRC, + ["moe_router_softmax_top10_f32"] + ); + add!( + "moe_scatter_fused_top10", + kernels::MOE_SCATTER_FUSED_TOP10_SRC, + ["moe_scatter_fused_top10"] + ); + add!( + "mq_rotate_x_128_v2", + kernels::MQ_ROTATE_X_128_V2_SRC, + [ + "mq_rotate_x_128_v2", + "mq_rotate_x_128_v2_f16", + "mq_rotate_x_128_v2_silu_bf16" + ] + ); + add!( + "qwen4_gemv_bf16_xf32", + kernels::QWEN4_GEMV_BF16_XF32_SRC, + [ + "gemv_bf16_xf32", + "gemv_bf16_xf32_bf16_scaled_add", + "gemv_bf16_xf32_k4", + "gemv_bf16_xf32_k4_rows_r2", + "gemv_bf16_xf32_k4_rows_r3", + "gemv_bf16_xf32_k4_rows_r4", + "gemv_bf16_xf32_k4_rows_r5", + "gemv_bf16_xf32_k4_rows_r6", + "gemv_bf16_xf32_k4_rows_r7", + "gemv_bf16_xf32_k4_rows_r8", + "gemv_bf16_xf32_k4_rows_tiled_r2", + "gemv_bf16_xf32_k4_rows_tiled_r3", + "gemv_bf16_xf32_k4_rows_tiled_r4", + "gemv_bf16_xf32_k4_rows_tiled_r5", + "gemv_bf16_xf32_k4_rows_tiled_r6", + "gemv_bf16_xf32_k4_rows_tiled_r7", + "gemv_bf16_xf32_k4_rows_tiled_r8", + "gemv_bf16_xf32_x4", + "gemv_bf16_xf32_x4_rows_r2", + "gemv_bf16_xf32_x4_rows_r3", + "gemv_bf16_xf32_x4_rows_r4", + "gemv_bf16_xf32_x4_rows_r5", + "gemv_bf16_xf32_x4_rows_r6", + "gemv_bf16_xf32_x4_rows_r7", + "gemv_bf16_xf32_x4_rows_r8", + "hyper_write_norm_f32", + ] + ); + add!( + "qwen4_gemv_mq4g256", + kernels::QWEN4_GEMV_MQ4G256_SRC, + [ + "hyper_read_projected_rotate_f32", + "mq_rotate_x_bf16_f16", + "mq_rotate_x_f16" + ] + ); + add!( + "qwen4_gemv_mq6g256v2", + kernels::QWEN4_GEMV_MQ6G256V2_SRC, + ["gemv_mq6g256v2", "gemv_mq6g256v2_x4"] + ); + add!( + "qwen4_gemv_q8_0", + kernels::QWEN4_GEMV_Q8_0_SRC, + [ + "gemv_q8_0_k2560_staged", + "gemv_q8_0_k2560_staged_pair", + "gemv_q8_0_k2560_staged_rows", + "gemv_q8_0_k2560_staged_rows_tiled", + "gemv_q8_0_k320_staged", + "gemv_q8_0_k320_staged_rows", + "gemv_q8_0_k8", + "gemv_q8_0_k8_rows_r2", + "gemv_q8_0_k8_rows_r3", + "gemv_q8_0_k8_rows_r4", + "gemv_q8_0_k8_rows_r5", + "gemv_q8_0_k8_rows_r6", + "gemv_q8_0_k8_rows_r7", + "gemv_q8_0_k8_rows_r8", + "quantize_bf16_q8_0", + "topk8_partial_f32", + "topk8_rescore_q8_0_k2560", + ] + ); + add!( + "qwen4_hc_streams_init_from_embed_batched", + kernels::QWEN4_HC_STREAMS_INIT_FROM_EMBED_BATCHED_SRC, + ["hc_streams_init_from_embed_batched_bf16"] + ); + add!( + "requant_g256", + kernels::REQUANT_G256_SRC, + [ + "requant_bf16_to_f32", + "requant_mqg256v2_to_f32", + "requant_pack_mqg256v2", + "requant_q8_0_to_f32" + ] + ); + add!( + "topk8_rescore_mq6g256v2", + kernels::TOPK8_RESCORE_MQ6G256V2_SRC, + [ + "gemv_mq6g256v2", + "gemv_mq6g256v2_x4", + "topk8_rescore_mq6g256v2_k2560" + ] + ); add!("zero_f32", kernels::ZERO_F32_SRC, ["zero_f32"]); } if matches!(arch, "gfx1201" | "gfx1100") { - add!("gemv_bf16_xf32", kernels::gemv_bf16_xf32_src(arch == "gfx1151"), ["gemv_bf16_xf32"]); + add!( + "gemv_bf16_xf32", + kernels::gemv_bf16_xf32_src(arch == "gfx1151"), + ["gemv_bf16_xf32"] + ); } if matches!(arch, "gfx1201" | "gfx1151") { - add!("gemv_mq2g256v2", kernels::GEMV_MQ2G256V2_SRC, ["gemv_mq2g256v2"]); - add!("gemv_mq6g256v2_mq6v2", kernels::GEMV_MQ6G256V2_SRC, ["gemv_mq6g256v2"]); + add!( + "gemv_mq2g256v2", + kernels::GEMV_MQ2G256V2_SRC, + ["gemv_mq2g256v2"] + ); + add!( + "gemv_mq6g256v2_mq6v2", + kernels::GEMV_MQ6G256V2_SRC, + ["gemv_mq6g256v2"] + ); } if matches!(arch, "gfx1100" | "gfx1151") { - add!("copy_rows_strided_f32", kernels::COPY_ROWS_STRIDED_F32_SRC, ["copy_rows_strided_f32"]); - add!("gated_delta_chunk_wmma", crate::tensor_ops::GATED_DELTA_CHUNK_WMMA_SRC, ["gated_delta_chunk_gate_wmma"]); - add!("gemm_mq4g128v2_moe_grouped_wmma_gfx1151", kernels::GEMM_MQ4G128V2_MOE_GROUPED_WMMA_GFX1151_SRC, [ - "gemm_mq4g128v2_moe_grouped_wmma_gfx1151", "gemm_mq4g128v2_moe_grouped_wmma_gfx1151_bf16out", - ]); - add!("gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", kernels::QWEN4_GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, [ - "gemm_mq4g256v2_moe_grouped_wmma_k2", "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", + add!( + "copy_rows_strided_f32", + kernels::COPY_ROWS_STRIDED_F32_SRC, + ["copy_rows_strided_f32"] + ); + add!( + "gated_delta_chunk_wmma", + crate::tensor_ops::GATED_DELTA_CHUNK_WMMA_SRC, + ["gated_delta_chunk_gate_wmma"] + ); + add!( + "gemm_mq4g128v2_moe_grouped_wmma_gfx1151", + kernels::GEMM_MQ4G128V2_MOE_GROUPED_WMMA_GFX1151_SRC, + [ + "gemm_mq4g128v2_moe_grouped_wmma_gfx1151", + "gemm_mq4g128v2_moe_grouped_wmma_gfx1151_bf16out", + ] + ); + add!( "gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", - ]); - add!("gemm_mq6g256v2_residual_wmma", kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_SRC, ["gemm_mq6g256v2_residual_wmma"]); - add!("hyper_read_up_wmma", crate::tensor_ops::HYPER_READ_UP_WMMA_SRC, ["hyper_read_up_wmma_bf16", "hyper_read_up_wmma_bf16_swap"]); - add!("indexed_attention_dense_wmma", crate::tensor_ops::INDEXED_ATTENTION_DENSE_WMMA_SRC, [ - "indexed_attention_dense_wmma_f16", "indexed_attention_kv_f16", - ]); - add!("qwen4_bf16_round_trip", kernels::QWEN4_BF16_ROUND_TRIP_SRC, ["bf16_round_trip_f32_strided"]); - add!("qwen4_gemm_mqv2_wmma_gfx11_bt", kernels::QWEN4_GEMM_MQV2_WMMA_GFX11_BT_SRC, [ - "gemm_gate_up_mq2g256v2_wmma_gfx11_bt12", "gemm_gate_up_mq3g256v2_wmma_gfx11_bt12", "gemm_gate_up_mq3g256v2_wmma_gfx11_bt6", - "gemm_gate_up_mq5g256v2_wmma_gfx11_bt12", "gemm_gate_up_mq5g256v2_wmma_gfx11_bt6", "gemm_gate_up_mq6g256v2_wmma_gfx11_bt12", - "gemm_gate_up_mq6g256v2_wmma_gfx11_bt6", "gemm_mq2g256v2_residual_wmma_gfx11_bt4", "gemm_mq3g256v2_residual_wmma_gfx11_bt4", - "gemm_mq3g256v2_residual_wmma_gfx11_bt6", "gemm_mq3g256v2_residual_wmma_gfx11_bt8", "gemm_mq5g256v2_residual_wmma_gfx11_bt4", - "gemm_mq5g256v2_residual_wmma_gfx11_bt6", "gemm_mq5g256v2_residual_wmma_gfx11_bt8", "gemm_mq6g256v2_residual_wmma_gfx11_bt4", - "gemm_mq6g256v2_residual_wmma_gfx11_bt6", "gemm_mq6g256v2_residual_wmma_gfx11_bt8", "gemm_mq6g256v2_residual_wmma_gfx11_bt8_x4", - "gemm_mq6g256v2_wmma_gfx11_bt8_x4", "gemm_mq6g256v2_wmma_gfx11_bt8_x4_bf16out", "gemm_mq6g256v2_wmma_gfx11_bt8_x4_hcw", - "gemm_mq6g256v2_wmma_gfx11_bt8_x4_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p1", - "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p1_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p2", - "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p2_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p1", - "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p1_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p2", - "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p2_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_dq", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_dq_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_dq", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_dq_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_regions", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_dq", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_dq_bf16out", "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_hcw", - "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_regions", "gemm_qkv_mq2g256v2_wmma_gfx11_bt4", "gemm_qkv_mq3g256v2_wmma_gfx11_bt12", - "gemm_qkv_mq3g256v2_wmma_gfx11_bt4", "gemm_qkv_mq5g256v2_wmma_gfx11_bt12", "gemm_qkv_mq5g256v2_wmma_gfx11_bt4", - "gemm_qkv_mq6g256v2_wmma_gfx11_bt12", "gemm_qkv_mq6g256v2_wmma_gfx11_bt4", "gemm_qkvza_mq2g256v2_wmma_gfx11_bt4", - "gemm_qkvza_mq3g256v2_wmma_gfx11_bt12", "gemm_qkvza_mq3g256v2_wmma_gfx11_bt4", "gemm_qkvza_mq5g256v2_wmma_gfx11_bt12", - "gemm_qkvza_mq5g256v2_wmma_gfx11_bt4", "gemm_qkvza_mq6g256v2_wmma_gfx11_bt12", "gemm_qkvza_mq6g256v2_wmma_gfx11_bt4", - ]); - add!("qwen4_silu_mul", kernels::QWEN4_SILU_MUL_SRC, ["shared_expert_activation_bf16_f32", "silu_mul_bf16_rt_f32"]); - add!("tensor_ops", crate::tensor_ops::TENSOR_OPS_SRC, [ - "argmax_f32", "bf16_roundtrip_f32", "bf16_scaled_add_batched_f32", "bf16_scaled_add_f32", "copy_regions_u32", - "gated_delta_conv_bf16_f32", "gated_delta_conv_bf16_f32_batched_k4", "gated_delta_conv_params_bf16_f32", - "gated_delta_conv_qknorm_bf16_f32_batched_k4", "gated_delta_gate_bf16_f32", "gated_delta_gate_bf16_f32_batched", - "gated_delta_gate_rotate_bf16_f32_batched", "gated_delta_params_bf16_f32", "gated_delta_params_bf16_f32_batched", - "gated_delta_params_f32", "gated_delta_qk_norm_bf16_batched", "gated_delta_rollback_layers_f32", "gated_delta_rollback_layers_q8", - "gated_delta_step_f32", "gated_delta_step_gate_norm128_gfx1151", "gated_delta_step_gate_norm128_q8", - "gated_delta_step_halves_state128_persistent256_capture_f32", "gated_delta_step_halves_state128_persistent256_capture_q8", - "gated_delta_step_halves_state128_persistent256_f32", "gated_delta_step_halves_state128_persistent256_q8", - "gated_delta_step_norm128_q8", "gated_delta_step_shared_norm128_gfx1151", "gdn_state_f32_to_q8", "gdn_state_q8_to_f32", - "hc_activation_fused_f32", "hc_state_bf16_add_f32", "hc_state_bf16_to_f32", "hyper_norm_f32", "hyper_norm_gate_f32", - "hyper_norm_gate_outputs", "hyper_norm_gate_outputs_f32", "hyper_read_f32", "hyper_read_projected_f32", "hyper_read_up_fused_f32", - "hyper_write_bf16x2", "hyper_write_f32", "indexed_attention_attention_f32", "indexed_attention_attention_f32_batched", - "indexed_attention_attention_f32_batched_hg12", "indexed_attention_attention_f32_batched_hg4", - "indexed_attention_attention_f32_batched_serial", - "indexed_attention_attention_f32_serial", "indexed_attention_cache_append_f32", "indexed_attention_cache_append_f32_batched", - "indexed_attention_decode_prologue_f32", "indexed_attention_index_key_append_bf16_batched", "indexed_attention_norm_rope_f32", - "indexed_attention_norm_rope_f32_batched", "indexed_attention_pool_rope_bf16", "indexed_attention_pool_rope_f32", - "indexed_attention_reuse_selection", "indexed_attention_select_bf16_batched", "indexed_attention_select_bf16_batched_serial", - "indexed_attention_select_f32", "indexed_attention_select_f32_batched", "indexed_attention_select_f32_batched_serial", - "indexed_attention_select_f32_serial", "indexed_attention_select_from_scores", "indexed_attention_select_scores_rows16_f32", - "indexed_attention_select_scores_rows8_f32", "scale_f32", - ]); + kernels::QWEN4_GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, + [ + "gemm_mq4g256v2_moe_grouped_wmma_k2", + "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", + "gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", + ] + ); + add!( + "gemm_mq6g256v2_residual_wmma", + kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_SRC, + ["gemm_mq6g256v2_residual_wmma"] + ); + add!( + "hyper_read_up_wmma", + crate::tensor_ops::HYPER_READ_UP_WMMA_SRC, + ["hyper_read_up_wmma_bf16", "hyper_read_up_wmma_bf16_swap"] + ); + add!( + "indexed_attention_dense_wmma", + crate::tensor_ops::INDEXED_ATTENTION_DENSE_WMMA_SRC, + [ + "indexed_attention_dense_wmma_f16", + "indexed_attention_kv_f16", + ] + ); + add!( + "qwen4_bf16_round_trip", + kernels::QWEN4_BF16_ROUND_TRIP_SRC, + ["bf16_round_trip_f32_strided"] + ); + add!( + "qwen4_gemm_mqv2_wmma_gfx11_bt", + kernels::QWEN4_GEMM_MQV2_WMMA_GFX11_BT_SRC, + [ + "gemm_gate_up_mq2g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq3g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq3g256v2_wmma_gfx11_bt6", + "gemm_gate_up_mq5g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq5g256v2_wmma_gfx11_bt6", + "gemm_gate_up_mq6g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq6g256v2_wmma_gfx11_bt6", + "gemm_mq2g256v2_residual_wmma_gfx11_bt4", + "gemm_mq3g256v2_residual_wmma_gfx11_bt4", + "gemm_mq3g256v2_residual_wmma_gfx11_bt6", + "gemm_mq3g256v2_residual_wmma_gfx11_bt8", + "gemm_mq5g256v2_residual_wmma_gfx11_bt4", + "gemm_mq5g256v2_residual_wmma_gfx11_bt6", + "gemm_mq5g256v2_residual_wmma_gfx11_bt8", + "gemm_mq6g256v2_residual_wmma_gfx11_bt4", + "gemm_mq6g256v2_residual_wmma_gfx11_bt6", + "gemm_mq6g256v2_residual_wmma_gfx11_bt8", + "gemm_mq6g256v2_residual_wmma_gfx11_bt8_x4", + "gemm_mq6g256v2_wmma_gfx11_bt8_x4", + "gemm_mq6g256v2_wmma_gfx11_bt8_x4_bf16out", + "gemm_mq6g256v2_wmma_gfx11_bt8_x4_hcw", + "gemm_mq6g256v2_wmma_gfx11_bt8_x4_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p1", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p1_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p2", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r4_p2_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p1", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p1_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p2", + "gemm_mq6g256v2_wmma_gfx11_u3_b12_r8_p2_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p1_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r4_p2_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p1_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b16_r8_p2_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_dq_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p1_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_dq", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_dq_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r4_p2_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_dq", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_dq_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p1_regions", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_dq", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_dq_bf16out", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_hcw", + "gemm_mq6g256v2_wmma_gfx11_u3_b8_r8_p2_regions", + "gemm_qkv_mq2g256v2_wmma_gfx11_bt4", + "gemm_qkv_mq3g256v2_wmma_gfx11_bt12", + "gemm_qkv_mq3g256v2_wmma_gfx11_bt4", + "gemm_qkv_mq5g256v2_wmma_gfx11_bt12", + "gemm_qkv_mq5g256v2_wmma_gfx11_bt4", + "gemm_qkv_mq6g256v2_wmma_gfx11_bt12", + "gemm_qkv_mq6g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq2g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq3g256v2_wmma_gfx11_bt12", + "gemm_qkvza_mq3g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq5g256v2_wmma_gfx11_bt12", + "gemm_qkvza_mq5g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq6g256v2_wmma_gfx11_bt12", + "gemm_qkvza_mq6g256v2_wmma_gfx11_bt4", + ] + ); + add!( + "qwen4_silu_mul", + kernels::QWEN4_SILU_MUL_SRC, + ["shared_expert_activation_bf16_f32", "silu_mul_bf16_rt_f32"] + ); + add!( + "tensor_ops", + crate::tensor_ops::TENSOR_OPS_SRC, + [ + "argmax_f32", + "bf16_roundtrip_f32", + "bf16_scaled_add_batched_f32", + "bf16_scaled_add_f32", + "copy_regions_u32", + "gated_delta_conv_bf16_f32", + "gated_delta_conv_bf16_f32_batched_k4", + "gated_delta_conv_params_bf16_f32", + "gated_delta_conv_qknorm_bf16_f32_batched_k4", + "gated_delta_gate_bf16_f32", + "gated_delta_gate_bf16_f32_batched", + "gated_delta_gate_rotate_bf16_f32_batched", + "gated_delta_params_bf16_f32", + "gated_delta_params_bf16_f32_batched", + "gated_delta_params_f32", + "gated_delta_qk_norm_bf16_batched", + "gated_delta_rollback_layers_f32", + "gated_delta_rollback_layers_q8", + "gated_delta_step_f32", + "gated_delta_step_gate_norm128_gfx1151", + "gated_delta_step_gate_norm128_q8", + "gated_delta_step_halves_state128_persistent256_capture_f32", + "gated_delta_step_halves_state128_persistent256_capture_q8", + "gated_delta_step_halves_state128_persistent256_f32", + "gated_delta_step_halves_state128_persistent256_q8", + "gated_delta_step_norm128_q8", + "gated_delta_step_shared_norm128_gfx1151", + "gdn_state_f32_to_q8", + "gdn_state_q8_to_f32", + "hc_activation_fused_f32", + "hc_state_bf16_add_f32", + "hc_state_bf16_to_f32", + "hyper_norm_f32", + "hyper_norm_gate_f32", + "hyper_norm_gate_outputs", + "hyper_norm_gate_outputs_f32", + "hyper_read_f32", + "hyper_read_projected_f32", + "hyper_read_up_fused_f32", + "hyper_write_bf16x2", + "hyper_write_f32", + "indexed_attention_attention_f32", + "indexed_attention_attention_f32_batched", + "indexed_attention_attention_f32_batched_hg12", + "indexed_attention_attention_f32_batched_hg4", + "indexed_attention_attention_f32_batched_serial", + "indexed_attention_attention_f32_serial", + "indexed_attention_cache_append_f32", + "indexed_attention_cache_append_f32_batched", + "indexed_attention_decode_prologue_f32", + "indexed_attention_index_key_append_bf16_batched", + "indexed_attention_norm_rope_f32", + "indexed_attention_norm_rope_f32_batched", + "indexed_attention_pool_rope_bf16", + "indexed_attention_pool_rope_f32", + "indexed_attention_reuse_selection", + "indexed_attention_select_bf16_batched", + "indexed_attention_select_bf16_batched_serial", + "indexed_attention_select_f32", + "indexed_attention_select_f32_batched", + "indexed_attention_select_f32_batched_serial", + "indexed_attention_select_f32_serial", + "indexed_attention_select_from_scores", + "indexed_attention_select_scores_rows16_f32", + "indexed_attention_select_scores_rows8_f32", + "scale_f32", + ] + ); } if arch == "gfx1201" { - add!("gemm_mq4g256v2_moe_grouped_wmma_gfx12", kernels::GEMM_MQ4G256V2_MOE_GROUPED_WMMA_GFX12_SRC, ["gemm_mq4g256v2_moe_grouped_wmma_gfx12"]); - add!("gemm_mq6g256v2_residual_wmma_gfx12_bt12_mq5v2", kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_BT_SRC, [ - "gemm_mq6g256v2_residual_wmma_gfx12_bt12", "gemm_mq6g256v2_residual_wmma_gfx12_bt4", "gemm_mq6g256v2_residual_wmma_gfx12_bt8", - "gemm_mq6g256v2_wmma_gfx12_bt12_hcw", "gemm_mq6g256v2_wmma_gfx12_bt4_hcw", "gemm_mq6g256v2_wmma_gfx12_bt8_hcw", - ]); - add!("gemm_mq6g256v2_residual_wmma_gfx12_bt4_mq5v2", kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_BT_SRC, [ - "gemm_mq6g256v2_residual_wmma_gfx12_bt12", "gemm_mq6g256v2_residual_wmma_gfx12_bt4", "gemm_mq6g256v2_residual_wmma_gfx12_bt8", - "gemm_mq6g256v2_wmma_gfx12_bt12_hcw", "gemm_mq6g256v2_wmma_gfx12_bt4_hcw", "gemm_mq6g256v2_wmma_gfx12_bt8_hcw", - ]); - add!("gemm_mq6g256v2_residual_wmma_gfx12_bt8_mq5v2", kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_BT_SRC, [ - "gemm_mq6g256v2_residual_wmma_gfx12_bt12", "gemm_mq6g256v2_residual_wmma_gfx12_bt4", "gemm_mq6g256v2_residual_wmma_gfx12_bt8", - "gemm_mq6g256v2_wmma_gfx12_bt12_hcw", "gemm_mq6g256v2_wmma_gfx12_bt4_hcw", "gemm_mq6g256v2_wmma_gfx12_bt8_hcw", - ]); - add!("gemm_mq6g256v2_residual_wmma_gfx12_mq5v2", kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_SRC, ["gemm_mq6g256v2_residual_wmma_gfx12"]); - add!("qwen4_gemm_mq6g256v2_wmma_gfx12_x4", kernels::QWEN4_GEMM_MQ6G256V2_WMMA_GFX12_X4_SRC, [ - "gemm_mq6g256v2_residual_wmma_gfx12_bt8_x4", "gemm_mq6g256v2_wmma_gfx12_bt8_x4", "gemm_mq6g256v2_wmma_gfx12_bt8_x4_bf16out", - "gemm_mq6g256v2_wmma_gfx12_bt8_x4_regions", - ]); - add!("tensor_ops", crate::tensor_ops::TENSOR_OPS_SRC, [ - "argmax_f32", "bf16_roundtrip_f32", "bf16_scaled_add_batched_f32", "bf16_scaled_add_f32", "copy_regions_u32", - "gated_delta_conv_bf16_f32", "gated_delta_conv_bf16_f32_batched_k4", "gated_delta_conv_params_bf16_f32", - "gated_delta_conv_qknorm_bf16_f32_batched_k4", "gated_delta_gate_bf16_f32", "gated_delta_gate_bf16_f32_batched", - "gated_delta_gate_rotate_bf16_f32_batched", "gated_delta_params_bf16_f32", "gated_delta_params_bf16_f32_batched", - "gated_delta_params_f32", "gated_delta_qk_norm_bf16_batched", "gated_delta_rollback_layers_f32", "gated_delta_rollback_layers_q8", - "gated_delta_step_f32", "gated_delta_step_gate_norm128_gfx1151", "gated_delta_step_gate_norm128_q8", - "gated_delta_step_halves_state128_persistent256_capture_f32", "gated_delta_step_halves_state128_persistent256_capture_q8", - "gated_delta_step_halves_state128_persistent256_f32", "gated_delta_step_halves_state128_persistent256_q8", - "gated_delta_step_norm128_q8", "gated_delta_step_pipe128_f32", "gated_delta_step_pipe128_q8", - "gated_delta_step_shared_norm128_gfx1151", "gdn_state_f32_to_q8", "gdn_state_q8_to_f32", - "hc_activation_fused_f32", "hc_state_bf16_add_f32", "hc_state_bf16_to_f32", "hyper_norm_f32", "hyper_norm_gate_f32", - "hyper_norm_gate_outputs", "hyper_norm_gate_outputs_f32", "hyper_read_f32", "hyper_read_projected_f32", "hyper_read_up_fused_f32", - "hyper_write_bf16x2", "hyper_write_f32", "indexed_attention_attention_f32", "indexed_attention_attention_f32_batched", - "indexed_attention_attention_f32_batched_hg4", "indexed_attention_attention_f32_batched_serial", - "indexed_attention_attention_f32_serial", "indexed_attention_attention_fp8", "indexed_attention_attention_fp8_batched", - "indexed_attention_attention_fp8_batched_hg4", "indexed_attention_attention_fp8_batched_serial", - "indexed_attention_attention_fp8_serial", "indexed_attention_cache_append_f32", "indexed_attention_cache_append_f32_batched", - "indexed_attention_cache_append_fp8_batched", "indexed_attention_decode_prologue_f32", "indexed_attention_decode_prologue_fp8", - "indexed_attention_index_key_append_bf16_batched", "indexed_attention_norm_rope_f32", "indexed_attention_norm_rope_f32_batched", - "indexed_attention_pool_rope_bf16", "indexed_attention_pool_rope_f32", "indexed_attention_reuse_selection", - "indexed_attention_select_bf16_batched", "indexed_attention_select_bf16_batched_serial", "indexed_attention_select_f32", - "indexed_attention_select_f32_batched", "indexed_attention_select_f32_batched_serial", "indexed_attention_select_f32_serial", - "scale_f32", - ]); + add!( + "gemm_mq4g256v2_moe_grouped_wmma_gfx12", + kernels::GEMM_MQ4G256V2_MOE_GROUPED_WMMA_GFX12_SRC, + ["gemm_mq4g256v2_moe_grouped_wmma_gfx12"] + ); + add!( + "gemm_mq6g256v2_residual_wmma_gfx12_bt12_mq5v2", + kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_BT_SRC, + [ + "gemm_mq6g256v2_residual_wmma_gfx12_bt12", + "gemm_mq6g256v2_residual_wmma_gfx12_bt4", + "gemm_mq6g256v2_residual_wmma_gfx12_bt8", + "gemm_mq6g256v2_wmma_gfx12_bt12_hcw", + "gemm_mq6g256v2_wmma_gfx12_bt4_hcw", + "gemm_mq6g256v2_wmma_gfx12_bt8_hcw", + ] + ); + add!( + "gemm_mq6g256v2_residual_wmma_gfx12_bt4_mq5v2", + kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_BT_SRC, + [ + "gemm_mq6g256v2_residual_wmma_gfx12_bt12", + "gemm_mq6g256v2_residual_wmma_gfx12_bt4", + "gemm_mq6g256v2_residual_wmma_gfx12_bt8", + "gemm_mq6g256v2_wmma_gfx12_bt12_hcw", + "gemm_mq6g256v2_wmma_gfx12_bt4_hcw", + "gemm_mq6g256v2_wmma_gfx12_bt8_hcw", + ] + ); + add!( + "gemm_mq6g256v2_residual_wmma_gfx12_bt8_mq5v2", + kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_BT_SRC, + [ + "gemm_mq6g256v2_residual_wmma_gfx12_bt12", + "gemm_mq6g256v2_residual_wmma_gfx12_bt4", + "gemm_mq6g256v2_residual_wmma_gfx12_bt8", + "gemm_mq6g256v2_wmma_gfx12_bt12_hcw", + "gemm_mq6g256v2_wmma_gfx12_bt4_hcw", + "gemm_mq6g256v2_wmma_gfx12_bt8_hcw", + ] + ); + add!( + "gemm_mq6g256v2_residual_wmma_gfx12_mq5v2", + kernels::GEMM_MQ6G256V2_RESIDUAL_WMMA_GFX12_SRC, + ["gemm_mq6g256v2_residual_wmma_gfx12"] + ); + add!( + "qwen4_gemm_mq6g256v2_wmma_gfx12_x4", + kernels::QWEN4_GEMM_MQ6G256V2_WMMA_GFX12_X4_SRC, + [ + "gemm_mq6g256v2_residual_wmma_gfx12_bt8_x4", + "gemm_mq6g256v2_wmma_gfx12_bt8_x4", + "gemm_mq6g256v2_wmma_gfx12_bt8_x4_bf16out", + "gemm_mq6g256v2_wmma_gfx12_bt8_x4_regions", + ] + ); + add!( + "tensor_ops", + crate::tensor_ops::TENSOR_OPS_SRC, + [ + "argmax_f32", + "bf16_roundtrip_f32", + "bf16_scaled_add_batched_f32", + "bf16_scaled_add_f32", + "copy_regions_u32", + "gated_delta_conv_bf16_f32", + "gated_delta_conv_bf16_f32_batched_k4", + "gated_delta_conv_params_bf16_f32", + "gated_delta_conv_qknorm_bf16_f32_batched_k4", + "gated_delta_gate_bf16_f32", + "gated_delta_gate_bf16_f32_batched", + "gated_delta_gate_rotate_bf16_f32_batched", + "gated_delta_params_bf16_f32", + "gated_delta_params_bf16_f32_batched", + "gated_delta_params_f32", + "gated_delta_qk_norm_bf16_batched", + "gated_delta_rollback_layers_f32", + "gated_delta_rollback_layers_q8", + "gated_delta_step_f32", + "gated_delta_step_gate_norm128_gfx1151", + "gated_delta_step_gate_norm128_q8", + "gated_delta_step_halves_state128_persistent256_capture_f32", + "gated_delta_step_halves_state128_persistent256_capture_q8", + "gated_delta_step_halves_state128_persistent256_f32", + "gated_delta_step_halves_state128_persistent256_q8", + "gated_delta_step_norm128_q8", + "gated_delta_step_pipe128_f32", + "gated_delta_step_pipe128_q8", + "gated_delta_step_shared_norm128_gfx1151", + "gdn_state_f32_to_q8", + "gdn_state_q8_to_f32", + "hc_activation_fused_f32", + "hc_state_bf16_add_f32", + "hc_state_bf16_to_f32", + "hyper_norm_f32", + "hyper_norm_gate_f32", + "hyper_norm_gate_outputs", + "hyper_norm_gate_outputs_f32", + "hyper_read_f32", + "hyper_read_projected_f32", + "hyper_read_up_fused_f32", + "hyper_write_bf16x2", + "hyper_write_f32", + "indexed_attention_attention_f32", + "indexed_attention_attention_f32_batched", + "indexed_attention_attention_f32_batched_hg4", + "indexed_attention_attention_f32_batched_serial", + "indexed_attention_attention_f32_serial", + "indexed_attention_attention_fp8", + "indexed_attention_attention_fp8_batched", + "indexed_attention_attention_fp8_batched_hg4", + "indexed_attention_attention_fp8_batched_serial", + "indexed_attention_attention_fp8_serial", + "indexed_attention_cache_append_f32", + "indexed_attention_cache_append_f32_batched", + "indexed_attention_cache_append_fp8_batched", + "indexed_attention_decode_prologue_f32", + "indexed_attention_decode_prologue_fp8", + "indexed_attention_index_key_append_bf16_batched", + "indexed_attention_norm_rope_f32", + "indexed_attention_norm_rope_f32_batched", + "indexed_attention_pool_rope_bf16", + "indexed_attention_pool_rope_f32", + "indexed_attention_reuse_selection", + "indexed_attention_select_bf16_batched", + "indexed_attention_select_bf16_batched_serial", + "indexed_attention_select_f32", + "indexed_attention_select_f32_batched", + "indexed_attention_select_f32_batched_serial", + "indexed_attention_select_f32_serial", + "scale_f32", + ] + ); } if arch == "gfx1100" { - add!("add", kernels::ADD_SRC, ["add_f32", "broadcast_add_rows_f32"]); - add!("gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", kernels::QWEN4_GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, [ - "gemm_mq4g256v2_moe_grouped_wmma_k2", "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", - "gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", - ]); - add!("gemm_mqv2_wmma_gfx1100_bt", kernels::gemm_mqv2_wmma_gfx11_bt_src(arch == "gfx1151"), [ - "gemm_gate_up_mq2g256v2_wmma_gfx11_bt12", "gemm_gate_up_mq3g256v2_wmma_gfx11_bt12", "gemm_gate_up_mq3g256v2_wmma_gfx11_bt6", - "gemm_gate_up_mq5g256v2_wmma_gfx11_bt12", "gemm_gate_up_mq5g256v2_wmma_gfx11_bt6", "gemm_gate_up_mq6g256v2_wmma_gfx11_bt12", - "gemm_gate_up_mq6g256v2_wmma_gfx11_bt6", "gemm_mq2g256v2_residual_wmma_gfx11_bt4", "gemm_mq3g256v2_residual_wmma_gfx11_bt4", - "gemm_mq3g256v2_residual_wmma_gfx11_bt6", "gemm_mq3g256v2_residual_wmma_gfx11_bt8", "gemm_mq5g256v2_residual_wmma_gfx11_bt4", - "gemm_mq5g256v2_residual_wmma_gfx11_bt6", "gemm_mq5g256v2_residual_wmma_gfx11_bt8", "gemm_mq6g256v2_residual_wmma_gfx11_bt4", - "gemm_mq6g256v2_residual_wmma_gfx11_bt6", "gemm_mq6g256v2_residual_wmma_gfx11_bt8", "gemm_qkv_mq2g256v2_wmma_gfx11_bt4", - "gemm_qkv_mq3g256v2_wmma_gfx11_bt12", "gemm_qkv_mq3g256v2_wmma_gfx11_bt4", "gemm_qkv_mq5g256v2_wmma_gfx11_bt12", - "gemm_qkv_mq5g256v2_wmma_gfx11_bt4", "gemm_qkv_mq6g256v2_wmma_gfx11_bt12", "gemm_qkv_mq6g256v2_wmma_gfx11_bt4", - "gemm_qkvza_mq2g256v2_wmma_gfx11_bt4", "gemm_qkvza_mq3g256v2_wmma_gfx11_bt12", "gemm_qkvza_mq3g256v2_wmma_gfx11_bt4", - "gemm_qkvza_mq5g256v2_wmma_gfx11_bt12", "gemm_qkvza_mq5g256v2_wmma_gfx11_bt4", "gemm_qkvza_mq6g256v2_wmma_gfx11_bt12", - "gemm_qkvza_mq6g256v2_wmma_gfx11_bt4", - ]); - add!("gemm_wmma_lds256", kernels::gemm_f16_x_f16_wmma_lds256_src(arch == "gfx1151"), [ - "gemm_wmma_lds_128_128_32_64_k64", "gemm_wmma_lds_128_128_32_64_k64_a", "gemm_wmma_lds_128_128_32_64_k64_gr", - "gemm_wmma_lds_128_128_32_64_k64_gra", "gemm_wmma_lds_128_128_32_64_k64_o16", "gemm_wmma_lds_128_128_32_64_k64_o16g", - "gemm_wmma_lds_128_128_32_64_k64_p", "gemm_wmma_lds_128_128_32_64_k64_p_a", "gemm_wmma_lds_128_128_32_64_k64_p_gr", - "gemm_wmma_lds_128_128_32_64_k64_p_gra", "gemm_wmma_lds_128_128_32_64_k64_p_o16", "gemm_wmma_lds_128_128_32_64_k64_p_o16g", - "gemm_wmma_lds_128_128_64_64_k64", "gemm_wmma_lds_128_128_64_64_k64_a", "gemm_wmma_lds_128_128_64_64_k64_gr", - "gemm_wmma_lds_128_128_64_64_k64_gra", "gemm_wmma_lds_128_128_64_64_k64_o16", "gemm_wmma_lds_128_128_64_64_k64_o16g", - "gemm_wmma_lds_128_256_32_64_k64", "gemm_wmma_lds_128_256_32_64_k64_a", "gemm_wmma_lds_128_256_32_64_k64_gr", - "gemm_wmma_lds_128_256_32_64_k64_gra", "gemm_wmma_lds_128_256_32_64_k64_o16", "gemm_wmma_lds_128_256_32_64_k64_o16g", - "gemm_wmma_lds_128_256_32_64_k64_p", "gemm_wmma_lds_128_256_32_64_k64_p_a", "gemm_wmma_lds_128_256_32_64_k64_p_gr", - "gemm_wmma_lds_128_256_32_64_k64_p_gra", "gemm_wmma_lds_128_256_32_64_k64_p_o16", "gemm_wmma_lds_128_256_32_64_k64_p_o16g", - "gemm_wmma_lds_128_256_64_64_k32", "gemm_wmma_lds_128_256_64_64_k32_a", "gemm_wmma_lds_128_256_64_64_k32_gr", - "gemm_wmma_lds_128_256_64_64_k32_gra", "gemm_wmma_lds_128_256_64_64_k32_o16", "gemm_wmma_lds_128_256_64_64_k32_o16g", - "gemm_wmma_lds_128_256_64_64_k32_p", "gemm_wmma_lds_128_256_64_64_k32_p_a", "gemm_wmma_lds_128_256_64_64_k32_p_gr", - "gemm_wmma_lds_128_256_64_64_k32_p_gra", "gemm_wmma_lds_128_256_64_64_k32_p_o16", "gemm_wmma_lds_128_256_64_64_k32_p_o16g", - "gemm_wmma_lds_128_256_64_64_k64", "gemm_wmma_lds_128_256_64_64_k64_sw", "gemm_wmma_lds_128_512_32_64_k32", - "gemm_wmma_lds_128_512_64_64_k32", "gemm_wmma_lds_128_512_64_64_k32_sw", "gemm_wmma_lds_256_128_32_64_k64", - "gemm_wmma_lds_256_128_64_64_k64", "gemm_wmma_lds_256_128_64_64_k64_sw", "gemm_wmma_lds_256_256_32_64_k64", - "gemm_wmma_lds_256_256_64_64_k64", "gemm_wmma_lds_256_256_64_64_k64_a", "gemm_wmma_lds_256_256_64_64_k64_gr", - "gemm_wmma_lds_256_256_64_64_k64_gra", "gemm_wmma_lds_256_256_64_64_k64_o16", "gemm_wmma_lds_256_256_64_64_k64_o16g", - "gemm_wmma_lds_256_256_64_64_k64_p", "gemm_wmma_lds_256_256_64_64_k64_p_a", "gemm_wmma_lds_256_256_64_64_k64_p_gr", - "gemm_wmma_lds_256_256_64_64_k64_p_gra", "gemm_wmma_lds_256_256_64_64_k64_p_o16", "gemm_wmma_lds_256_256_64_64_k64_p_o16g", - "gemm_wmma_lds_64_512_64_64_k32", - ]); - add!("gemv_mq2g256v2_rdna3", kernels::GEMV_MQ2G256V2_SRC, ["gemv_mq2g256v2"]); - add!("gemv_mq6g256v2_rdna3_mq6v2", kernels::GEMV_MQ6G256V2_SRC, ["gemv_mq6g256v2"]); - add!("qwen4_gemm_wmma_lds256", kernels::QWEN4_GEMM_F16_X_F16_WMMA_LDS256_SRC, [ - "gemm_f16_x_f16_wmma_lds_regions_128_128_32_64_k64_p", "gemm_wmma_lds_128_128_32_64_k64", "gemm_wmma_lds_128_128_32_64_k64_a", - "gemm_wmma_lds_128_128_32_64_k64_bsr", "gemm_wmma_lds_128_128_32_64_k64_gr", "gemm_wmma_lds_128_128_32_64_k64_gra", - "gemm_wmma_lds_128_128_32_64_k64_hcsd", "gemm_wmma_lds_128_128_32_64_k64_o16", "gemm_wmma_lds_128_128_32_64_k64_o16g", - "gemm_wmma_lds_128_128_32_64_k64_p", "gemm_wmma_lds_128_128_32_64_k64_p_a", "gemm_wmma_lds_128_128_32_64_k64_p_bsr", - "gemm_wmma_lds_128_128_32_64_k64_p_gr", "gemm_wmma_lds_128_128_32_64_k64_p_gra", "gemm_wmma_lds_128_128_32_64_k64_p_hcsd", - "gemm_wmma_lds_128_128_32_64_k64_p_o16", "gemm_wmma_lds_128_128_32_64_k64_p_o16g", "gemm_wmma_lds_128_128_64_64_k64", - "gemm_wmma_lds_128_128_64_64_k64_a", "gemm_wmma_lds_128_128_64_64_k64_bsr", "gemm_wmma_lds_128_128_64_64_k64_gr", - "gemm_wmma_lds_128_128_64_64_k64_gra", "gemm_wmma_lds_128_128_64_64_k64_o16", "gemm_wmma_lds_128_128_64_64_k64_o16g", - "gemm_wmma_lds_128_256_32_64_k64", "gemm_wmma_lds_128_256_32_64_k64_a", "gemm_wmma_lds_128_256_32_64_k64_bf16st", - "gemm_wmma_lds_128_256_32_64_k64_bsr", "gemm_wmma_lds_128_256_32_64_k64_gr", "gemm_wmma_lds_128_256_32_64_k64_gra", - "gemm_wmma_lds_128_256_32_64_k64_hcsd", "gemm_wmma_lds_128_256_32_64_k64_o16", "gemm_wmma_lds_128_256_32_64_k64_o16g", - "gemm_wmma_lds_128_256_32_64_k64_p", "gemm_wmma_lds_128_256_32_64_k64_p_a", "gemm_wmma_lds_128_256_32_64_k64_p_bsr", - "gemm_wmma_lds_128_256_32_64_k64_p_gr", "gemm_wmma_lds_128_256_32_64_k64_p_gra", "gemm_wmma_lds_128_256_32_64_k64_p_hcsd", - "gemm_wmma_lds_128_256_32_64_k64_p_o16", "gemm_wmma_lds_128_256_32_64_k64_p_o16g", "gemm_wmma_lds_128_256_64_64_k32", - "gemm_wmma_lds_128_256_64_64_k32_a", "gemm_wmma_lds_128_256_64_64_k32_bsr", "gemm_wmma_lds_128_256_64_64_k32_gr", - "gemm_wmma_lds_128_256_64_64_k32_gra", "gemm_wmma_lds_128_256_64_64_k32_o16", "gemm_wmma_lds_128_256_64_64_k32_o16g", - "gemm_wmma_lds_128_256_64_64_k32_p", "gemm_wmma_lds_128_256_64_64_k32_p_a", "gemm_wmma_lds_128_256_64_64_k32_p_bsr", - "gemm_wmma_lds_128_256_64_64_k32_p_gr", "gemm_wmma_lds_128_256_64_64_k32_p_gra", "gemm_wmma_lds_128_256_64_64_k32_p_o16", - "gemm_wmma_lds_128_256_64_64_k32_p_o16g", "gemm_wmma_lds_128_256_64_64_k64", "gemm_wmma_lds_128_256_64_64_k64_sw", - "gemm_wmma_lds_128_512_32_64_k32", "gemm_wmma_lds_128_512_64_64_k32", "gemm_wmma_lds_128_512_64_64_k32_sw", - "gemm_wmma_lds_160_64_32_64_k64", "gemm_wmma_lds_160_64_32_64_k64_p", "gemm_wmma_lds_256_128_32_64_k64", - "gemm_wmma_lds_256_128_64_64_k64", "gemm_wmma_lds_256_128_64_64_k64_sw", "gemm_wmma_lds_256_256_32_64_k64", - "gemm_wmma_lds_256_256_64_64_k64", "gemm_wmma_lds_256_256_64_64_k64_a", "gemm_wmma_lds_256_256_64_64_k64_bsr", - "gemm_wmma_lds_256_256_64_64_k64_gr", "gemm_wmma_lds_256_256_64_64_k64_gra", "gemm_wmma_lds_256_256_64_64_k64_hcsd", - "gemm_wmma_lds_256_256_64_64_k64_o16", "gemm_wmma_lds_256_256_64_64_k64_o16g", "gemm_wmma_lds_256_256_64_64_k64_p", - "gemm_wmma_lds_256_256_64_64_k64_p_a", "gemm_wmma_lds_256_256_64_64_k64_p_bsr", "gemm_wmma_lds_256_256_64_64_k64_p_gr", - "gemm_wmma_lds_256_256_64_64_k64_p_gra", "gemm_wmma_lds_256_256_64_64_k64_p_hcsd", "gemm_wmma_lds_256_256_64_64_k64_p_o16", - "gemm_wmma_lds_256_256_64_64_k64_p_o16g", "gemm_wmma_lds_64_512_64_64_k32", "gemm_wmma_lds_64_64_32_64_k64", - "gemm_wmma_lds_64_64_32_64_k64_p", - ]); + add!( + "add", + kernels::ADD_SRC, + ["add_f32", "broadcast_add_rows_f32"] + ); + add!( + "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", + kernels::QWEN4_GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, + [ + "gemm_mq4g256v2_moe_grouped_wmma_k2", + "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", + "gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", + ] + ); + add!( + "gemm_mqv2_wmma_gfx1100_bt", + kernels::gemm_mqv2_wmma_gfx11_bt_src(arch == "gfx1151"), + [ + "gemm_gate_up_mq2g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq3g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq3g256v2_wmma_gfx11_bt6", + "gemm_gate_up_mq5g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq5g256v2_wmma_gfx11_bt6", + "gemm_gate_up_mq6g256v2_wmma_gfx11_bt12", + "gemm_gate_up_mq6g256v2_wmma_gfx11_bt6", + "gemm_mq2g256v2_residual_wmma_gfx11_bt4", + "gemm_mq3g256v2_residual_wmma_gfx11_bt4", + "gemm_mq3g256v2_residual_wmma_gfx11_bt6", + "gemm_mq3g256v2_residual_wmma_gfx11_bt8", + "gemm_mq5g256v2_residual_wmma_gfx11_bt4", + "gemm_mq5g256v2_residual_wmma_gfx11_bt6", + "gemm_mq5g256v2_residual_wmma_gfx11_bt8", + "gemm_mq6g256v2_residual_wmma_gfx11_bt4", + "gemm_mq6g256v2_residual_wmma_gfx11_bt6", + "gemm_mq6g256v2_residual_wmma_gfx11_bt8", + "gemm_qkv_mq2g256v2_wmma_gfx11_bt4", + "gemm_qkv_mq3g256v2_wmma_gfx11_bt12", + "gemm_qkv_mq3g256v2_wmma_gfx11_bt4", + "gemm_qkv_mq5g256v2_wmma_gfx11_bt12", + "gemm_qkv_mq5g256v2_wmma_gfx11_bt4", + "gemm_qkv_mq6g256v2_wmma_gfx11_bt12", + "gemm_qkv_mq6g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq2g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq3g256v2_wmma_gfx11_bt12", + "gemm_qkvza_mq3g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq5g256v2_wmma_gfx11_bt12", + "gemm_qkvza_mq5g256v2_wmma_gfx11_bt4", + "gemm_qkvza_mq6g256v2_wmma_gfx11_bt12", + "gemm_qkvza_mq6g256v2_wmma_gfx11_bt4", + ] + ); + add!( + "gemm_wmma_lds256", + kernels::gemm_f16_x_f16_wmma_lds256_src(arch == "gfx1151"), + [ + "gemm_wmma_lds_128_128_32_64_k64", + "gemm_wmma_lds_128_128_32_64_k64_a", + "gemm_wmma_lds_128_128_32_64_k64_gr", + "gemm_wmma_lds_128_128_32_64_k64_gra", + "gemm_wmma_lds_128_128_32_64_k64_o16", + "gemm_wmma_lds_128_128_32_64_k64_o16g", + "gemm_wmma_lds_128_128_32_64_k64_p", + "gemm_wmma_lds_128_128_32_64_k64_p_a", + "gemm_wmma_lds_128_128_32_64_k64_p_gr", + "gemm_wmma_lds_128_128_32_64_k64_p_gra", + "gemm_wmma_lds_128_128_32_64_k64_p_o16", + "gemm_wmma_lds_128_128_32_64_k64_p_o16g", + "gemm_wmma_lds_128_128_64_64_k64", + "gemm_wmma_lds_128_128_64_64_k64_a", + "gemm_wmma_lds_128_128_64_64_k64_gr", + "gemm_wmma_lds_128_128_64_64_k64_gra", + "gemm_wmma_lds_128_128_64_64_k64_o16", + "gemm_wmma_lds_128_128_64_64_k64_o16g", + "gemm_wmma_lds_128_256_32_64_k64", + "gemm_wmma_lds_128_256_32_64_k64_a", + "gemm_wmma_lds_128_256_32_64_k64_gr", + "gemm_wmma_lds_128_256_32_64_k64_gra", + "gemm_wmma_lds_128_256_32_64_k64_o16", + "gemm_wmma_lds_128_256_32_64_k64_o16g", + "gemm_wmma_lds_128_256_32_64_k64_p", + "gemm_wmma_lds_128_256_32_64_k64_p_a", + "gemm_wmma_lds_128_256_32_64_k64_p_gr", + "gemm_wmma_lds_128_256_32_64_k64_p_gra", + "gemm_wmma_lds_128_256_32_64_k64_p_o16", + "gemm_wmma_lds_128_256_32_64_k64_p_o16g", + "gemm_wmma_lds_128_256_64_64_k32", + "gemm_wmma_lds_128_256_64_64_k32_a", + "gemm_wmma_lds_128_256_64_64_k32_gr", + "gemm_wmma_lds_128_256_64_64_k32_gra", + "gemm_wmma_lds_128_256_64_64_k32_o16", + "gemm_wmma_lds_128_256_64_64_k32_o16g", + "gemm_wmma_lds_128_256_64_64_k32_p", + "gemm_wmma_lds_128_256_64_64_k32_p_a", + "gemm_wmma_lds_128_256_64_64_k32_p_gr", + "gemm_wmma_lds_128_256_64_64_k32_p_gra", + "gemm_wmma_lds_128_256_64_64_k32_p_o16", + "gemm_wmma_lds_128_256_64_64_k32_p_o16g", + "gemm_wmma_lds_128_256_64_64_k64", + "gemm_wmma_lds_128_256_64_64_k64_sw", + "gemm_wmma_lds_128_512_32_64_k32", + "gemm_wmma_lds_128_512_64_64_k32", + "gemm_wmma_lds_128_512_64_64_k32_sw", + "gemm_wmma_lds_256_128_32_64_k64", + "gemm_wmma_lds_256_128_64_64_k64", + "gemm_wmma_lds_256_128_64_64_k64_sw", + "gemm_wmma_lds_256_256_32_64_k64", + "gemm_wmma_lds_256_256_64_64_k64", + "gemm_wmma_lds_256_256_64_64_k64_a", + "gemm_wmma_lds_256_256_64_64_k64_gr", + "gemm_wmma_lds_256_256_64_64_k64_gra", + "gemm_wmma_lds_256_256_64_64_k64_o16", + "gemm_wmma_lds_256_256_64_64_k64_o16g", + "gemm_wmma_lds_256_256_64_64_k64_p", + "gemm_wmma_lds_256_256_64_64_k64_p_a", + "gemm_wmma_lds_256_256_64_64_k64_p_gr", + "gemm_wmma_lds_256_256_64_64_k64_p_gra", + "gemm_wmma_lds_256_256_64_64_k64_p_o16", + "gemm_wmma_lds_256_256_64_64_k64_p_o16g", + "gemm_wmma_lds_64_512_64_64_k32", + ] + ); + add!( + "gemv_mq2g256v2_rdna3", + kernels::GEMV_MQ2G256V2_SRC, + ["gemv_mq2g256v2"] + ); + add!( + "gemv_mq6g256v2_rdna3_mq6v2", + kernels::GEMV_MQ6G256V2_SRC, + ["gemv_mq6g256v2"] + ); + add!( + "qwen4_gemm_wmma_lds256", + kernels::QWEN4_GEMM_F16_X_F16_WMMA_LDS256_SRC, + [ + "gemm_f16_x_f16_wmma_lds_regions_128_128_32_64_k64_p", + "gemm_wmma_lds_128_128_32_64_k64", + "gemm_wmma_lds_128_128_32_64_k64_a", + "gemm_wmma_lds_128_128_32_64_k64_bsr", + "gemm_wmma_lds_128_128_32_64_k64_gr", + "gemm_wmma_lds_128_128_32_64_k64_gra", + "gemm_wmma_lds_128_128_32_64_k64_hcsd", + "gemm_wmma_lds_128_128_32_64_k64_o16", + "gemm_wmma_lds_128_128_32_64_k64_o16g", + "gemm_wmma_lds_128_128_32_64_k64_p", + "gemm_wmma_lds_128_128_32_64_k64_p_a", + "gemm_wmma_lds_128_128_32_64_k64_p_bsr", + "gemm_wmma_lds_128_128_32_64_k64_p_gr", + "gemm_wmma_lds_128_128_32_64_k64_p_gra", + "gemm_wmma_lds_128_128_32_64_k64_p_hcsd", + "gemm_wmma_lds_128_128_32_64_k64_p_o16", + "gemm_wmma_lds_128_128_32_64_k64_p_o16g", + "gemm_wmma_lds_128_128_64_64_k64", + "gemm_wmma_lds_128_128_64_64_k64_a", + "gemm_wmma_lds_128_128_64_64_k64_bsr", + "gemm_wmma_lds_128_128_64_64_k64_gr", + "gemm_wmma_lds_128_128_64_64_k64_gra", + "gemm_wmma_lds_128_128_64_64_k64_o16", + "gemm_wmma_lds_128_128_64_64_k64_o16g", + "gemm_wmma_lds_128_256_32_64_k64", + "gemm_wmma_lds_128_256_32_64_k64_a", + "gemm_wmma_lds_128_256_32_64_k64_bf16st", + "gemm_wmma_lds_128_256_32_64_k64_bsr", + "gemm_wmma_lds_128_256_32_64_k64_gr", + "gemm_wmma_lds_128_256_32_64_k64_gra", + "gemm_wmma_lds_128_256_32_64_k64_hcsd", + "gemm_wmma_lds_128_256_32_64_k64_o16", + "gemm_wmma_lds_128_256_32_64_k64_o16g", + "gemm_wmma_lds_128_256_32_64_k64_p", + "gemm_wmma_lds_128_256_32_64_k64_p_a", + "gemm_wmma_lds_128_256_32_64_k64_p_bsr", + "gemm_wmma_lds_128_256_32_64_k64_p_gr", + "gemm_wmma_lds_128_256_32_64_k64_p_gra", + "gemm_wmma_lds_128_256_32_64_k64_p_hcsd", + "gemm_wmma_lds_128_256_32_64_k64_p_o16", + "gemm_wmma_lds_128_256_32_64_k64_p_o16g", + "gemm_wmma_lds_128_256_64_64_k32", + "gemm_wmma_lds_128_256_64_64_k32_a", + "gemm_wmma_lds_128_256_64_64_k32_bsr", + "gemm_wmma_lds_128_256_64_64_k32_gr", + "gemm_wmma_lds_128_256_64_64_k32_gra", + "gemm_wmma_lds_128_256_64_64_k32_o16", + "gemm_wmma_lds_128_256_64_64_k32_o16g", + "gemm_wmma_lds_128_256_64_64_k32_p", + "gemm_wmma_lds_128_256_64_64_k32_p_a", + "gemm_wmma_lds_128_256_64_64_k32_p_bsr", + "gemm_wmma_lds_128_256_64_64_k32_p_gr", + "gemm_wmma_lds_128_256_64_64_k32_p_gra", + "gemm_wmma_lds_128_256_64_64_k32_p_o16", + "gemm_wmma_lds_128_256_64_64_k32_p_o16g", + "gemm_wmma_lds_128_256_64_64_k64", + "gemm_wmma_lds_128_256_64_64_k64_sw", + "gemm_wmma_lds_128_512_32_64_k32", + "gemm_wmma_lds_128_512_64_64_k32", + "gemm_wmma_lds_128_512_64_64_k32_sw", + "gemm_wmma_lds_160_64_32_64_k64", + "gemm_wmma_lds_160_64_32_64_k64_p", + "gemm_wmma_lds_256_128_32_64_k64", + "gemm_wmma_lds_256_128_64_64_k64", + "gemm_wmma_lds_256_128_64_64_k64_sw", + "gemm_wmma_lds_256_256_32_64_k64", + "gemm_wmma_lds_256_256_64_64_k64", + "gemm_wmma_lds_256_256_64_64_k64_a", + "gemm_wmma_lds_256_256_64_64_k64_bsr", + "gemm_wmma_lds_256_256_64_64_k64_gr", + "gemm_wmma_lds_256_256_64_64_k64_gra", + "gemm_wmma_lds_256_256_64_64_k64_hcsd", + "gemm_wmma_lds_256_256_64_64_k64_o16", + "gemm_wmma_lds_256_256_64_64_k64_o16g", + "gemm_wmma_lds_256_256_64_64_k64_p", + "gemm_wmma_lds_256_256_64_64_k64_p_a", + "gemm_wmma_lds_256_256_64_64_k64_p_bsr", + "gemm_wmma_lds_256_256_64_64_k64_p_gr", + "gemm_wmma_lds_256_256_64_64_k64_p_gra", + "gemm_wmma_lds_256_256_64_64_k64_p_hcsd", + "gemm_wmma_lds_256_256_64_64_k64_p_o16", + "gemm_wmma_lds_256_256_64_64_k64_p_o16g", + "gemm_wmma_lds_64_512_64_64_k32", + "gemm_wmma_lds_64_64_32_64_k64", + "gemm_wmma_lds_64_64_32_64_k64_p", + ] + ); } if arch == "gfx1151" { - add!("gemm_wmma_lds256", kernels::gemm_f16_x_f16_wmma_lds256_src(arch == "gfx1151"), [ - "gemm_f16_x_f16_wmma_lds_regions_128_128_32_64_k64_p", "gemm_wmma_lds_128_128_32_64_k64", "gemm_wmma_lds_128_128_32_64_k64_a", - "gemm_wmma_lds_128_128_32_64_k64_bsr", "gemm_wmma_lds_128_128_32_64_k64_gr", "gemm_wmma_lds_128_128_32_64_k64_gra", - "gemm_wmma_lds_128_128_32_64_k64_hcsd", "gemm_wmma_lds_128_128_32_64_k64_o16", "gemm_wmma_lds_128_128_32_64_k64_o16g", - "gemm_wmma_lds_128_128_32_64_k64_p", "gemm_wmma_lds_128_128_32_64_k64_p_a", "gemm_wmma_lds_128_128_32_64_k64_p_bsr", - "gemm_wmma_lds_128_128_32_64_k64_p_gr", "gemm_wmma_lds_128_128_32_64_k64_p_gra", "gemm_wmma_lds_128_128_32_64_k64_p_hcsd", - "gemm_wmma_lds_128_128_32_64_k64_p_o16", "gemm_wmma_lds_128_128_32_64_k64_p_o16g", "gemm_wmma_lds_128_128_64_64_k64", - "gemm_wmma_lds_128_128_64_64_k64_a", "gemm_wmma_lds_128_128_64_64_k64_bsr", "gemm_wmma_lds_128_128_64_64_k64_gr", - "gemm_wmma_lds_128_128_64_64_k64_gra", "gemm_wmma_lds_128_128_64_64_k64_o16", "gemm_wmma_lds_128_128_64_64_k64_o16g", - "gemm_wmma_lds_128_256_32_64_k64", "gemm_wmma_lds_128_256_32_64_k64_a", "gemm_wmma_lds_128_256_32_64_k64_bf16st", - "gemm_wmma_lds_128_256_32_64_k64_bsr", "gemm_wmma_lds_128_256_32_64_k64_gr", "gemm_wmma_lds_128_256_32_64_k64_gra", - "gemm_wmma_lds_128_256_32_64_k64_hcsd", "gemm_wmma_lds_128_256_32_64_k64_o16", "gemm_wmma_lds_128_256_32_64_k64_o16g", - "gemm_wmma_lds_128_256_32_64_k64_p", "gemm_wmma_lds_128_256_32_64_k64_p_a", "gemm_wmma_lds_128_256_32_64_k64_p_bsr", - "gemm_wmma_lds_128_256_32_64_k64_p_gr", "gemm_wmma_lds_128_256_32_64_k64_p_gra", "gemm_wmma_lds_128_256_32_64_k64_p_hcsd", - "gemm_wmma_lds_128_256_32_64_k64_p_o16", "gemm_wmma_lds_128_256_32_64_k64_p_o16g", "gemm_wmma_lds_128_256_64_64_k32", - "gemm_wmma_lds_128_256_64_64_k32_a", "gemm_wmma_lds_128_256_64_64_k32_bsr", "gemm_wmma_lds_128_256_64_64_k32_gr", - "gemm_wmma_lds_128_256_64_64_k32_gra", "gemm_wmma_lds_128_256_64_64_k32_o16", "gemm_wmma_lds_128_256_64_64_k32_o16g", - "gemm_wmma_lds_128_256_64_64_k32_p", "gemm_wmma_lds_128_256_64_64_k32_p_a", "gemm_wmma_lds_128_256_64_64_k32_p_bsr", - "gemm_wmma_lds_128_256_64_64_k32_p_gr", "gemm_wmma_lds_128_256_64_64_k32_p_gra", "gemm_wmma_lds_128_256_64_64_k32_p_o16", - "gemm_wmma_lds_128_256_64_64_k32_p_o16g", "gemm_wmma_lds_128_256_64_64_k64", "gemm_wmma_lds_128_256_64_64_k64_sw", - "gemm_wmma_lds_128_512_32_64_k32", "gemm_wmma_lds_128_512_64_64_k32", "gemm_wmma_lds_128_512_64_64_k32_sw", - "gemm_wmma_lds_160_64_32_64_k64", "gemm_wmma_lds_160_64_32_64_k64_p", "gemm_wmma_lds_256_128_32_64_k64", - "gemm_wmma_lds_256_128_64_64_k64", "gemm_wmma_lds_256_128_64_64_k64_sw", "gemm_wmma_lds_256_256_32_64_k64", - "gemm_wmma_lds_256_256_64_64_k64", "gemm_wmma_lds_256_256_64_64_k64_a", "gemm_wmma_lds_256_256_64_64_k64_bsr", - "gemm_wmma_lds_256_256_64_64_k64_gr", "gemm_wmma_lds_256_256_64_64_k64_gra", "gemm_wmma_lds_256_256_64_64_k64_hcsd", - "gemm_wmma_lds_256_256_64_64_k64_o16", "gemm_wmma_lds_256_256_64_64_k64_o16g", "gemm_wmma_lds_256_256_64_64_k64_p", - "gemm_wmma_lds_256_256_64_64_k64_p_a", "gemm_wmma_lds_256_256_64_64_k64_p_bsr", "gemm_wmma_lds_256_256_64_64_k64_p_gr", - "gemm_wmma_lds_256_256_64_64_k64_p_gra", "gemm_wmma_lds_256_256_64_64_k64_p_hcsd", "gemm_wmma_lds_256_256_64_64_k64_p_o16", - "gemm_wmma_lds_256_256_64_64_k64_p_o16g", "gemm_wmma_lds_64_512_64_64_k32", "gemm_wmma_lds_64_64_32_64_k64", - "gemm_wmma_lds_64_64_32_64_k64_p", - ]); - add!("gemv_bf16_xf32", kernels::gemv_bf16_xf32_src(arch == "gfx1151"), [ - "gemv_bf16_xf32", "gemv_bf16_xf32_bf16_scaled_add", "gemv_bf16_xf32_k4", "gemv_bf16_xf32_k4_rows_r2", "gemv_bf16_xf32_k4_rows_r3", - "gemv_bf16_xf32_k4_rows_r4", "gemv_bf16_xf32_k4_rows_r5", "gemv_bf16_xf32_k4_rows_r6", "gemv_bf16_xf32_k4_rows_r7", - "gemv_bf16_xf32_k4_rows_r8", "gemv_bf16_xf32_k4_rows_tiled_r2", "gemv_bf16_xf32_k4_rows_tiled_r3", - "gemv_bf16_xf32_k4_rows_tiled_r4", "gemv_bf16_xf32_k4_rows_tiled_r5", "gemv_bf16_xf32_k4_rows_tiled_r6", - "gemv_bf16_xf32_k4_rows_tiled_r7", "gemv_bf16_xf32_k4_rows_tiled_r8", "gemv_bf16_xf32_x4", "gemv_bf16_xf32_x4_rows_r2", - "gemv_bf16_xf32_x4_rows_r3", - "gemv_bf16_xf32_x4_rows_r4", "gemv_bf16_xf32_x4_rows_r5", "gemv_bf16_xf32_x4_rows_r6", "gemv_bf16_xf32_x4_rows_r7", - "gemv_bf16_xf32_x4_rows_r8", "hyper_write_norm_f32", - ]); - add!("mq_rotate_x_i4", kernels::MQ_ROTATE_X_I4_SRC, ["mq_rotate_x_i4"]); - add!("qwen4_gemv_q8_0_wide", kernels::QWEN4_GEMV_Q8_0_WIDE_SRC, ["gemv_q8_0_wide_k640_staged", "gemv_q8_0_wide_rows"]); - add!("qwen4_moe_iu4_sym_gfx1151", kernels::QWEN4_MOE_IU4_SYM_GFX1151_SRC, [ - "qwen4_moe_down_iu4_sym_gfx1151", "qwen4_moe_down_iu4_sym_gfx1151_nt4", "qwen4_moe_gate_up_silu_iu4_sym_gfx1151", - "qwen4_moe_gate_up_silu_iu4_sym_gfx1151_nt4", "qwen4_moe_sym_check_gfx1151", - ]); - add!("qwen4_moe_rotate128_i4", kernels::QWEN4_MOE_ROTATE128_I4_SRC, ["qwen4_moe_rotate128_i4"]); - add!("qwen4_moe_scatter_stable_top10", kernels::QWEN4_MOE_SCATTER_STABLE_TOP10_SRC, [ - "qwen4_moe_group_prefix", "qwen4_moe_group_ranks", "qwen4_moe_group_scatter", - ]); + add!( + "gemm_wmma_lds256", + kernels::gemm_f16_x_f16_wmma_lds256_src(arch == "gfx1151"), + [ + "gemm_f16_x_f16_wmma_lds_regions_128_128_32_64_k64_p", + "gemm_wmma_lds_128_128_32_64_k64", + "gemm_wmma_lds_128_128_32_64_k64_a", + "gemm_wmma_lds_128_128_32_64_k64_bsr", + "gemm_wmma_lds_128_128_32_64_k64_gr", + "gemm_wmma_lds_128_128_32_64_k64_gra", + "gemm_wmma_lds_128_128_32_64_k64_hcsd", + "gemm_wmma_lds_128_128_32_64_k64_o16", + "gemm_wmma_lds_128_128_32_64_k64_o16g", + "gemm_wmma_lds_128_128_32_64_k64_p", + "gemm_wmma_lds_128_128_32_64_k64_p_a", + "gemm_wmma_lds_128_128_32_64_k64_p_bsr", + "gemm_wmma_lds_128_128_32_64_k64_p_gr", + "gemm_wmma_lds_128_128_32_64_k64_p_gra", + "gemm_wmma_lds_128_128_32_64_k64_p_hcsd", + "gemm_wmma_lds_128_128_32_64_k64_p_o16", + "gemm_wmma_lds_128_128_32_64_k64_p_o16g", + "gemm_wmma_lds_128_128_64_64_k64", + "gemm_wmma_lds_128_128_64_64_k64_a", + "gemm_wmma_lds_128_128_64_64_k64_bsr", + "gemm_wmma_lds_128_128_64_64_k64_gr", + "gemm_wmma_lds_128_128_64_64_k64_gra", + "gemm_wmma_lds_128_128_64_64_k64_o16", + "gemm_wmma_lds_128_128_64_64_k64_o16g", + "gemm_wmma_lds_128_256_32_64_k64", + "gemm_wmma_lds_128_256_32_64_k64_a", + "gemm_wmma_lds_128_256_32_64_k64_bf16st", + "gemm_wmma_lds_128_256_32_64_k64_bsr", + "gemm_wmma_lds_128_256_32_64_k64_gr", + "gemm_wmma_lds_128_256_32_64_k64_gra", + "gemm_wmma_lds_128_256_32_64_k64_hcsd", + "gemm_wmma_lds_128_256_32_64_k64_o16", + "gemm_wmma_lds_128_256_32_64_k64_o16g", + "gemm_wmma_lds_128_256_32_64_k64_p", + "gemm_wmma_lds_128_256_32_64_k64_p_a", + "gemm_wmma_lds_128_256_32_64_k64_p_bsr", + "gemm_wmma_lds_128_256_32_64_k64_p_gr", + "gemm_wmma_lds_128_256_32_64_k64_p_gra", + "gemm_wmma_lds_128_256_32_64_k64_p_hcsd", + "gemm_wmma_lds_128_256_32_64_k64_p_o16", + "gemm_wmma_lds_128_256_32_64_k64_p_o16g", + "gemm_wmma_lds_128_256_64_64_k32", + "gemm_wmma_lds_128_256_64_64_k32_a", + "gemm_wmma_lds_128_256_64_64_k32_bsr", + "gemm_wmma_lds_128_256_64_64_k32_gr", + "gemm_wmma_lds_128_256_64_64_k32_gra", + "gemm_wmma_lds_128_256_64_64_k32_o16", + "gemm_wmma_lds_128_256_64_64_k32_o16g", + "gemm_wmma_lds_128_256_64_64_k32_p", + "gemm_wmma_lds_128_256_64_64_k32_p_a", + "gemm_wmma_lds_128_256_64_64_k32_p_bsr", + "gemm_wmma_lds_128_256_64_64_k32_p_gr", + "gemm_wmma_lds_128_256_64_64_k32_p_gra", + "gemm_wmma_lds_128_256_64_64_k32_p_o16", + "gemm_wmma_lds_128_256_64_64_k32_p_o16g", + "gemm_wmma_lds_128_256_64_64_k64", + "gemm_wmma_lds_128_256_64_64_k64_sw", + "gemm_wmma_lds_128_512_32_64_k32", + "gemm_wmma_lds_128_512_64_64_k32", + "gemm_wmma_lds_128_512_64_64_k32_sw", + "gemm_wmma_lds_160_64_32_64_k64", + "gemm_wmma_lds_160_64_32_64_k64_p", + "gemm_wmma_lds_256_128_32_64_k64", + "gemm_wmma_lds_256_128_64_64_k64", + "gemm_wmma_lds_256_128_64_64_k64_sw", + "gemm_wmma_lds_256_256_32_64_k64", + "gemm_wmma_lds_256_256_64_64_k64", + "gemm_wmma_lds_256_256_64_64_k64_a", + "gemm_wmma_lds_256_256_64_64_k64_bsr", + "gemm_wmma_lds_256_256_64_64_k64_gr", + "gemm_wmma_lds_256_256_64_64_k64_gra", + "gemm_wmma_lds_256_256_64_64_k64_hcsd", + "gemm_wmma_lds_256_256_64_64_k64_o16", + "gemm_wmma_lds_256_256_64_64_k64_o16g", + "gemm_wmma_lds_256_256_64_64_k64_p", + "gemm_wmma_lds_256_256_64_64_k64_p_a", + "gemm_wmma_lds_256_256_64_64_k64_p_bsr", + "gemm_wmma_lds_256_256_64_64_k64_p_gr", + "gemm_wmma_lds_256_256_64_64_k64_p_gra", + "gemm_wmma_lds_256_256_64_64_k64_p_hcsd", + "gemm_wmma_lds_256_256_64_64_k64_p_o16", + "gemm_wmma_lds_256_256_64_64_k64_p_o16g", + "gemm_wmma_lds_64_512_64_64_k32", + "gemm_wmma_lds_64_64_32_64_k64", + "gemm_wmma_lds_64_64_32_64_k64_p", + ] + ); + add!( + "gemv_bf16_xf32", + kernels::gemv_bf16_xf32_src(arch == "gfx1151"), + [ + "gemv_bf16_xf32", + "gemv_bf16_xf32_bf16_scaled_add", + "gemv_bf16_xf32_k4", + "gemv_bf16_xf32_k4_rows_r2", + "gemv_bf16_xf32_k4_rows_r3", + "gemv_bf16_xf32_k4_rows_r4", + "gemv_bf16_xf32_k4_rows_r5", + "gemv_bf16_xf32_k4_rows_r6", + "gemv_bf16_xf32_k4_rows_r7", + "gemv_bf16_xf32_k4_rows_r8", + "gemv_bf16_xf32_k4_rows_tiled_r2", + "gemv_bf16_xf32_k4_rows_tiled_r3", + "gemv_bf16_xf32_k4_rows_tiled_r4", + "gemv_bf16_xf32_k4_rows_tiled_r5", + "gemv_bf16_xf32_k4_rows_tiled_r6", + "gemv_bf16_xf32_k4_rows_tiled_r7", + "gemv_bf16_xf32_k4_rows_tiled_r8", + "gemv_bf16_xf32_x4", + "gemv_bf16_xf32_x4_rows_r2", + "gemv_bf16_xf32_x4_rows_r3", + "gemv_bf16_xf32_x4_rows_r4", + "gemv_bf16_xf32_x4_rows_r5", + "gemv_bf16_xf32_x4_rows_r6", + "gemv_bf16_xf32_x4_rows_r7", + "gemv_bf16_xf32_x4_rows_r8", + "hyper_write_norm_f32", + ] + ); + add!( + "mq_rotate_x_i4", + kernels::MQ_ROTATE_X_I4_SRC, + ["mq_rotate_x_i4"] + ); + add!( + "qwen4_gemv_q8_0_wide", + kernels::QWEN4_GEMV_Q8_0_WIDE_SRC, + ["gemv_q8_0_wide_k640_staged", "gemv_q8_0_wide_rows"] + ); + add!( + "qwen4_moe_iu4_sym_gfx1151", + kernels::QWEN4_MOE_IU4_SYM_GFX1151_SRC, + [ + "qwen4_moe_down_iu4_sym_gfx1151", + "qwen4_moe_down_iu4_sym_gfx1151_nt4", + "qwen4_moe_gate_up_silu_iu4_sym_gfx1151", + "qwen4_moe_gate_up_silu_iu4_sym_gfx1151_nt4", + "qwen4_moe_sym_check_gfx1151", + ] + ); + add!( + "qwen4_moe_rotate128_i4", + kernels::QWEN4_MOE_ROTATE128_I4_SRC, + ["qwen4_moe_rotate128_i4"] + ); + add!( + "qwen4_moe_scatter_stable_top10", + kernels::QWEN4_MOE_SCATTER_STABLE_TOP10_SRC, + [ + "qwen4_moe_group_prefix", + "qwen4_moe_group_ranks", + "qwen4_moe_group_scatter", + ] + ); } // Qwen3.5-MoE (ornith-1.5-35b-a3b) load, AR prefill/decode and native MTP, // plus Qwen3.5 dense (H2): every module a kernel-load trace @@ -813,130 +2632,383 @@ pub fn entries(arch: &str, extra_flags: &str) -> Result, Regist // expressions the callsites pass to `ensure_kernel`; symbols are every // kernel the arch's object defines. if matches!(arch, "gfx1201" | "gfx1100" | "gfx1151") { - add!("gated_delta_net_q8_compact2_b2", kernels::GATED_DELTA_NET_Q8_COMPACT2_B2_SRC, [ - "gated_delta_net_q8_compact2_b2", "gated_delta_net_q8_fast_independent_masked", - ]); - add!("gemm_qkv_hfq6g256", kernels::GEMM_QKV_HFQ6G256_SRC, ["gemm_qkv_hfq6g256"]); - add!("gemm_qkvza_hfq6g256", kernels::GEMM_QKVZA_HFQ6G256_SRC, ["gemm_qkvza_hfq6g256"]); - add!("gemv_hfq6g256_residual", kernels::GEMV_HFQ6G256_RESIDUAL_SRC, ["gemv_hfq6g256_residual"]); - add!("gemv_hfq6g256_residual_sigmoid_scaled", kernels::GEMV_HFQ6G256_RESIDUAL_SIGMOID_SCALED_SRC, [ - "gemv_hfq6g256_residual_sigmoid_scaled_gpu_batched", - ]); - add!("gemv_mq4g256v2_moe_down_k8_indexed_batched_expanded", kernels::GEMV_MQ4G256V2_MOE_DOWN_K8_INDEXED_BATCHED_EXPANDED_SRC, [ + add!( + "gated_delta_net_q8_compact2_b2", + kernels::GATED_DELTA_NET_Q8_COMPACT2_B2_SRC, + [ + "gated_delta_net_q8_compact2_b2", + "gated_delta_net_q8_fast_independent_masked", + ] + ); + add!( + "gemm_qkv_hfq6g256", + kernels::GEMM_QKV_HFQ6G256_SRC, + ["gemm_qkv_hfq6g256"] + ); + add!( + "gemm_qkvza_hfq6g256", + kernels::GEMM_QKVZA_HFQ6G256_SRC, + ["gemm_qkvza_hfq6g256"] + ); + add!( + "gemv_hfq6g256_residual", + kernels::GEMV_HFQ6G256_RESIDUAL_SRC, + ["gemv_hfq6g256_residual"] + ); + add!( + "gemv_hfq6g256_residual_sigmoid_scaled", + kernels::GEMV_HFQ6G256_RESIDUAL_SIGMOID_SCALED_SRC, + ["gemv_hfq6g256_residual_sigmoid_scaled_gpu_batched",] + ); + add!( "gemv_mq4g256v2_moe_down_k8_indexed_batched_expanded", - ]); - add!("moe_down_combine_grouped_k8", kernels::MOE_DOWN_COMBINE_GROUPED_K8_SRC, ["moe_down_combine_grouped_k8"]); - add!("moe_down_combine_k8_batched", kernels::MOE_DOWN_COMBINE_K8_BATCHED_SRC, ["moe_down_combine_k8_batched"]); - add!("moe_gate_up_unscatter_k8", kernels::MOE_GATE_UP_UNSCATTER_K8_SRC, ["moe_gate_up_unscatter_k8"]); - add!("moe_scatter_fused_k8", kernels::moe_scatter_fused_k8_src(arch == "gfx1151"), ["moe_scatter_fused_k8"]); - add!("moe_topk_renorm_k8", kernels::MOE_TOPK_RENORM_K8_SRC, ["moe_topk_renorm_k8"]); - add!("moe_topk_renorm_k8_batched", kernels::MOE_TOPK_RENORM_K8_BATCHED_SRC, ["moe_topk_renorm_k8_batched"]); - add!("scaled_add_inplace", kernels::SCALED_ADD_INPLACE_SRC, ["scaled_add_inplace_cpu_scalar_f32", "scaled_add_inplace_gpu_scalar_f32"]); - add!("sigmoid_scaled_residual_add_batched", kernels::SIGMOID_SCALED_RESIDUAL_ADD_BATCHED_SRC, ["sigmoid_scaled_residual_add_batched_f32"]); + kernels::GEMV_MQ4G256V2_MOE_DOWN_K8_INDEXED_BATCHED_EXPANDED_SRC, + ["gemv_mq4g256v2_moe_down_k8_indexed_batched_expanded",] + ); + add!( + "moe_down_combine_grouped_k8", + kernels::MOE_DOWN_COMBINE_GROUPED_K8_SRC, + ["moe_down_combine_grouped_k8"] + ); + add!( + "moe_down_combine_k8_batched", + kernels::MOE_DOWN_COMBINE_K8_BATCHED_SRC, + ["moe_down_combine_k8_batched"] + ); + add!( + "moe_gate_up_unscatter_k8", + kernels::MOE_GATE_UP_UNSCATTER_K8_SRC, + ["moe_gate_up_unscatter_k8"] + ); + add!( + "moe_scatter_fused_k8", + kernels::moe_scatter_fused_k8_src(arch == "gfx1151"), + ["moe_scatter_fused_k8"] + ); + add!( + "moe_topk_renorm_k8", + kernels::MOE_TOPK_RENORM_K8_SRC, + ["moe_topk_renorm_k8"] + ); + add!( + "moe_topk_renorm_k8_batched", + kernels::MOE_TOPK_RENORM_K8_BATCHED_SRC, + ["moe_topk_renorm_k8_batched"] + ); + add!( + "scaled_add_inplace", + kernels::SCALED_ADD_INPLACE_SRC, + [ + "scaled_add_inplace_cpu_scalar_f32", + "scaled_add_inplace_gpu_scalar_f32" + ] + ); + add!( + "sigmoid_scaled_residual_add_batched", + kernels::SIGMOID_SCALED_RESIDUAL_ADD_BATCHED_SRC, + ["sigmoid_scaled_residual_add_batched_f32"] + ); add!("softmax", kernels::SOFTMAX_SRC, ["softmax_f32"]); } if matches!(arch, "gfx1201" | "gfx1100") { - add!("gemv_mq4g256v2_residual_sigmoid_scaled_k512", kernels::GEMV_MQ4G256V2_RESIDUAL_SIGMOID_SCALED_K512_SRC, [ + add!( "gemv_mq4g256v2_residual_sigmoid_scaled_k512", - ]); - add!("repeat_interleave_qk_batched", kernels::REPEAT_INTERLEAVE_QK_BATCHED_SRC, ["repeat_interleave_qk_f32_batched"]); + kernels::GEMV_MQ4G256V2_RESIDUAL_SIGMOID_SCALED_K512_SRC, + ["gemv_mq4g256v2_residual_sigmoid_scaled_k512",] + ); + add!( + "repeat_interleave_qk_batched", + kernels::REPEAT_INTERLEAVE_QK_BATCHED_SRC, + ["repeat_interleave_qk_f32_batched"] + ); } if matches!(arch, "gfx1201" | "gfx1151") { - add!("gemv_mq4g256v2_moe_gate_up_k8_indexed", kernels::GEMV_MQ4G256V2_MOE_GATE_UP_K8_INDEXED_SRC, ["gemv_mq4g256v2_moe_gate_up_k8_indexed"]); - add!("gemv_mq4g256v2_moe_ninepath_d4", kernels::GEMV_MQ4G256V2_MOE_NINEPATH_D4_SRC, ["gemv_mq4g256v2_moe_ninepath_d4"]); - add!("kv_cache_write_q8_0_pair", kernels::KV_CACHE_WRITE_Q8_0_PAIR_GFX1100_SRC, ["kv_cache_write_q8_0_pair"]); + add!( + "gemv_mq4g256v2_moe_gate_up_k8_indexed", + kernels::GEMV_MQ4G256V2_MOE_GATE_UP_K8_INDEXED_SRC, + ["gemv_mq4g256v2_moe_gate_up_k8_indexed"] + ); + add!( + "gemv_mq4g256v2_moe_ninepath_d4", + kernels::GEMV_MQ4G256V2_MOE_NINEPATH_D4_SRC, + ["gemv_mq4g256v2_moe_ninepath_d4"] + ); + add!( + "kv_cache_write_q8_0_pair", + kernels::KV_CACHE_WRITE_Q8_0_PAIR_GFX1100_SRC, + ["kv_cache_write_q8_0_pair"] + ); } if matches!(arch, "gfx1100" | "gfx1151") { - add!("fused_rmsnorm_mq_rotate_vecsum", kernels::FUSED_RMSNORM_MQ_ROTATE_VECSUM_GFX1100_SRC, ["fused_rmsnorm_mq_rotate_vecsum"]); - add!("gemm_gate_up_hfq6g256_wmma", kernels::GEMM_GATE_UP_HFQ6G256_WMMA_SRC, ["gemm_gate_up_hfq6g256_wmma"]); - add!("gemm_hfq6g256_residual_wmma_k2", kernels::GEMM_HFQ6G256_RESIDUAL_WMMA_K2_SRC, ["gemm_hfq6g256_residual_wmma_k2"]); - add!("gemm_q8_0_wmma", kernels::GEMM_Q8_0_WMMA_SRC, ["gemm_q8_0_wmma"]); - add!("gemm_qkv_hfq6g256_wmma", kernels::GEMM_QKV_HFQ6G256_WMMA_SRC, ["gemm_qkv_hfq6g256_wmma"]); - add!("gemm_qkvza_hfq6g256_wmma", kernels::GEMM_QKVZA_HFQ6G256_WMMA_SRC, ["gemm_qkvza_hfq6g256_wmma"]); - add!("moe_router_softmax_topk_k8_wave64_exact", kernels::MOE_ROUTER_SOFTMAX_TOPK_K8_WAVE64_EXACT_SRC, ["moe_router_softmax_topk_k8_wave64_exact"]); + add!( + "fused_rmsnorm_mq_rotate_vecsum", + kernels::FUSED_RMSNORM_MQ_ROTATE_VECSUM_GFX1100_SRC, + ["fused_rmsnorm_mq_rotate_vecsum"] + ); + add!( + "gemm_gate_up_hfq6g256_wmma", + kernels::GEMM_GATE_UP_HFQ6G256_WMMA_SRC, + ["gemm_gate_up_hfq6g256_wmma"] + ); + add!( + "gemm_hfq6g256_residual_wmma_k2", + kernels::GEMM_HFQ6G256_RESIDUAL_WMMA_K2_SRC, + ["gemm_hfq6g256_residual_wmma_k2"] + ); + add!( + "gemm_q8_0_wmma", + kernels::GEMM_Q8_0_WMMA_SRC, + ["gemm_q8_0_wmma"] + ); + add!( + "gemm_qkv_hfq6g256_wmma", + kernels::GEMM_QKV_HFQ6G256_WMMA_SRC, + ["gemm_qkv_hfq6g256_wmma"] + ); + add!( + "gemm_qkvza_hfq6g256_wmma", + kernels::GEMM_QKVZA_HFQ6G256_WMMA_SRC, + ["gemm_qkvza_hfq6g256_wmma"] + ); + add!( + "moe_router_softmax_topk_k8_wave64_exact", + kernels::MOE_ROUTER_SOFTMAX_TOPK_K8_WAVE64_EXACT_SRC, + ["moe_router_softmax_topk_k8_wave64_exact"] + ); } if arch == "gfx1201" { - add!("attention_flash_q8_0_reduce_gated_mq_rotate_gfx1201", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1201_SRC, [ + add!( "attention_flash_q8_0_reduce_gated_mq_rotate_gfx1201", - ]); - add!("attention_q8_0_flash_prefill_wmma_gfx12_hd256", format!("#define SPLIT_Q 0\n#define FIXED_HEAD_DIM 256\n#define PREFETCH_V 1\n{}", kernels::kv_slot_desc_source(kernels::ATTENTION_Q8_0_FLASH_PREFILL_WMMA_GFX12_SRC, false)), [ - "attention_q8_0_flash_prefill_wmma", - ]); - add!("gated_norm_mq_rotate_gfx1201", kernels::GATED_NORM_MQ_ROTATE_GFX1201_SRC, ["gated_norm_mq_rotate_gfx1201"]); - add!("gated_norm_mq_rotate_i4_gfx12_v2", kernels::GATED_NORM_MQ_ROTATE_I4_GFX12_V2_SRC, ["gated_norm_mq_rotate_i4_gfx12_v2"]); - add!("gemm_gate_up_hfq6g256_wmma_gfx12", kernels::GEMM_GATE_UP_HFQ6G256_WMMA_GFX12_SRC, ["gemm_gate_up_hfq6g256_wmma_gfx12"]); - add!("gemm_hfq6g256_residual_wmma_gfx12", kernels::GEMM_HFQ6G256_RESIDUAL_WMMA_GFX12_SRC, ["gemm_hfq6g256_residual_wmma_gfx12"]); - add!("gemm_mq4g256v2_residual_mmq_iu4_gfx12_g12r", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GFX12_G12R_SRC, [ - "gemm_mq4g256v2_gate_up_silu_mmq_iu4_g12r", "gemm_mq4g256v2_residual_mmq_iu4_full_add_g12r", - "gemm_mq4g256v2_residual_mmq_iu4_full_set_g12r", "gemm_mq4g256v2_residual_mmq_iu4_g12r", "quantize_int4_mmq_ds128", - ]); - add!("gemm_q8_0_wmma_gfx12", kernels::GEMM_Q8_0_WMMA_GFX12_SRC, ["gemm_q8_0_wmma_gfx12"]); - add!("gemm_qkv_hfq6g256_wmma_gfx12", kernels::GEMM_QKV_HFQ6G256_WMMA_GFX12_SRC, ["gemm_qkv_hfq6g256_wmma_gfx12"]); - add!("gemm_qkvza_hfq6g256_wmma_gfx12", kernels::GEMM_QKVZA_HFQ6G256_WMMA_GFX12_SRC, ["gemm_qkvza_hfq6g256_wmma_gfx12"]); - add!("moe_router_softmax_topk_k8_wave64", kernels::MOE_ROUTER_SOFTMAX_TOPK_K8_WAVE64_SRC, ["moe_router_softmax_topk_k8_wave64"]); - add!("qwen35_fa_prep_gfx1201", kernels::QWEN35_FA_PREP_GFX1201_SRC, ["qwen35_fa_prep_gfx1201"]); + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1201_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1201",] + ); + add!( + "attention_q8_0_flash_prefill_wmma_gfx12_hd256", + format!( + "#define SPLIT_Q 0\n#define FIXED_HEAD_DIM 256\n#define PREFETCH_V 1\n{}", + kernels::kv_slot_desc_source( + kernels::ATTENTION_Q8_0_FLASH_PREFILL_WMMA_GFX12_SRC, + false + ) + ), + ["attention_q8_0_flash_prefill_wmma",] + ); + add!( + "gated_norm_mq_rotate_gfx1201", + kernels::GATED_NORM_MQ_ROTATE_GFX1201_SRC, + ["gated_norm_mq_rotate_gfx1201"] + ); + add!( + "gated_norm_mq_rotate_i4_gfx12_v2", + kernels::GATED_NORM_MQ_ROTATE_I4_GFX12_V2_SRC, + ["gated_norm_mq_rotate_i4_gfx12_v2"] + ); + add!( + "gemm_gate_up_hfq6g256_wmma_gfx12", + kernels::GEMM_GATE_UP_HFQ6G256_WMMA_GFX12_SRC, + ["gemm_gate_up_hfq6g256_wmma_gfx12"] + ); + add!( + "gemm_hfq6g256_residual_wmma_gfx12", + kernels::GEMM_HFQ6G256_RESIDUAL_WMMA_GFX12_SRC, + ["gemm_hfq6g256_residual_wmma_gfx12"] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4_gfx12_g12r", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GFX12_G12R_SRC, + [ + "gemm_mq4g256v2_gate_up_silu_mmq_iu4_g12r", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_g12r", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_g12r", + "gemm_mq4g256v2_residual_mmq_iu4_g12r", + "quantize_int4_mmq_ds128", + ] + ); + add!( + "gemm_q8_0_wmma_gfx12", + kernels::GEMM_Q8_0_WMMA_GFX12_SRC, + ["gemm_q8_0_wmma_gfx12"] + ); + add!( + "gemm_qkv_hfq6g256_wmma_gfx12", + kernels::GEMM_QKV_HFQ6G256_WMMA_GFX12_SRC, + ["gemm_qkv_hfq6g256_wmma_gfx12"] + ); + add!( + "gemm_qkvza_hfq6g256_wmma_gfx12", + kernels::GEMM_QKVZA_HFQ6G256_WMMA_GFX12_SRC, + ["gemm_qkvza_hfq6g256_wmma_gfx12"] + ); + add!( + "moe_router_softmax_topk_k8_wave64", + kernels::MOE_ROUTER_SOFTMAX_TOPK_K8_WAVE64_SRC, + ["moe_router_softmax_topk_k8_wave64"] + ); + add!( + "qwen35_fa_prep_gfx1201", + kernels::QWEN35_FA_PREP_GFX1201_SRC, + ["qwen35_fa_prep_gfx1201"] + ); } if arch == "gfx1100" { - add!("argmax_token_chain", kernels::ARGMAX_TOKEN_CHAIN_SRC, ["argmax_token_chain_f32"]); - add!("attention_flash_q8_0_reduce_gated_mq_rotate_awq_gfx1100", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_AWQ_GFX1100_SRC, [ + add!( + "argmax_token_chain", + kernels::ARGMAX_TOKEN_CHAIN_SRC, + ["argmax_token_chain_f32"] + ); + add!( "attention_flash_q8_0_reduce_gated_mq_rotate_awq_gfx1100", - ]); - add!("attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1100_SRC, [ + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_AWQ_GFX1100_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_awq_gfx1100",] + ); + add!( "attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100", - ]); - add!("conv1d_silu_split_qknorm_b256", kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SRC, ["conv1d_silu_split_qknorm_b256"]); - add!("deinterleave_q_rmsnorm_f32_batched", kernels::DEINTERLEAVE_Q_RMSNORM_BATCHED_SRC, ["deinterleave_q_rmsnorm_f32_batched"]); - add!("fused_qkv_mq4g256v2_k2048_x_buffer_gfx1100", kernels::FUSED_QKV_MQ4G256V2_K2048_X_BUFFER_GFX1100_SRC, [ + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1100_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100",] + ); + add!( + "conv1d_silu_split_qknorm_b256", + kernels::CONV1D_SILU_SPLIT_QKNORM_B256_SRC, + ["conv1d_silu_split_qknorm_b256"] + ); + add!( + "deinterleave_q_rmsnorm_f32_batched", + kernels::DEINTERLEAVE_Q_RMSNORM_BATCHED_SRC, + ["deinterleave_q_rmsnorm_f32_batched"] + ); + add!( "fused_qkv_mq4g256v2_k2048_x_buffer_gfx1100", - ]); - add!("fused_qkvza_mq4g256v2_k2048_hoist_x32_gfx1100", kernels::FUSED_QKVZA_MQ4G256V2_K2048_HOIST_X32_GFX1100_SRC, [ + kernels::FUSED_QKV_MQ4G256V2_K2048_X_BUFFER_GFX1100_SRC, + ["fused_qkv_mq4g256v2_k2048_x_buffer_gfx1100",] + ); + add!( "fused_qkvza_mq4g256v2_k2048_hoist_x32_gfx1100", - ]); - add!("gated_norm_mq_rotate_gfx1100", kernels::GATED_NORM_MQ_ROTATE_GFX1100_SRC, ["gated_norm_mq_rotate_gfx1100"]); - add!("gated_norm_mq_rotate_i4_gfx1100_v2", kernels::GATED_NORM_MQ_ROTATE_I4_GFX1100_V2_SRC, ["gated_norm_mq_rotate_i4_gfx1100_v2"]); - add!("gemm_mq4g256v2_moe_grouped_wmma_k2", kernels::GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, ["gemm_mq4g256v2_moe_grouped_wmma_k2"]); - add!("gemm_mq4g256v2_residual_mmq_iu4_gridspec", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SRC, [ - "gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", "gemm_mq4g256v2_residual_mmq_iu4_full_add", - "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", - "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", - "gemm_mq4g256v2_residual_mmq_iu4_full_set", "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", - "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", - "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", - "quantize_int4_mmq_ds128", - ]); - add!("gemv_mq4g256v2_moe_gate_up_k8_indexed_k2048_nolds_gfx1100", kernels::GEMV_MQ4G256V2_MOE_GATE_UP_K8_INDEXED_K2048_NOLDS_GFX1100_SRC, [ + kernels::FUSED_QKVZA_MQ4G256V2_K2048_HOIST_X32_GFX1100_SRC, + ["fused_qkvza_mq4g256v2_k2048_hoist_x32_gfx1100",] + ); + add!( + "gated_norm_mq_rotate_gfx1100", + kernels::GATED_NORM_MQ_ROTATE_GFX1100_SRC, + ["gated_norm_mq_rotate_gfx1100"] + ); + add!( + "gated_norm_mq_rotate_i4_gfx1100_v2", + kernels::GATED_NORM_MQ_ROTATE_I4_GFX1100_V2_SRC, + ["gated_norm_mq_rotate_i4_gfx1100_v2"] + ); + add!( + "gemm_mq4g256v2_moe_grouped_wmma_k2", + kernels::GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, + ["gemm_mq4g256v2_moe_grouped_wmma_k2"] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4_gridspec", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_gfx1100", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", + "quantize_int4_mmq_ds128", + ] + ); + add!( "gemv_mq4g256v2_moe_gate_up_k8_indexed_k2048_nolds_gfx1100", - ]); - add!("gemv_mq4g256v2_moe_ninepath_rpb8_gfx1100", kernels::GEMV_MQ4G256V2_MOE_NINEPATH_RPB8_GFX1100_SRC, ["gemv_mq4g256v2_moe_ninepath_rpb8_gfx1100"]); - add!("gemv_mq4g256v2_residual_r1_k4096_gfx1100_noscratch", kernels::GEMV_MQ4G256V2_RESIDUAL_R1_K4096_GFX1100_NOSCRATCH_SRC, [ + kernels::GEMV_MQ4G256V2_MOE_GATE_UP_K8_INDEXED_K2048_NOLDS_GFX1100_SRC, + ["gemv_mq4g256v2_moe_gate_up_k8_indexed_k2048_nolds_gfx1100",] + ); + add!( + "gemv_mq4g256v2_moe_ninepath_rpb8_gfx1100", + kernels::GEMV_MQ4G256V2_MOE_NINEPATH_RPB8_GFX1100_SRC, + ["gemv_mq4g256v2_moe_ninepath_rpb8_gfx1100"] + ); + add!( "gemv_mq4g256v2_residual_r1_k4096_gfx1100_noscratch", - ]); - add!("greedy_accept", kernels::GREEDY_ACCEPT_SRC, ["greedy_accept_from_argmax_i32"]); - add!("mq_rotate_x_i4", kernels::MQ_ROTATE_X_I4_SRC, ["mq_rotate_x_i4"]); - add!("qwen35_fa_prep_gfx1100", kernels::QWEN35_FA_PREP_GFX1100_SRC, ["qwen35_fa_prep_gfx1100"]); + kernels::GEMV_MQ4G256V2_RESIDUAL_R1_K4096_GFX1100_NOSCRATCH_SRC, + ["gemv_mq4g256v2_residual_r1_k4096_gfx1100_noscratch",] + ); + add!( + "greedy_accept", + kernels::GREEDY_ACCEPT_SRC, + ["greedy_accept_from_argmax_i32"] + ); + add!( + "mq_rotate_x_i4", + kernels::MQ_ROTATE_X_I4_SRC, + ["mq_rotate_x_i4"] + ); + add!( + "qwen35_fa_prep_gfx1100", + kernels::QWEN35_FA_PREP_GFX1100_SRC, + ["qwen35_fa_prep_gfx1100"] + ); } if arch == "gfx1151" { - add!("attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1151_SRC, [ + add!( "attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151", - ]); - add!("attention_verify_wmma_gfx1151", kernels::ATTENTION_VERIFY_WMMA_GFX1151_SRC, [ - "attention_verify_wmma_pv1s_d4_gfx1151", "attention_verify_wmma_pv2_d2_gfx1151", "attention_verify_wmma_qk_gfx1151", - ]); - add!("gated_norm_mq_rotate_gfx1151", kernels::GATED_NORM_MQ_ROTATE_GFX1151_SRC, ["gated_norm_mq_rotate_gfx1151"]); - add!("gated_norm_mq_rotate_i4_gfx1151_v2", kernels::GATED_NORM_MQ_ROTATE_I4_GFX1151_V2_SRC, ["gated_norm_mq_rotate_i4_gfx1151_v2"]); - add!("gemm_mq4g256v2_moe_grouped_wmma_k2", kernels::QWEN4_GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, [ - "gemm_mq4g256v2_moe_grouped_wmma_k2", "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", - "gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", - ]); - add!("gemm_mq4g256v2_residual_mmq_iu4_gridspec", kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SRC, [ - "gemm_mq4g256v2_residual_mmq_iu4", "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", "gemm_mq4g256v2_residual_mmq_iu4_full_add", - "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", - "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set", - "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", - "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", - "quantize_int4_mmq_ds128", - ]); - add!("qwen35_fa_prep_gfx1151", kernels::QWEN35_FA_PREP_GFX1151_SRC, ["qwen35_fa_prep_gfx1151"]); + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1151_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151",] + ); + add!( + "attention_verify_wmma_gfx1151", + kernels::ATTENTION_VERIFY_WMMA_GFX1151_SRC, + [ + "attention_verify_wmma_pv1s_d4_gfx1151", + "attention_verify_wmma_pv2_d2_gfx1151", + "attention_verify_wmma_qk_gfx1151", + ] + ); + add!( + "gated_norm_mq_rotate_gfx1151", + kernels::GATED_NORM_MQ_ROTATE_GFX1151_SRC, + ["gated_norm_mq_rotate_gfx1151"] + ); + add!( + "gated_norm_mq_rotate_i4_gfx1151_v2", + kernels::GATED_NORM_MQ_ROTATE_I4_GFX1151_V2_SRC, + ["gated_norm_mq_rotate_i4_gfx1151_v2"] + ); + add!( + "gemm_mq4g256v2_moe_grouped_wmma_k2", + kernels::QWEN4_GEMM_MQ4G256V2_MOE_GROUPED_WMMA_K2_SRC, + [ + "gemm_mq4g256v2_moe_grouped_wmma_k2", + "gemm_mq4g256v2_moe_grouped_wmma_k2_bf16out", + "gemm_mq4g256v2_moe_grouped_wmma_k2_silu_bf16out", + ] + ); + add!( + "gemm_mq4g256v2_residual_mmq_iu4_gridspec", + kernels::GEMM_MQ4G256V2_RESIDUAL_MMQ_IU4_GRIDSPEC_SRC, + [ + "gemm_mq4g256v2_residual_mmq_iu4", + "gemm_mq4g256v2_residual_mmq_iu4_branch_gridspec", + "gemm_mq4g256v2_residual_mmq_iu4_full_add", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_add_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_lf16_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3", + "gemm_mq4g256v2_residual_mmq_iu4_full_set_occ3_col_gfx1151", + "gemm_mq4g256v2_residual_mmq_iu4_tail_gridspec", + "quantize_int4_mmq_ds128", + ] + ); + add!( + "qwen35_fa_prep_gfx1151", + kernels::QWEN35_FA_PREP_GFX1151_SRC, + ["qwen35_fa_prep_gfx1151"] + ); } Ok(entries) } @@ -966,7 +3038,10 @@ fn entry( /// the programs' captures record; each source is the expression the cited /// launcher passes to `ensure_kernel` on the default route. railgun's JIT /// receipt corpus compiles [`corpus_entries`], never a hand-picked source. -pub fn default_route_entries(arch: &str, extra_flags: &str) -> Result, RegistryError> { +pub fn default_route_entries( + arch: &str, + extra_flags: &str, +) -> Result, RegistryError> { let arch: &'static str = SUPPORTED_ARCHES .iter() .copied() @@ -981,84 +3056,331 @@ pub fn default_route_entries(arch: &str, extra_flags: &str) -> Result { - add!("attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1100_SRC, ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100"]); // attention.rs:10626 - add!("gated_norm_mq_rotate_gfx1100", kernels::GATED_NORM_MQ_ROTATE_GFX1100_SRC, ["gated_norm_mq_rotate_gfx1100"]); // gemv.rs:4796 - add!("gemv_hfq4g256_moe_gate_up_indexed_cpol_slc", kernels::GEMV_HFQ4G256_MOE_GATE_UP_INDEXED_CPOL_SLC_GFX1100_SRC, ["gemv_hfq4g256_moe_gate_up_k8_indexed_cpol_slc"]); // gemv.rs:13770 - add!("gemv_hfq4g256_residual_sigmoid_buffer_gfx1100", kernels::GEMV_HFQ4G256_RESIDUAL_SIGMOID_BUFFER_GFX1100_SRC, ["gemv_hfq4g256_residual_sigmoid_scaled_gpu"]); // gemv.rs:12369 - add!("gemv_hfq4g256_residual_stage_x32_gfx1100", kernels::GEMV_HFQ4G256_RESIDUAL_STAGE_X32_GFX1100_SRC, ["gemv_hfq4g256_residual"]); // gemv.rs:10262 - add!("moe_down_combine_rmsnorm_mq_rotate_vecsum_gfx1100", kernels::MOE_DOWN_COMBINE_RMSNORM_MQ_ROTATE_VECSUM_GFX1100_SRC, ["moe_down_combine_rmsnorm_mq_rotate_vecsum"]); // moe.rs:119 - add!("qwen35_fa_prep_gfx1100", kernels::QWEN35_FA_PREP_GFX1100_SRC, ["qwen35_fa_prep_gfx1100"]); // norm.rs:1152 + add!( + "attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100", + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1100_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1100"] + ); // attention.rs:10626 + add!( + "gated_norm_mq_rotate_gfx1100", + kernels::GATED_NORM_MQ_ROTATE_GFX1100_SRC, + ["gated_norm_mq_rotate_gfx1100"] + ); // gemv.rs:4796 + add!( + "gemv_hfq4g256_moe_gate_up_indexed_cpol_slc", + kernels::GEMV_HFQ4G256_MOE_GATE_UP_INDEXED_CPOL_SLC_GFX1100_SRC, + ["gemv_hfq4g256_moe_gate_up_k8_indexed_cpol_slc"] + ); // gemv.rs:13770 + add!( + "gemv_hfq4g256_residual_sigmoid_buffer_gfx1100", + kernels::GEMV_HFQ4G256_RESIDUAL_SIGMOID_BUFFER_GFX1100_SRC, + ["gemv_hfq4g256_residual_sigmoid_scaled_gpu"] + ); // gemv.rs:12369 + add!( + "gemv_hfq4g256_residual_stage_x32_gfx1100", + kernels::GEMV_HFQ4G256_RESIDUAL_STAGE_X32_GFX1100_SRC, + ["gemv_hfq4g256_residual"] + ); // gemv.rs:10262 + add!( + "moe_down_combine_rmsnorm_mq_rotate_vecsum_gfx1100", + kernels::MOE_DOWN_COMBINE_RMSNORM_MQ_ROTATE_VECSUM_GFX1100_SRC, + ["moe_down_combine_rmsnorm_mq_rotate_vecsum"] + ); // moe.rs:119 + add!( + "qwen35_fa_prep_gfx1100", + kernels::QWEN35_FA_PREP_GFX1100_SRC, + ["qwen35_fa_prep_gfx1100"] + ); // norm.rs:1152 } "gfx1151" => { - add!("attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151", kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1151_SRC, ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151"]); // attention.rs:10620 - add!("fused_qkvza_hfq4g256_k2048_all_buffer_gfx1151", kernels::FUSED_QKVZA_HFQ4G256_K2048_ALL_BUFFER_GFX1151_SRC, ["fused_qkvza_hfq4g256_k2048_all_buffer_gfx1151"]); // gemm.rs:3811 - add!("fused_qkvza_hfq4g256_k2048_hybrid_buffer_gfx1151", kernels::FUSED_QKVZA_HFQ4G256_K2048_HYBRID_BUFFER_GFX1151_SRC, ["fused_qkvza_hfq4g256_k2048_hybrid_buffer_gfx1151"]); // gemm.rs:3800 - add!("gated_norm_mq_rotate_gfx1151", kernels::GATED_NORM_MQ_ROTATE_GFX1151_SRC, ["gated_norm_mq_rotate_gfx1151"]); // gemv.rs:4783 - add!("gemv_hfq4g256_lm_head_r1_hybrid_buffer_gfx1151", kernels::GEMV_HFQ4G256_LM_HEAD_R1_HYBRID_BUFFER_GFX1151_SRC, ["gemv_hfq4g256_lm_head_r1_hybrid_buffer_gfx1151"]); // gemv.rs:9892 - add!("gemv_hfq4g256_residual_rt_low_gfx1151", kernels::GEMV_HFQ4G256_RESIDUAL_RT_LOW_GFX1151_SRC, ["gemv_hfq4g256_residual_rt_low_gfx1151"]); // gemv.rs:10226 - add!("moe_down_combine_rmsnorm_mq_rotate_vecsum_gfx1151", kernels::MOE_DOWN_COMBINE_RMSNORM_MQ_ROTATE_VECSUM_GFX1151_SRC, ["moe_down_combine_rmsnorm_mq_rotate_vecsum_gfx1151"]); // moe.rs:115 - add!("qwen35_fa_prep_gfx1151", kernels::QWEN35_FA_PREP_GFX1151_SRC, ["qwen35_fa_prep_gfx1151"]); // norm.rs:1146 - // DeepSeek4 MQ2R gfx1151 v2 route (registry/deepseek4-mq2r-gfx1151-v2.json). - add!("compressor_add_ape_buf", kernels::COMPRESSOR_ADD_APE_BATCHED_SRC, ["compressor_add_ape_f32_buf"]); // attention.rs:14675 - add!("compressor_overlap_concat", kernels::COMPRESSOR_OVERLAP_CONCAT_SRC, ["compressor_overlap_concat_f32"]); // attention.rs:14720 - add!("compressor_softmax_pool_f32_buf", kernels::COMPRESSOR_SOFTMAX_POOL_BUF_SRC, ["compressor_softmax_pool_f32_buf"]); // attention.rs:14852 - add!("deepseek4_attn_swa_buf", kernels::V4F_ATTN_SWA_BUF_SRC, ["deepseek4_attn_swa_buf"]); // attention.rs:18544 - add!("deepseek4_attn_swa_topk_scoregrid_f32_buf", kernels::V4F_ATTN_SWA_TOPK_BUF_SRC, ["deepseek4_attn_swa_topk_scoregrid_f32_buf"]); // attention.rs:19068 - add!("deepseek4_fused_silu_mul_clamp_mq_rotate", kernels::V4F_FUSED_SILU_MUL_CLAMP_MQ_ROTATE_SRC, ["deepseek4_fused_silu_mul_clamp_mq_rotate"]); // norm.rs:6165 - add!("deepseek4_moe_topk_bias_aware", kernels::V4F_MOE_TOPK_BIAS_AWARE_SRC, ["deepseek4_moe_topk_bias_aware_f32"]); // moe.rs:878 - add!("deepseek4_silu_mul_clamp", kernels::V4F_SILU_MUL_CLAMP_SRC, ["deepseek4_silu_mul_clamp_f32"]); // norm.rs:6224 - add!("deepseek4_topk_kv_gather_f32_buf", kernels::V4F_TOPK_KV_GATHER_BUF_SRC, ["deepseek4_topk_kv_gather_f32_buf"]); // moe.rs:1159 - add!("deepseek4_topk_kv_gather_identity_f32_buf", kernels::V4F_TOPK_KV_GATHER_IDENTITY_BUF_SRC, ["deepseek4_topk_kv_gather_identity_f32_buf"]); // moe.rs:1410 - add!("fused_rmsnorm_mq_rotate_plain_nox", kernels::FUSED_RMSNORM_MQ_ROTATE_PLAIN_SRC, ["fused_rmsnorm_mq_rotate_plain_nox"]); // norm.rs:5946 - add!("gemv_mfp4g32_e8_soa_grouped_gfx1151", kernels::GEMV_MFP4G32_E8_SOA_GROUPED_GFX1151_SRC, ["gemv_mfp4g32_e8_soa_grouped_gfx1151"]); // gemv.rs:7564,8025 - add!("gemv_mfp4g32_e8_soa_u4", kernels::GEMV_MFP4G32_E8_SOA_U4_SRC, ["gemv_mfp4g32_e8_soa_u4"]); // gemv.rs:8075,7448 - add!("gemv_mq2g256_lloyd_moe_down_residual_scaled_k8all_indexed", kernels::GEMV_MQ2G256_LLOYD_MOE_DOWN_INDEXED_SRC, ["gemv_mq2g256_lloyd_moe_down_residual_scaled_k8all_indexed"]); // gemv.rs:16927 - add!("gemv_mq2g256_lloyd_moe_gate_up_indexed", kernels::GEMV_MQ2G256_LLOYD_MOE_GATE_UP_INDEXED_SRC, ["gemv_mq2g256_lloyd_moe_gate_up_k8_indexed"]); // gemv.rs:17258 - add!("hash_router_normalize_f32_buf", kernels::HASH_ROUTER_NORMALIZE_BUF_SRC, ["hash_router_normalize_f32_buf"]); // moe.rs:770 - add!("hc_compute_control_vec4_finalize", kernels::HC_COMPUTE_CONTROL_SRC, ["hc_compute_control_vec4_finalize"]); // attention.rs:15414 - add!("hc_head_compute_pre", kernels::HC_HEAD_COMPUTE_PRE_SRC, ["hc_head_compute_pre"]); // attention.rs:15697 - add!("hc_input_map_4stream", kernels::HC_INPUT_MAP_SRC, ["hc_input_map_4stream"]); // attention.rs:15751 - add!("hc_mix_4stream", kernels::HC_MIX_4STREAM_SRC, ["hc_mix_4stream"]); // attention.rs:15835 - add!("indexer_relu_score_f32_buf", kernels::INDEXER_RELU_SCORE_BUF_SRC, ["indexer_relu_score_f32_buf"]); // attention.rs:16916 - add!("indexer_top_k_buf_parallel", kernels::INDEXER_TOP_K_BUF_PARALLEL_GFX1151_SRC, ["indexer_top_k_buf_parallel"]); // attention.rs:17320 - add!("rmsnorm_f32_at_slot_buf", kernels::RMSNORM_AT_SLOT_BUF_SRC, ["rmsnorm_f32_at_slot_buf"]); // norm.rs:6089 - add!("rope_tail_interleaved", kernels::ROPE_TAIL_INTERLEAVED_SRC, ["rope_tail_interleaved_f32"]); // attention.rs:17562 - add!("rope_tail_yarn_interleaved_at_slot_buf", kernels::ROPE_TAIL_YARN_INTERLEAVED_AT_SLOT_BUF_SRC, ["rope_tail_yarn_interleaved_at_slot_buf_f32"]); // attention.rs:17838 - add!("rope_tail_yarn_interleaved_wide", kernels::ROPE_TAIL_YARN_INTERLEAVED_SRC, ["rope_tail_yarn_interleaved_wide_f32"]); // attention.rs:17768 - add!("sqrt_softplus_f32", kernels::SQRT_SOFTPLUS_F32_SRC, ["sqrt_softplus_f32"]); // norm.rs:6128 - add!("state_overlap_shift_f32_buf", kernels::STATE_OVERLAP_SHIFT_F32_BUF_SRC, ["state_overlap_shift_f32_buf"]); // attention.rs:17988 - add!("state_ring_write_f32_buf", kernels::STATE_RING_WRITE_F32_BUF_SRC, ["state_ring_write_f32_buf"]); // attention.rs:18031 - add!("swa_ring_write_f32_buf", kernels::SWA_RING_WRITE_BUF_SRC, ["swa_ring_write_f32_buf"]); // attention.rs:18171 + add!( + "attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151", + kernels::ATTENTION_FLASH_Q8_0_REDUCE_GATED_MQ_ROTATE_GFX1151_SRC, + ["attention_flash_q8_0_reduce_gated_mq_rotate_gfx1151"] + ); // attention.rs:10620 + add!( + "fused_qkvza_hfq4g256_k2048_all_buffer_gfx1151", + kernels::FUSED_QKVZA_HFQ4G256_K2048_ALL_BUFFER_GFX1151_SRC, + ["fused_qkvza_hfq4g256_k2048_all_buffer_gfx1151"] + ); // gemm.rs:3811 + add!( + "fused_qkvza_hfq4g256_k2048_hybrid_buffer_gfx1151", + kernels::FUSED_QKVZA_HFQ4G256_K2048_HYBRID_BUFFER_GFX1151_SRC, + ["fused_qkvza_hfq4g256_k2048_hybrid_buffer_gfx1151"] + ); // gemm.rs:3800 + add!( + "gated_norm_mq_rotate_gfx1151", + kernels::GATED_NORM_MQ_ROTATE_GFX1151_SRC, + ["gated_norm_mq_rotate_gfx1151"] + ); // gemv.rs:4783 + add!( + "gemv_hfq4g256_lm_head_r1_hybrid_buffer_gfx1151", + kernels::GEMV_HFQ4G256_LM_HEAD_R1_HYBRID_BUFFER_GFX1151_SRC, + ["gemv_hfq4g256_lm_head_r1_hybrid_buffer_gfx1151"] + ); // gemv.rs:9892 + add!( + "gemv_hfq4g256_residual_rt_low_gfx1151", + kernels::GEMV_HFQ4G256_RESIDUAL_RT_LOW_GFX1151_SRC, + ["gemv_hfq4g256_residual_rt_low_gfx1151"] + ); // gemv.rs:10226 + add!( + "moe_down_combine_rmsnorm_mq_rotate_vecsum_gfx1151", + kernels::MOE_DOWN_COMBINE_RMSNORM_MQ_ROTATE_VECSUM_GFX1151_SRC, + ["moe_down_combine_rmsnorm_mq_rotate_vecsum_gfx1151"] + ); // moe.rs:115 + add!( + "qwen35_fa_prep_gfx1151", + kernels::QWEN35_FA_PREP_GFX1151_SRC, + ["qwen35_fa_prep_gfx1151"] + ); // norm.rs:1146 + // DeepSeek4 MQ2R gfx1151 v2 route (registry/deepseek4-mq2r-gfx1151-v2.json). + add!( + "compressor_add_ape_buf", + kernels::COMPRESSOR_ADD_APE_BATCHED_SRC, + ["compressor_add_ape_f32_buf"] + ); // attention.rs:14675 + add!( + "compressor_overlap_concat", + kernels::COMPRESSOR_OVERLAP_CONCAT_SRC, + ["compressor_overlap_concat_f32"] + ); // attention.rs:14720 + add!( + "compressor_softmax_pool_f32_buf", + kernels::COMPRESSOR_SOFTMAX_POOL_BUF_SRC, + ["compressor_softmax_pool_f32_buf"] + ); // attention.rs:14852 + add!( + "deepseek4_attn_swa_buf", + kernels::V4F_ATTN_SWA_BUF_SRC, + ["deepseek4_attn_swa_buf"] + ); // attention.rs:18544 + add!( + "deepseek4_attn_swa_topk_scoregrid_f32_buf", + kernels::V4F_ATTN_SWA_TOPK_BUF_SRC, + ["deepseek4_attn_swa_topk_scoregrid_f32_buf"] + ); // attention.rs:19068 + add!( + "deepseek4_fused_silu_mul_clamp_mq_rotate", + kernels::V4F_FUSED_SILU_MUL_CLAMP_MQ_ROTATE_SRC, + ["deepseek4_fused_silu_mul_clamp_mq_rotate"] + ); // norm.rs:6165 + add!( + "deepseek4_moe_topk_bias_aware", + kernels::V4F_MOE_TOPK_BIAS_AWARE_SRC, + ["deepseek4_moe_topk_bias_aware_f32"] + ); // moe.rs:878 + add!( + "deepseek4_silu_mul_clamp", + kernels::V4F_SILU_MUL_CLAMP_SRC, + ["deepseek4_silu_mul_clamp_f32"] + ); // norm.rs:6224 + add!( + "deepseek4_topk_kv_gather_f32_buf", + kernels::V4F_TOPK_KV_GATHER_BUF_SRC, + ["deepseek4_topk_kv_gather_f32_buf"] + ); // moe.rs:1159 + add!( + "deepseek4_topk_kv_gather_identity_f32_buf", + kernels::V4F_TOPK_KV_GATHER_IDENTITY_BUF_SRC, + ["deepseek4_topk_kv_gather_identity_f32_buf"] + ); // moe.rs:1410 + add!( + "fused_rmsnorm_mq_rotate_plain_nox", + kernels::FUSED_RMSNORM_MQ_ROTATE_PLAIN_SRC, + ["fused_rmsnorm_mq_rotate_plain_nox"] + ); // norm.rs:5946 + add!( + "gemv_mfp4g32_e8_soa_grouped_gfx1151", + kernels::GEMV_MFP4G32_E8_SOA_GROUPED_GFX1151_SRC, + ["gemv_mfp4g32_e8_soa_grouped_gfx1151"] + ); // gemv.rs:7564,8025 + add!( + "gemv_mfp4g32_e8_soa_u4", + kernels::GEMV_MFP4G32_E8_SOA_U4_SRC, + ["gemv_mfp4g32_e8_soa_u4"] + ); // gemv.rs:8075,7448 + add!( + "gemv_mq2g256_lloyd_moe_down_residual_scaled_k8all_indexed", + kernels::GEMV_MQ2G256_LLOYD_MOE_DOWN_INDEXED_SRC, + ["gemv_mq2g256_lloyd_moe_down_residual_scaled_k8all_indexed"] + ); // gemv.rs:16927 + add!( + "gemv_mq2g256_lloyd_moe_gate_up_indexed", + kernels::GEMV_MQ2G256_LLOYD_MOE_GATE_UP_INDEXED_SRC, + ["gemv_mq2g256_lloyd_moe_gate_up_k8_indexed"] + ); // gemv.rs:17258 + add!( + "hash_router_normalize_f32_buf", + kernels::HASH_ROUTER_NORMALIZE_BUF_SRC, + ["hash_router_normalize_f32_buf"] + ); // moe.rs:770 + add!( + "hc_compute_control_vec4_finalize", + kernels::HC_COMPUTE_CONTROL_SRC, + ["hc_compute_control_vec4_finalize"] + ); // attention.rs:15414 + add!( + "hc_head_compute_pre", + kernels::HC_HEAD_COMPUTE_PRE_SRC, + ["hc_head_compute_pre"] + ); // attention.rs:15697 + add!( + "hc_input_map_4stream", + kernels::HC_INPUT_MAP_SRC, + ["hc_input_map_4stream"] + ); // attention.rs:15751 + add!( + "hc_mix_4stream", + kernels::HC_MIX_4STREAM_SRC, + ["hc_mix_4stream"] + ); // attention.rs:15835 + add!( + "indexer_relu_score_f32_buf", + kernels::INDEXER_RELU_SCORE_BUF_SRC, + ["indexer_relu_score_f32_buf"] + ); // attention.rs:16916 + add!( + "indexer_top_k_buf_parallel", + kernels::INDEXER_TOP_K_BUF_PARALLEL_GFX1151_SRC, + ["indexer_top_k_buf_parallel"] + ); // attention.rs:17320 + add!( + "rmsnorm_f32_at_slot_buf", + kernels::RMSNORM_AT_SLOT_BUF_SRC, + ["rmsnorm_f32_at_slot_buf"] + ); // norm.rs:6089 + add!( + "rope_tail_interleaved", + kernels::ROPE_TAIL_INTERLEAVED_SRC, + ["rope_tail_interleaved_f32"] + ); // attention.rs:17562 + add!( + "rope_tail_yarn_interleaved_at_slot_buf", + kernels::ROPE_TAIL_YARN_INTERLEAVED_AT_SLOT_BUF_SRC, + ["rope_tail_yarn_interleaved_at_slot_buf_f32"] + ); // attention.rs:17838 + add!( + "rope_tail_yarn_interleaved_wide", + kernels::ROPE_TAIL_YARN_INTERLEAVED_SRC, + ["rope_tail_yarn_interleaved_wide_f32"] + ); // attention.rs:17768 + add!( + "sqrt_softplus_f32", + kernels::SQRT_SOFTPLUS_F32_SRC, + ["sqrt_softplus_f32"] + ); // norm.rs:6128 + add!( + "state_overlap_shift_f32_buf", + kernels::STATE_OVERLAP_SHIFT_F32_BUF_SRC, + ["state_overlap_shift_f32_buf"] + ); // attention.rs:17988 + add!( + "state_ring_write_f32_buf", + kernels::STATE_RING_WRITE_F32_BUF_SRC, + ["state_ring_write_f32_buf"] + ); // attention.rs:18031 + add!( + "swa_ring_write_f32_buf", + kernels::SWA_RING_WRITE_BUF_SRC, + ["swa_ring_write_f32_buf"] + ); // attention.rs:18171 } "gfx1201" => { - add!("attention_flash_fp8_e4m3_tile_gqa_gfx1201", kernels::ATTENTION_FLASH_FP8_E4M3_TILE_GQA_GFX1201_SRC, ["attention_flash_fp8_e4m3_tile_gqa_gfx1201"]); // attention.rs:7783 - add!("attention_flash_reduce_dsplit_gfx1201", kernels::ATTENTION_FLASH_REDUCE_DSPLIT_GFX1201_SRC, ["attention_flash_reduce_dsplit_gfx1201"]); // attention.rs:7784 - add!("gemv_hfq4g256_multirow_default", kernels::GEMV_HFQ4G256_MULTIROW_SRC, ["gemv_hfq4g256_multirow_r2"]); // gemv.rs:10031 - // Qwen3.6-27B decode fusions on exact gfx1201 (3be1cdbe9): the H2 tape's compact-3 GDN, - // 48-head AWQ gated-norm/MQ rotation and 24Q/4K FA prep. - add!("gated_delta_net_q8_compact3_b2", kernels::GATED_DELTA_NET_Q8_COMPACT3_B2_SRC, ["gated_delta_net_q8_compact3_b2"]); // norm.rs:3045 - add!("gated_norm_mq_rotate_awq_k6144_gfx1201", kernels::gated_norm_mq_rotate_awq_k6144_gfx1201_src(), ["gated_norm_mq_rotate_awq_k6144_gfx1201"]); // gemv.rs:4751 + add!( + "attention_flash_fp8_e4m3_tile_gqa_gfx1201", + kernels::ATTENTION_FLASH_FP8_E4M3_TILE_GQA_GFX1201_SRC, + ["attention_flash_fp8_e4m3_tile_gqa_gfx1201"] + ); // attention.rs:7783 + add!( + "attention_flash_reduce_dsplit_gfx1201", + kernels::ATTENTION_FLASH_REDUCE_DSPLIT_GFX1201_SRC, + ["attention_flash_reduce_dsplit_gfx1201"] + ); // attention.rs:7784 + add!( + "gemv_hfq4g256_multirow_default", + kernels::GEMV_HFQ4G256_MULTIROW_SRC, + ["gemv_hfq4g256_multirow_r2"] + ); // gemv.rs:10031 + // Qwen3.6-27B decode fusions on exact gfx1201 (3be1cdbe9): the H2 tape's compact-3 GDN, + // 48-head AWQ gated-norm/MQ rotation and 24Q/4K FA prep. + add!( + "gated_delta_net_q8_compact3_b2", + kernels::GATED_DELTA_NET_Q8_COMPACT3_B2_SRC, + ["gated_delta_net_q8_compact3_b2"] + ); // norm.rs:3045 + add!( + "gated_norm_mq_rotate_awq_k6144_gfx1201", + kernels::gated_norm_mq_rotate_awq_k6144_gfx1201_src(), + ["gated_norm_mq_rotate_awq_k6144_gfx1201"] + ); // gemv.rs:4751 #[cfg(feature = "deltanet")] - add!("qwen36_27b_fa_prep_gfx1201", kernels::qwen36_27b_fa_prep_gfx1201_src(), ["qwen36_27b_fa_prep_gfx1201"]); // norm.rs:1149 - add!("moe_router_softmax_topk_k8_wave64", kernels::MOE_ROUTER_SOFTMAX_TOPK_K8_WAVE64_SRC, ["moe_router_softmax_topk_k8_wave64"]); // gemv.rs:12992 + add!( + "qwen36_27b_fa_prep_gfx1201", + kernels::qwen36_27b_fa_prep_gfx1201_src(), + ["qwen36_27b_fa_prep_gfx1201"] + ); // norm.rs:1149 + add!( + "moe_router_softmax_topk_k8_wave64", + kernels::MOE_ROUTER_SOFTMAX_TOPK_K8_WAVE64_SRC, + ["moe_router_softmax_topk_k8_wave64"] + ); // gemv.rs:12992 } _ => {} } @@ -1077,7 +3399,13 @@ pub fn railgun_entries(arch: &str, extra_flags: &str) -> Result .ok_or_else(|| RegistryError::UnsupportedArch(arch.to_owned()))?; let mut entries = Vec::new(); if matches!(arch, "gfx1100" | "gfx1151" | "gfx1201") { - entries.push(entry(arch, "railgun_copy", &["railgun_copy", "railgun_copy_batch"], kernels::RAILGUN_COPY_SRC.into(), extra_flags)); + entries.push(entry( + arch, + "railgun_copy", + &["railgun_copy", "railgun_copy_batch"], + kernels::RAILGUN_COPY_SRC.into(), + extra_flags, + )); } Ok(entries) } @@ -1090,7 +3418,10 @@ pub fn railgun_entries(arch: &str, extra_flags: &str) -> Result pub fn corpus_entries(arch: &str, extra_flags: &str) -> Result, RegistryError> { let mut all = entries(arch, extra_flags)?; for entry in default_route_entries(arch, extra_flags)? { - if !all.iter().any(|e| e.module == entry.module && e.source == entry.source) { + if !all + .iter() + .any(|e| e.module == entry.module && e.source == entry.source) + { all.push(entry); } } @@ -1119,17 +3450,29 @@ mod tests { fn mq4v2_k5120_inventory_preserves_generic_and_limits_arch() { let module = "fused_gate_up_mq4g256v2_k5120_gfx1100"; let registry = entries("gfx1100", "").unwrap(); - let candidate = registry.iter().find(|entry| entry.module == module).unwrap(); + let candidate = registry + .iter() + .find(|entry| entry.module == module) + .unwrap(); assert_eq!(candidate.symbols, [module]); - let body = candidate.source().strip_prefix( - "#define HIPFIRE_FUSED_GATE_UP_KERNEL fused_gate_up_mq4g256v2_k5120_gfx1100\n" - ).unwrap(); + let body = candidate + .source() + .strip_prefix( + "#define HIPFIRE_FUSED_GATE_UP_KERNEL fused_gate_up_mq4g256v2_k5120_gfx1100\n", + ) + .unwrap(); assert_eq!( - body.replace("const int groups_per_row = 20;", "const int groups_per_row = K / 256;"), + body.replace( + "const int groups_per_row = 20;", + "const int groups_per_row = K / 256;" + ), kernels::FUSED_GATE_UP_MQ4G256V2_SRC ); for arch in ["gfx1151", "gfx1201", "gfx906", "gfx942"] { - assert!(!entries(arch, "").unwrap().iter().any(|entry| entry.module == module)); + assert!(!entries(arch, "") + .unwrap() + .iter() + .any(|entry| entry.module == module)); } } @@ -1147,8 +3490,14 @@ mod tests { let mut count = 0; for line in expected.lines().filter(|line| !line.starts_with('#')) { let (module, digest) = line.split_once('\t').unwrap(); - let entry = by_name.get(module).unwrap_or_else(|| panic!("missing P0 module {module}")); - assert_eq!(format!("{:x}", Sha256::digest(entry.source().as_bytes())), digest, "{module}"); + let entry = by_name + .get(module) + .unwrap_or_else(|| panic!("missing P0 module {module}")); + assert_eq!( + format!("{:x}", Sha256::digest(entry.source().as_bytes())), + digest, + "{module}" + ); count += 1; } assert_eq!(count, 92); @@ -1161,12 +3510,18 @@ mod tests { // (tests/fixtures/kernel-trace-qwen4-flash-next.tsv) and the 32 // Qwen3.5-MoE modules (tests/fixtures/kernel-trace-qwen35.tsv). Those // 113 keys are additional to P0's 92. - assert_eq!(registry.len(), count + 113, "unexpected gfx1201 inventory size"); - let default_prefill = by_name.get("attention_q8_0_flash_prefill_br8_bc16").unwrap(); + assert_eq!( + registry.len(), + count + 113, + "unexpected gfx1201 inventory size" + ); + let default_prefill = by_name + .get("attention_q8_0_flash_prefill_br8_bc16") + .unwrap(); assert_eq!(default_prefill.symbols, ["attention_q8_0_flash_prefill"]); - assert!(default_prefill.source().starts_with( - "#define BR 8\n#define BC 16\n#define NTHREADS 256\n" - )); + assert!(default_prefill + .source() + .starts_with("#define BR 8\n#define BC 16\n#define NTHREADS 256\n")); assert_eq!( format!("{:x}", Sha256::digest(default_prefill.source().as_bytes())), "680c37bf2f4f978d0361bd5c1b48c1fbd6430d1777663d31b45c9b6ad510bb95", @@ -1174,7 +3529,10 @@ mod tests { ); assert_ne!( default_prefill.source(), - by_name.get("attention_q8_0_flash_prefill").unwrap().source() + by_name + .get("attention_q8_0_flash_prefill") + .unwrap() + .source() ); if !root.is_dir() { // P0 cache paths are machine-local. The 92 digests above are @@ -1213,37 +3571,66 @@ mod tests { const FLAGS_ADDED_SINCE_P0: [&str; 1] = ["-fuse-cuid=none"]; let mut repins_seen = HashSet::new(); for (trace, expected_count) in [("cold1", 35), ("smokecold1", 35), ("preinstall", 58)] { - let records = std::fs::read_to_string(root.join(format!("{trace}.hipcc.jsonl"))).unwrap(); + let records = + std::fs::read_to_string(root.join(format!("{trace}.hipcc.jsonl"))).unwrap(); let mut observed = 0; for line in records.lines() { let record: serde_json::Value = serde_json::from_str(line).unwrap(); - let Some(module) = record["module"].as_str() else { continue }; + let Some(module) = record["module"].as_str() else { + continue; + }; let source_path = record["source"].as_str().unwrap(); let expected = std::fs::read(source_path).unwrap(); - let entry = by_name.get(module).unwrap_or_else(|| panic!("missing {trace} module {module}")); + let entry = by_name + .get(module) + .unwrap_or_else(|| panic!("missing {trace} module {module}")); if let Some(&repinned) = REPINNED_SINCE_P0.iter().find(|name| **name == module) { - assert_ne!(entry.source().as_bytes(), expected, - "{trace}: {module} matches the trace again; drop it from REPINNED_SINCE_P0"); + assert_ne!( + entry.source().as_bytes(), + expected, + "{trace}: {module} matches the trace again; drop it from REPINNED_SINCE_P0" + ); repins_seen.insert(repinned); } else { - assert_eq!(entry.source().as_bytes(), expected, "{trace}: {module} source differs"); + assert_eq!( + entry.source().as_bytes(), + expected, + "{trace}: {module} source differs" + ); } let argv = record["argv"].as_array().unwrap(); - let expected_flags = argv.iter().map(|arg| arg.as_str().unwrap()) + let expected_flags = argv + .iter() + .map(|arg| arg.as_str().unwrap()) .take_while(|arg| *arg != "-o") - .filter(|arg| !arg.starts_with("--rocm-path=") - && !arg.starts_with("--hip-path=") && !arg.starts_with("-I")) + .filter(|arg| { + !arg.starts_with("--rocm-path=") + && !arg.starts_with("--hip-path=") + && !arg.starts_with("-I") + }) .collect::>(); - let flags = entry.flags.iter().map(String::as_str) + let flags = entry + .flags + .iter() + .map(String::as_str) .filter(|flag| !FLAGS_ADDED_SINCE_P0.contains(flag)) .collect::>(); - assert_eq!(flags, expected_flags, "{trace}: {module} core hipcc flags differ"); + assert_eq!( + flags, expected_flags, + "{trace}: {module} core hipcc flags differ" + ); observed += 1; } - assert_eq!(observed, expected_count, "{trace}: incomplete compiler trace"); + assert_eq!( + observed, expected_count, + "{trace}: incomplete compiler trace" + ); } - assert_eq!(repins_seen.len(), REPINNED_SINCE_P0.len(), - "REPINNED_SINCE_P0 names a module no trace compiled: {repins_seen:?}"); + assert_eq!( + repins_seen.len(), + REPINNED_SINCE_P0.len(), + "REPINNED_SINCE_P0 names a module no trace compiled: {repins_seen:?}" + ); } #[test] @@ -1252,15 +3639,25 @@ mod tests { let entries = corpus_entries(arch, "").unwrap(); let mut seen = std::collections::HashSet::new(); for entry in &entries { - assert!(seen.insert(entry.module), "{arch}: module {} inventoried twice", entry.module); + assert!( + seen.insert(entry.module), + "{arch}: module {} inventoried twice", + entry.module + ); } } } #[test] fn kernel_registry_explicit_arch_and_module_rejection() { - assert!(matches!(lookup("gfx1200", "rmsnorm", ""), Err(RegistryError::UnsupportedArch(_)))); - assert!(matches!(lookup("gfx942", "qwen35_fa_prep_batched_gfx1201", ""), Err(RegistryError::UnsupportedModule { .. }))); + assert!(matches!( + lookup("gfx1200", "rmsnorm", ""), + Err(RegistryError::UnsupportedArch(_)) + )); + assert!(matches!( + lookup("gfx942", "qwen35_fa_prep_batched_gfx1201", ""), + Err(RegistryError::UnsupportedModule { .. }) + )); assert!(lookup("gfx942", "fused_qkv_hfq4g256_wave64", "").is_ok()); assert!(lookup("gfx906", "gemm_qkv_hfq4g256_wmma_gfx12", "").is_err()); for arch in SUPPORTED_ARCHES { @@ -1269,6 +3666,37 @@ mod tests { } } + #[test] + fn gfx1100_fa2_split_verifier_is_compiler_free_packaged() { + let entry = lookup("gfx1100", "attention_q8_0_fa2_gqa_gfx1100", "").unwrap(); + for symbol in [ + "attention_fa2_q_preconvert_gfx1100", + "attention_q8_0_fa2_gqa_gfx1100", + "attention_q8_0_fa2_gqa_partial_gfx1100", + "attention_q8_0_fa2_gqa_merge_gfx1100", + ] { + assert!( + entry.symbols.contains(&symbol), + "gfx1100 FA2 package is missing {symbol}" + ); + assert!( + entry.source().contains(symbol), + "gfx1100 FA2 source lacks {symbol}" + ); + } + for arch in SUPPORTED_ARCHES.iter().filter(|arch| **arch != "gfx1100") { + let symbols = entries(arch, "") + .unwrap() + .into_iter() + .find(|candidate| candidate.module == "attention_q8_0_fa2_gqa_gfx1100") + .map(|candidate| candidate.symbols); + assert!( + symbols.is_none(), + "{arch} must not package the gfx1100 module" + ); + } + } + /// The opt-in Hyper inline-Q8 chunk module is inventoried on exact gfx1151 /// only, carries the source `gated_delta_step_gate_wmma` compiles, and /// exports the symbol it launches. @@ -1276,8 +3704,13 @@ mod tests { fn hyper_q8_chunk_module_is_exact_gfx1151_only() { let entry = lookup("gfx1151", "gated_delta_chunk_q8_wmma", "").unwrap(); assert_eq!(entry.symbols, ["gated_delta_chunk_gate_q8_wmma"]); - assert_eq!(entry.source(), crate::tensor_ops::GATED_DELTA_CHUNK_Q8_WMMA_SRC); - assert!(entry.source().contains("void gated_delta_chunk_gate_q8_wmma(")); + assert_eq!( + entry.source(), + crate::tensor_ops::GATED_DELTA_CHUNK_Q8_WMMA_SRC + ); + assert!(entry + .source() + .contains("void gated_delta_chunk_gate_q8_wmma(")); for arch in SUPPORTED_ARCHES.iter().filter(|arch| **arch != "gfx1151") { assert!( lookup(arch, "gated_delta_chunk_q8_wmma", "").is_err(), @@ -1296,13 +3729,18 @@ mod tests { let [arch, module, symbols] = line.split('\t').collect::>()[..] else { panic!("malformed trace row {line:?}"); }; - let inventory = by_arch.entry(arch).or_insert_with(|| entries(arch, "").unwrap()); + let inventory = by_arch + .entry(arch) + .or_insert_with(|| entries(arch, "").unwrap()); let entry = inventory .iter() .find(|entry| entry.module == module) .unwrap_or_else(|| panic!("{arch}: traced module {module} is not packaged")); for symbol in symbols.split(',') { - assert!(entry.symbols.contains(&symbol), "{arch} {module}: packaged symbols lack {symbol}"); + assert!( + entry.symbols.contains(&symbol), + "{arch} {module}: packaged symbols lack {symbol}" + ); } rows += 1; } @@ -1311,7 +3749,9 @@ mod tests { #[test] fn qwen4_flash_next_trace_is_packaged() { - assert_trace_packaged(include_str!("../tests/fixtures/kernel-trace-qwen4-flash-next.tsv")); + assert_trace_packaged(include_str!( + "../tests/fixtures/kernel-trace-qwen4-flash-next.tsv" + )); } #[test] diff --git a/crates/rdna-compute/src/mq_f16_producers.rs b/crates/rdna-compute/src/mq_f16_producers.rs index bb3e7e4bb3..236f561c33 100644 --- a/crates/rdna-compute/src/mq_f16_producers.rs +++ b/crates/rdna-compute/src/mq_f16_producers.rs @@ -471,6 +471,7 @@ impl Gpu { up_m: usize, k: usize, batch_size: usize, + split_verify_capture: bool, ) -> HipResult<()> { if !self.arch_caps.is_gfx1100() { return Err(hip_bridge::HipError::new( @@ -492,8 +493,7 @@ impl Gpu { // contracts for this exact launch, so its capture must preserve the eager // reduction route instead of silently switching to the base kernel. let recording = self.replay.is_recording() || self.graphs.capture_mode; - let recording_supported = - !recording || hipfire_config::developer_bool("HIPFIRE_GFX1100_FA2_SPLIT_VERIFY", false); + let recording_supported = !recording || split_verify_capture; let (kname, ksrc, block_x) = if recording_supported && self.arch_caps.is_gfx1100() && self.arch == "gfx1100" diff --git a/docs/CONFIG.md b/docs/CONFIG.md index 054f173d53..184959c165 100644 --- a/docs/CONFIG.md +++ b/docs/CONFIG.md @@ -836,7 +836,7 @@ Deprecated since 0.4.0, removal in 0.5.0: -**Count:** 272 schema keys, plus the `developer.` namespace. +**Count:** 273 schema keys, plus the `developer.` namespace. | Key | Legacy key | Env | Lifecycle | |---|---|---|---| @@ -936,6 +936,7 @@ Deprecated since 0.4.0, removal in 0.5.0: | `kernel.gemv_dp4a` | `gemv_dp4a` | `HIPFIRE_GEMV_DP4A` | stable | | `kernel.gemv_prefetch` | `gemv_prefetch` | `HIPFIRE_GEMV_PREFETCH` | stable | | `kernel.gfx1100_dec_norm` | `gfx1100_dec_norm` | `HIPFIRE_GFX1100_DEC_NORM` | stable | +| `kernel.gfx1100_fa2_split_verify` | `gfx1100_fa2_split_verify` | `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | stable | | `kernel.gfx11_a4_candidates` | `gfx11_a4_candidates` | `HIPFIRE_GFX11_A4_CANDIDATES` | experimental | | `kernel.gfx11_fa2_prefill` | `gfx11_fa2_prefill` | `HIPFIRE_GFX11_FA2_PREFILL` | stable | | `kernel.gfx11_iu4_gridspec` | `gfx11_iu4_gridspec` | `HIPFIRE_GFX11_IU4_GRIDSPEC` | stable | diff --git a/docs/env-vars.md b/docs/env-vars.md index d67aa9b121..e38ca0fc7d 100644 --- a/docs/env-vars.md +++ b/docs/env-vars.md @@ -216,7 +216,7 @@ Read only by the Qwen4 carrier and its kernels; no other model reads them. | `HIPFIRE_GFX11_FA2_PREFILL` | GQA-fused FA2 prefill on gfx1100/gfx1151 (Qwen NH24/NKV4/HD256, N 64..512 step 16, ctx 64..32768) — default ON (`kernel.gfx11_fa2_prefill`); `=0` opts out toward the byte-identical incumbent | | `HIPFIRE_FA2_FILL` | Warp-specialized K/V fill in that FA2 kernel on gfx1100/gfx1151 (bit-exact; helper waves dequantize the next K/V tile while compute waves run QK/PV) — default ON; `=0` restores the all-wave per-tile fill | | `HIPFIRE_GFX1100_FA2_R3` | Exact-gfx1100 variant of that FA2 fill body (bit-exact; CU mode, bank-conflict-free helper plane stores, O rescale skipped when alpha is exactly 1, heaviest q tiles first; symbols `attention_q8_0_fa2_gqa_gfx1100` / `attention_fa2_q_preconvert_gfx1100`) — default ON; `=0` restores the shared gfx11 body | -| `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | Experimental, default off, exact gfx1100: `1` replaces the Q8 DFlash verifier's eager R4/R8 attention with the FA2 split-KV S8 route for dense Qwen H24/NKV4/HD256 batches 4..32 once the **live logical context** exceeds 4,096. HipGraph/retained recording additionally require precompiled kernels and fully materialized fixed-address Q16 scratch; unsupported or not-yet-ready shapes fail closed to the established batched route. The Redline/PM4 product admission remains B=16 and inherits its existing single-GPU/state guards. | +| `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | Default **on** only on exact gfx1100 (`kernel.gfx1100_fa2_split_verify`); `0` opts out. Replaces the Q8 DFlash verifier's eager R4/R8 attention with the FA2 split-KV S8 route only for dense Qwen H24/NKV4/HD256 sequential batches 4..32 once the **live logical context** exceeds `HIPFIRE_FA_PERTOKEN_MIN_CTX` (4,096 by default; smaller non-zero overrides are clamped to the measured 4,096 crossover for this route). `HIPFIRE_VERIFY_ATTN=0` remains the parent opt-out. HipGraph/retained recording additionally require the compiler-free precompiled kernel set, fully materialized fixed-address Q16 scratch, and enough `flash_partials` capacity for the S8 record layout; a deliberately small `HIPFIRE_FLASH_PARTIALS_BATCH` therefore fails closed to the established batched route. The Redline/PM4 product admission remains B=16 and inherits its existing single-GPU/state guards. | | `HIPFIRE_GFX1151_FA2_TWIN` | Exact-gfx1151 twin of that FA2 fill kernel (CU mode, heaviest q-tile first, conflict-free helper V stores; bit-exact) — default ON; `=0` restores the gfx11 module | | `HIPFIRE_GFX12_FA2_PREFILL` | GQA-fused FA2 prefill on exact gfx1201 (same Qwen NH24/NKV4/HD256 envelope) — default ON (`kernel.gfx12_fa2_prefill`); `=0` opts out toward the byte-identical incumbent | | `HIPFIRE_GFX12_FA_PACKET` | Packet-minimal Q128 FA2 body on exact gfx1201 (same Qwen envelope as `HIPFIRE_GFX12_FA2_PREFILL`) — default ON (`kernel.gfx12_fa_packet`); `=0` opts out to the byte-identical route-N body | @@ -508,7 +508,7 @@ Presence in the inventory means the token appears in source; it does **not** mea **Generation method:** token scan over tracked `*.rs`, `*.py`, and `*.sh` (`scripts/check-lifecycle.py --write`). **Columns:** variable; up to two lexical source paths; lifecycle status (see [Lifecycle status](#lifecycle-status)). -**Count:** 1408 +**Count:** 1409 | Variable | Example source path(s) | Lifecycle | |---|---|---| @@ -1051,7 +1051,7 @@ Presence in the inventory means the token appears in source; it does **not** mea | `HIPFIRE_GFX1100_DENSE_GATE_UP_SETPRIO` | crates/rdna-compute/src/gemm.rs | developer | | `HIPFIRE_GFX1100_DENSE_GATE_UP_STAGE_X32` | crates/rdna-compute/src/gemm.rs | developer | | `HIPFIRE_GFX1100_FA2_R3` | crates/rdna-compute/src/attention.rs | developer | -| `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | crates/hipfire-arch-qwen35/src/dflash_spec.rs, crates/hipfire-arch-qwen35/src/qwen35/prefill.rs | developer | +| `HIPFIRE_GFX1100_FA2_SPLIT_VERIFY` | crates/hipfire-config/src/lib.rs, crates/rdna-compute/src/feature_flags.rs | stable | | `HIPFIRE_GFX1100_FA_PREP` | crates/hipfire-arch-qwen35/src/qwen35/prefill.rs | developer | | `HIPFIRE_GFX1100_GATED_NORM_V2` | crates/rdna-compute/src/gemv.rs, crates/rdna-compute/src/kernels.rs | developer | | `HIPFIRE_GFX1100_MQ4V2_NINEPATH_RPB8` | crates/hipfire-runtime/examples/mq4v2_fused_parity.rs | harness | @@ -1765,13 +1765,13 @@ Presence in the inventory means the token appears in source; it does **not** mea | `HIPFIRE_ROUTE_ORACLE_KV_B` | crates/hipfire-arch-qwen35/tests/route_oracle_single.rs | harness | | `HIPFIRE_S4_FLAG_PROBE_UNSET_OFF` | crates/hipfire-config/src/lib.rs | developer | | `HIPFIRE_S4_FLAG_PROBE_UNSET_ON` | crates/hipfire-config/src/lib.rs | developer | -| `HIPFIRE_SAMPLED_MTP_ARM` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | -| `HIPFIRE_SAMPLED_MTP_MIN_P` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | -| `HIPFIRE_SAMPLED_MTP_OUT` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | -| `HIPFIRE_SAMPLED_MTP_TEMP` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | -| `HIPFIRE_SAMPLED_MTP_TOP_K` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | -| `HIPFIRE_SAMPLED_MTP_TOP_P` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | -| `HIPFIRE_SAMPLED_MTP_TRIALS` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs, crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_ARM` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_MIN_P` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_OUT` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_TEMP` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_TOP_K` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_TOP_P` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | +| `HIPFIRE_SAMPLED_MTP_TRIALS` | crates/hipfire-arch-qwen4/tests/sampled_mtp_distribution_hw.rs | harness | | `HIPFIRE_SAMPLE_COMPARE` | crates/hipfire-runtime/src/llama.rs, crates/saddle-lab/examples/infer_qwen35.rs | developer | | `HIPFIRE_SAMPLE_FAST` | crates/rdna-compute/examples/sample_parallel_stable_parity.rs, crates/rdna-compute/src/sampling.rs | developer | | `HIPFIRE_SAMPLE_PARALLEL` | crates/rdna-compute/examples/sample_accept_parity.rs, crates/rdna-compute/src/sampling.rs | developer |