use crate::runtime::CudaDeviceCapabilities;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum ArchTier {
Legacy,
Volta,
Turing,
Ampere,
Ada,
Hopper,
Blackwell,
}
impl ArchTier {
#[must_use]
pub fn from_compute_capability((major, minor): (u32, u32)) -> Self {
match (major, minor) {
(0..=6, _) => ArchTier::Legacy,
(7, 0..=2) => ArchTier::Volta,
(7, _) => ArchTier::Turing, (8, 0..=7) => ArchTier::Ampere,
(8, _) => ArchTier::Ada, (9, _) => ArchTier::Hopper,
(_, _) => ArchTier::Blackwell,
}
}
#[must_use]
pub fn has_tensor_cores(self) -> bool {
matches!(
self,
ArchTier::Ampere | ArchTier::Ada | ArchTier::Hopper | ArchTier::Blackwell
)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ArchConfig {
pub tier: ArchTier,
pub qmoe_tile_hint: u32,
pub resident_warps_per_sm: u32,
pub prefers_tensor_core: bool,
pub smem_budget_bytes: u32,
pub l2_residency_candidate: bool,
}
const DEFAULT_SMEM_BUDGET_BYTES: u32 = 48 * 1024;
impl ArchConfig {
#[must_use]
pub fn for_tier(tier: ArchTier) -> Self {
let (qmoe_tile_hint, resident_warps_per_sm) = match tier {
ArchTier::Legacy => (2, 48),
ArchTier::Volta => (4, 64), ArchTier::Turing => (4, 48),
ArchTier::Ampere => (8, 64), ArchTier::Ada => (8, 48), ArchTier::Hopper => (8, 64), ArchTier::Blackwell => (8, 64),
};
Self {
tier,
qmoe_tile_hint,
resident_warps_per_sm,
prefers_tensor_core: tier.has_tensor_cores(),
smem_budget_bytes: DEFAULT_SMEM_BUDGET_BYTES,
l2_residency_candidate: matches!(tier, ArchTier::Ada),
}
}
#[must_use]
pub fn for_capabilities(capabilities: CudaDeviceCapabilities) -> Self {
Self::for_tier(capabilities.arch_tier())
}
}
#[must_use]
pub fn decode_resident_warps_per_sm((major, minor): (u32, u32)) -> u32 {
match (major, minor) {
(8, 0) | (9.., _) => 64,
_ => 48,
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DecodeTilingProfile {
pub tier: ArchTier,
pub multiprocessor_count: u32,
pub resident_warps_per_sm: u32,
pub sm_count_split_k: bool,
}
impl DecodeTilingProfile {
#[must_use]
pub fn for_capabilities(capabilities: CudaDeviceCapabilities) -> Self {
let tier = capabilities.arch_tier();
Self {
tier,
multiprocessor_count: capabilities.multiprocessor_count(),
resident_warps_per_sm: decode_resident_warps_per_sm(capabilities.compute_capability()),
sm_count_split_k: !matches!(tier, ArchTier::Hopper),
}
}
#[must_use]
pub fn one_wave_ctas(self, threads_per_cta: u32) -> usize {
let warps_per_cta = (threads_per_cta / 32).max(1) as usize;
let resident_ctas = (self.resident_warps_per_sm as usize / warps_per_cta).max(1);
(self.multiprocessor_count.max(1) as usize).saturating_mul(resident_ctas)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime::CudaDeviceCapabilities;
#[test]
fn arch_tier_mapping_is_total_and_panic_free() {
for major in 0u32..=20 {
for minor in 0u32..=16 {
let tier = ArchTier::from_compute_capability((major, minor));
let _ = ArchConfig::for_tier(tier);
}
}
}
#[test]
fn known_compute_capabilities_map_to_expected_tiers() {
let cases = [
((6, 1), ArchTier::Legacy), ((7, 0), ArchTier::Volta), ((7, 5), ArchTier::Turing), ((8, 0), ArchTier::Ampere), ((8, 6), ArchTier::Ampere), ((8, 9), ArchTier::Ada), ((9, 0), ArchTier::Hopper), ((10, 0), ArchTier::Blackwell), ((12, 0), ArchTier::Blackwell), ];
for (cc, expected) in cases {
assert_eq!(
ArchTier::from_compute_capability(cc),
expected,
"cc {cc:?} mapped to the wrong tier"
);
}
}
#[test]
fn sm_90_hopper_config_is_frozen() {
let cfg = ArchConfig::for_tier(ArchTier::Hopper);
assert_eq!(cfg.tier, ArchTier::Hopper);
assert_eq!(cfg.qmoe_tile_hint, 8, "sm_90 QMoE tile must stay 8");
assert_eq!(
cfg.resident_warps_per_sm, 64,
"sm_90 resident warps must stay 64"
);
assert!(cfg.prefers_tensor_core, "sm_90 stays tensor-core eligible");
assert_eq!(cfg.smem_budget_bytes, DEFAULT_SMEM_BUDGET_BYTES);
assert!(
!cfg.l2_residency_candidate,
"L2 residency is an Ada lever, not a Hopper one"
);
let caps = CudaDeviceCapabilities::for_test((9, 0), 132, 50 * 1024 * 1024);
assert_eq!(caps.arch_tier(), ArchTier::Hopper);
assert_eq!(caps.arch_config(), cfg);
}
#[test]
fn ada_consumer_config_differs_from_hopper() {
let ada = ArchConfig::for_tier(ArchTier::Ada);
assert_eq!(ada.resident_warps_per_sm, 48);
assert!(ada.l2_residency_candidate);
assert!(ada.prefers_tensor_core);
let hopper = ArchConfig::for_tier(ArchTier::Hopper);
assert_ne!(ada, hopper);
}
#[test]
fn pre_sm80_tiers_have_no_tensor_cores() {
for tier in [ArchTier::Legacy, ArchTier::Volta, ArchTier::Turing] {
assert!(!ArchConfig::for_tier(tier).prefers_tensor_core);
}
for tier in [
ArchTier::Ampere,
ArchTier::Ada,
ArchTier::Hopper,
ArchTier::Blackwell,
] {
assert!(ArchConfig::for_tier(tier).prefers_tensor_core);
}
}
#[test]
fn decode_resident_warps_ladder_matches_frozen_selection() {
for cc in [(8, 0), (9, 0), (10, 0), (12, 0)] {
assert_eq!(
decode_resident_warps_per_sm(cc),
64,
"cc {cc:?} must stay on the 64-warp rung"
);
}
for cc in [(8, 6), (8, 7), (8, 9), (7, 5), (7, 0), (6, 1)] {
assert_eq!(
decode_resident_warps_per_sm(cc),
48,
"cc {cc:?} must stay on the 48-warp rung"
);
}
}
#[test]
fn sm_90_decode_profile_is_frozen_out_of_rtx_splitk() {
let h200 = CudaDeviceCapabilities::for_test((9, 0), 132, 50 * 1024 * 1024);
let profile = DecodeTilingProfile::for_capabilities(h200);
assert_eq!(profile.tier, ArchTier::Hopper);
assert_eq!(profile.resident_warps_per_sm, 64);
assert_eq!(profile.multiprocessor_count, 132);
assert!(
!profile.sm_count_split_k,
"sm_90/H200 must NOT opt into the RTX SM-count split-K lever (frozen)"
);
}
#[test]
fn rtx_consumer_profiles_opt_into_sm_count_split_k() {
let ada = DecodeTilingProfile::for_capabilities(CudaDeviceCapabilities::for_test(
(8, 9),
128,
72 * 1024 * 1024,
));
assert_eq!(ada.tier, ArchTier::Ada);
assert_eq!(ada.resident_warps_per_sm, 48);
assert!(
ada.sm_count_split_k,
"Ada consumer opts into SM-count split-K"
);
let ada_l4 = DecodeTilingProfile::for_capabilities(CudaDeviceCapabilities::for_test(
(8, 9),
58,
48 * 1024 * 1024,
));
assert!(ada_l4.sm_count_split_k);
assert!(
ada_l4.one_wave_ctas(256) < ada.one_wave_ctas(256),
"fewer SMs => smaller one-wave CTA target (split-K fills the grid sooner)"
);
let ampere = DecodeTilingProfile::for_capabilities(CudaDeviceCapabilities::for_test(
(8, 6),
68,
5 * 1024 * 1024,
));
assert_eq!(ampere.tier, ArchTier::Ampere);
assert_eq!(ampere.resident_warps_per_sm, 48);
assert!(
ampere.sm_count_split_k,
"Ampere consumer opts into SM-count split-K"
);
let h200 = DecodeTilingProfile::for_capabilities(CudaDeviceCapabilities::for_test(
(9, 0),
132,
50 * 1024 * 1024,
));
assert!(!h200.sm_count_split_k);
}
#[test]
fn one_wave_ctas_tracks_sm_count_and_cta_width() {
let profile = DecodeTilingProfile::for_capabilities(CudaDeviceCapabilities::for_test(
(8, 9),
100,
48 * 1024 * 1024,
));
assert_eq!(profile.one_wave_ctas(256), 600);
assert_eq!(profile.one_wave_ctas(32), 4800);
let tiny =
DecodeTilingProfile::for_capabilities(CudaDeviceCapabilities::for_test((8, 9), 0, 0));
assert_eq!(tiny.one_wave_ctas(256), 6);
}
}