use std::fmt;
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum GpuVendor {
Nvidia,
Amd,
}
impl GpuVendor {
pub fn as_str(self) -> &'static str {
match self {
GpuVendor::Nvidia => "nvidia",
GpuVendor::Amd => "amd",
}
}
pub fn parse(s: &str) -> Option<Self> {
match s.trim().to_ascii_lowercase().as_str() {
"nvidia" | "cuda" => Some(GpuVendor::Nvidia),
"amd" | "rocm" | "hip" => Some(GpuVendor::Amd),
_ => None,
}
}
pub fn cargo_feature(self) -> &'static str {
match self {
GpuVendor::Nvidia => "cuda",
GpuVendor::Amd => "rocm",
}
}
}
impl fmt::Display for GpuVendor {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
GpuVendor::Nvidia => "NVIDIA",
GpuVendor::Amd => "AMD",
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VariantClass {
Cpu,
Vendor(GpuVendor),
Unknown,
}
pub fn classify_variant_label(label: &str) -> VariantClass {
let basename = std::path::Path::new(label)
.file_name()
.and_then(|n| n.to_str())
.unwrap_or("");
let tagged = |prefix: &str| {
basename
.strip_prefix(prefix)
.is_some_and(|rest| rest.starts_with(|c: char| c.is_ascii_digit()))
};
if basename == "cpu" || basename.starts_with("cpu-") {
return VariantClass::Cpu;
}
if tagged("cu") || tagged("sm") {
return VariantClass::Vendor(GpuVendor::Nvidia);
}
if tagged("rocm") || tagged("gfx") {
return VariantClass::Vendor(GpuVendor::Amd);
}
VariantClass::Unknown
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum GpuArch {
Sm { major: u32, minor: u32 },
Gfx(String),
}
impl GpuArch {
pub fn parse(vendor: GpuVendor, token: &str) -> Option<Self> {
let t = token.trim();
match vendor {
GpuVendor::Nvidia => Self::parse_sm(t),
GpuVendor::Amd => {
let bare = t.split(':').next()?.trim().to_ascii_lowercase();
if !bare.starts_with("gfx") || bare.len() <= 3 {
return None;
}
Some(GpuArch::Gfx(bare))
}
}
}
fn parse_sm(t: &str) -> Option<Self> {
let t = t.trim().split('+').next()?.trim();
if let Some((maj, min)) = t.split_once('.') {
return Some(GpuArch::Sm {
major: maj.trim().parse().ok()?,
minor: min.trim().parse().ok()?,
});
}
let digits = t.trim_start_matches("sm_").trim_start_matches("sm").trim();
if digits.len() < 2 || !digits.chars().all(|c| c.is_ascii_digit()) {
return None;
}
let (maj, min) = digits.split_at(digits.len() - 1);
Some(GpuArch::Sm {
major: maj.parse().ok()?,
minor: min.parse().ok()?,
})
}
pub fn vendor(&self) -> GpuVendor {
match self {
GpuArch::Sm { .. } => GpuVendor::Nvidia,
GpuArch::Gfx(_) => GpuVendor::Amd,
}
}
pub fn sm_major(&self) -> Option<u32> {
match self {
GpuArch::Sm { major, .. } => Some(*major),
_ => None,
}
}
pub fn sm_minor(&self) -> Option<u32> {
match self {
GpuArch::Sm { minor, .. } => Some(*minor),
_ => None,
}
}
pub fn generation(&self) -> u32 {
match self {
GpuArch::Sm { major, minor } => major * 10 + minor,
GpuArch::Gfx(g) => g
.trim_start_matches("gfx")
.chars()
.take_while(|c| c.is_ascii_digit())
.collect::<String>()
.parse()
.unwrap_or(0),
}
}
pub fn archs_token(&self) -> String {
match self {
GpuArch::Sm { major, minor } => format!("{major}.{minor}"),
GpuArch::Gfx(g) => g.clone(),
}
}
pub fn covered_by(&self, archs: &str) -> bool {
archs.split([';', ',', ' ']).any(|token| {
match (self, GpuArch::parse(self.vendor(), token)) {
(GpuArch::Sm { major, .. }, Some(GpuArch::Sm { major: m, .. })) => m == *major,
(GpuArch::Gfx(gfx), Some(GpuArch::Gfx(g))) => g == *gfx,
_ => false,
}
})
}
}
impl fmt::Display for GpuArch {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
GpuArch::Sm { major, minor } => write!(f, "sm_{major}{minor}"),
GpuArch::Gfx(g) => f.write_str(g),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn vendor_round_trips_and_accepts_stack_names() {
assert_eq!(GpuVendor::parse("nvidia"), Some(GpuVendor::Nvidia));
assert_eq!(GpuVendor::parse(" CUDA "), Some(GpuVendor::Nvidia));
assert_eq!(GpuVendor::parse("AMD"), Some(GpuVendor::Amd));
assert_eq!(GpuVendor::parse("rocm"), Some(GpuVendor::Amd));
assert_eq!(GpuVendor::parse("hip"), Some(GpuVendor::Amd));
assert_eq!(GpuVendor::parse("intel"), None);
for v in [GpuVendor::Nvidia, GpuVendor::Amd] {
assert_eq!(GpuVendor::parse(v.as_str()), Some(v));
}
}
#[test]
fn variant_labels_classify_three_ways() {
for (label, class) in [
("cpu", VariantClass::Cpu),
("precompiled/cpu", VariantClass::Cpu),
("cpu-static", VariantClass::Cpu),
("precompiled/cu128", VariantClass::Vendor(GpuVendor::Nvidia)),
("builds/sm61-sm120", VariantClass::Vendor(GpuVendor::Nvidia)),
("precompiled/rocm70", VariantClass::Vendor(GpuVendor::Amd)),
("builds/gfx1030", VariantClass::Vendor(GpuVendor::Amd)),
("builds/gfx", VariantClass::Unknown),
("builds/mybuild", VariantClass::Unknown),
("", VariantClass::Unknown),
] {
assert_eq!(classify_variant_label(label), class, "{label:?}");
}
}
#[test]
fn vendor_picks_its_cargo_feature() {
assert_eq!(GpuVendor::Nvidia.cargo_feature(), "cuda");
assert_eq!(GpuVendor::Amd.cargo_feature(), "rocm");
}
#[test]
fn parses_nvidia_capability_forms() {
let expect = GpuArch::Sm { major: 8, minor: 6 };
for form in ["8.6", "sm_86", "sm86", " 8.6 "] {
assert_eq!(
GpuArch::parse(GpuVendor::Nvidia, form).unwrap(),
expect,
"{form}"
);
}
}
#[test]
fn concatenated_form_takes_the_last_digit_as_minor() {
assert_eq!(
GpuArch::parse(GpuVendor::Nvidia, "sm_120").unwrap(),
GpuArch::Sm {
major: 12,
minor: 0
}
);
assert_eq!(
GpuArch::parse(GpuVendor::Nvidia, "sm_61").unwrap(),
GpuArch::Sm { major: 6, minor: 1 }
);
}
#[test]
fn rejects_malformed_capability() {
for bad in ["", "sm_", "x.y", "sm_1", "notacap", "8."] {
assert!(GpuArch::parse(GpuVendor::Nvidia, bad).is_none(), "{bad:?}");
}
}
#[test]
fn strips_the_gfx_feature_suffix() {
assert_eq!(
GpuArch::parse(GpuVendor::Amd, "gfx906:sramecc-:xnack-").unwrap(),
GpuArch::Gfx("gfx906".into())
);
assert_eq!(
GpuArch::parse(GpuVendor::Amd, "GFX1030").unwrap(),
GpuArch::Gfx("gfx1030".into())
);
}
#[test]
fn rejects_malformed_gfx() {
for bad in ["", "gfx", "1030", "radeon"] {
assert!(GpuArch::parse(GpuVendor::Amd, bad).is_none(), "{bad:?}");
}
}
#[test]
fn displays_in_vendor_form() {
assert_eq!(
GpuArch::Sm {
major: 12,
minor: 0
}
.to_string(),
"sm_120"
);
assert_eq!(GpuArch::Gfx("gfx1100".into()).to_string(), "gfx1100");
}
#[test]
fn archs_token_differs_from_display_on_nvidia_only() {
let sm = GpuArch::Sm {
major: 12,
minor: 0,
};
assert_eq!(sm.archs_token(), "12.0");
assert_ne!(sm.archs_token(), sm.to_string());
let gfx = GpuArch::Gfx("gfx1100".into());
assert_eq!(gfx.archs_token(), gfx.to_string());
}
#[test]
fn an_arch_is_covered_by_a_list_of_its_own_tokens() {
for a in [
GpuArch::Sm { major: 6, minor: 1 },
GpuArch::Sm {
major: 12,
minor: 0,
},
GpuArch::Gfx("gfx1030".into()),
] {
assert!(a.covered_by(&a.archs_token()), "{a}");
assert!(
a.covered_by(&format!("gfx900;{};8.9", a.archs_token())),
"{a}"
);
}
}
#[test]
fn nvidia_coverage_falls_back_to_major() {
let sm86 = GpuArch::Sm { major: 8, minor: 6 };
assert!(sm86.covered_by("6.1;8.6"));
assert!(
sm86.covered_by("8.0"),
"same major is forward-compatible via PTX"
);
assert!(!sm86.covered_by("6.1;12.0"));
}
#[test]
fn nvidia_coverage_does_not_match_a_digit_inside_another_token() {
let cu128 = "7.0 7.5 8.0 8.6 8.9 9.0 12.0";
for (maj, min) in [(5u32, 0u32), (5, 2)] {
let dev = GpuArch::Sm {
major: maj,
minor: min,
};
assert!(
!dev.covered_by(cu128),
"sm_{maj}{min} matched the 5 inside 7.5",
);
}
assert!(!GpuArch::Sm { major: 2, minor: 0 }.covered_by("12.0"));
assert!(
GpuArch::Sm {
major: 12,
minor: 0
}
.covered_by(cu128)
);
assert!(GpuArch::Sm { major: 6, minor: 1 }.covered_by("5.0 5.2 6.0 6.1 7.0"));
assert!(!GpuArch::Sm { major: 6, minor: 1 }.covered_by(cu128));
}
#[test]
fn a_ptx_suffixed_arch_list_entry_still_matches() {
let sm86 = GpuArch::Sm { major: 8, minor: 6 };
assert!(sm86.covered_by("6.1;8.6+PTX"));
assert_eq!(GpuArch::parse(GpuVendor::Nvidia, "8.6+PTX").unwrap(), sm86,);
}
#[test]
fn a_cpu_variant_covers_nothing() {
assert!(!GpuArch::Sm { major: 8, minor: 6 }.covered_by("cpu"));
assert!(!GpuArch::Gfx("gfx1030".into()).covered_by("cpu"));
}
#[test]
fn amd_coverage_is_exact_only() {
let gfx1030 = GpuArch::Gfx("gfx1030".into());
assert!(gfx1030.covered_by("gfx900;gfx1030;gfx1100"));
assert!(gfx1030.covered_by("gfx1030"));
assert!(!gfx1030.covered_by("gfx1031"));
assert!(!gfx1030.covered_by("gfx10"));
assert!(!gfx1030.covered_by("gfx1100"));
}
#[test]
fn arch_knows_its_own_vendor() {
assert_eq!(
GpuArch::Sm { major: 8, minor: 6 }.vendor(),
GpuVendor::Nvidia
);
assert_eq!(GpuArch::Gfx("gfx942".into()).vendor(), GpuVendor::Amd);
assert_eq!(GpuArch::Gfx("gfx942".into()).sm_major(), None);
}
}