use std::fmt;
use std::sync::OnceLock;
#[allow(non_camel_case_types)]
#[derive(Copy, Clone, PartialEq, Eq, Hash, Debug)]
pub enum Arch {
Arm,
Aarch64,
X86_64,
RiscV64,
Wasm32Simd128,
}
impl Arch {
pub const ALL: [Arch; 5] =
[Arch::Arm, Arch::Aarch64, Arch::X86_64, Arch::RiscV64, Arch::Wasm32Simd128];
pub fn is_native(&self) -> bool {
match self {
Arch::Arm => cfg!(target_arch = "arm"),
Arch::Aarch64 => cfg!(target_arch = "aarch64"),
Arch::X86_64 => cfg!(target_arch = "x86_64"),
Arch::RiscV64 => cfg!(target_arch = "riscv64"),
Arch::Wasm32Simd128 => {
cfg!(all(target_arch = "wasm32", target_feature = "simd128"))
}
}
}
}
impl fmt::Display for Arch {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str(match self {
Arch::Arm => "arm",
Arch::Aarch64 => "aarch64",
Arch::X86_64 => "x86_64",
Arch::RiscV64 => "riscv64",
Arch::Wasm32Simd128 => "wasm32+simd128",
})
}
}
#[derive(Copy, Clone, PartialEq, Eq, Hash, Debug)]
pub enum Isa {
Arm,
Aarch64,
X86_64,
RiscV64,
Wasm32,
ArmNeon,
Aarch64Fp16,
Aarch64DotProd,
Aarch64Sve2,
Aarch64Sme,
Aarch64Sme2,
Aarch64AppleAmx,
X86_64Avx,
X86_64Avx2,
X86_64Fma,
X86_64F16c,
X86_64Avx512f,
X86_64Avx512Vnni,
X86_64Avx512Fp16,
X86_64AvxVnni,
X86_64AmxInt8,
X86_64AmxBf16,
RiscV64V,
RiscV64Vlen256,
RiscV64Zvfh,
Wasm32Simd128,
Wasm32RelaxedSimd,
}
impl Isa {
pub const ALL: [Isa; 27] = [
Isa::Arm,
Isa::Aarch64,
Isa::X86_64,
Isa::RiscV64,
Isa::Wasm32,
Isa::ArmNeon,
Isa::Aarch64Fp16,
Isa::Aarch64DotProd,
Isa::Aarch64Sve2,
Isa::Aarch64Sme,
Isa::Aarch64Sme2,
Isa::Aarch64AppleAmx,
Isa::X86_64Avx,
Isa::X86_64Avx2,
Isa::X86_64Fma,
Isa::X86_64F16c,
Isa::X86_64Avx512f,
Isa::X86_64Avx512Vnni,
Isa::X86_64Avx512Fp16,
Isa::X86_64AvxVnni,
Isa::X86_64AmxInt8,
Isa::X86_64AmxBf16,
Isa::RiscV64V,
Isa::RiscV64Vlen256,
Isa::RiscV64Zvfh,
Isa::Wasm32Simd128,
Isa::Wasm32RelaxedSimd,
];
pub fn name(&self) -> &'static str {
match self {
Isa::Arm => "arm",
Isa::Aarch64 => "aarch64",
Isa::X86_64 => "x86_64",
Isa::RiscV64 => "riscv64",
Isa::Wasm32 => "wasm32",
Isa::ArmNeon => "neon",
Isa::Aarch64Fp16 => "fp16",
Isa::Aarch64DotProd => "dotprod",
Isa::Aarch64Sve2 => "sve2",
Isa::Aarch64Sme => "sme",
Isa::Aarch64Sme2 => "sme2",
Isa::Aarch64AppleAmx => "apple-amx",
Isa::X86_64Avx => "avx",
Isa::X86_64Avx2 => "avx2",
Isa::X86_64Fma => "fma",
Isa::X86_64F16c => "f16c",
Isa::X86_64Avx512f => "avx512f",
Isa::X86_64Avx512Vnni => "avx512vnni",
Isa::X86_64Avx512Fp16 => "avx512fp16",
Isa::X86_64AvxVnni => "avxvnni",
Isa::X86_64AmxInt8 => "amx-int8",
Isa::X86_64AmxBf16 => "amx-bf16",
Isa::RiscV64V => "rvv",
Isa::RiscV64Vlen256 => "vlen256",
Isa::RiscV64Zvfh => "zvfh",
Isa::Wasm32Simd128 => "simd128",
Isa::Wasm32RelaxedSimd => "relaxed-simd",
}
}
fn from_name(s: &str) -> Option<Isa> {
Isa::ALL.into_iter().find(|i| i.name() == s)
}
pub const fn level(&self) -> u8 {
match self {
Isa::Arm | Isa::Aarch64 | Isa::X86_64 | Isa::RiscV64 | Isa::Wasm32 => 0,
Isa::X86_64Avx => 1,
Isa::X86_64Avx2 | Isa::X86_64Fma | Isa::X86_64F16c => 2,
Isa::X86_64Avx512f | Isa::X86_64AvxVnni => 3,
Isa::X86_64Avx512Vnni => 4,
Isa::X86_64Avx512Fp16 => 4,
Isa::X86_64AmxInt8 | Isa::X86_64AmxBf16 => 5,
Isa::RiscV64V => 1,
Isa::RiscV64Vlen256 => 2,
Isa::RiscV64Zvfh => 3,
Isa::ArmNeon => 1,
Isa::Aarch64Fp16 | Isa::Aarch64DotProd => 2,
Isa::Aarch64Sve2 => 3,
Isa::Aarch64Sme | Isa::Aarch64Sme2 => 4,
Isa::Aarch64AppleAmx => 5,
Isa::Wasm32Simd128 => 0,
Isa::Wasm32RelaxedSimd => 1,
}
}
pub const fn fp16_arithmetic(&self) -> bool {
matches!(self, Isa::Aarch64Fp16 | Isa::X86_64Avx512Fp16 | Isa::RiscV64Zvfh)
}
pub const fn arch(&self) -> Arch {
match self {
Isa::Arm | Isa::ArmNeon => Arch::Arm,
Isa::Aarch64
| Isa::Aarch64Fp16
| Isa::Aarch64DotProd
| Isa::Aarch64Sve2
| Isa::Aarch64Sme
| Isa::Aarch64Sme2
| Isa::Aarch64AppleAmx => Arch::Aarch64,
Isa::X86_64
| Isa::X86_64Avx
| Isa::X86_64Avx2
| Isa::X86_64Fma
| Isa::X86_64F16c
| Isa::X86_64Avx512f
| Isa::X86_64Avx512Vnni
| Isa::X86_64Avx512Fp16
| Isa::X86_64AvxVnni
| Isa::X86_64AmxInt8
| Isa::X86_64AmxBf16 => Arch::X86_64,
Isa::RiscV64 | Isa::RiscV64V | Isa::RiscV64Vlen256 | Isa::RiscV64Zvfh => Arch::RiscV64,
Isa::Wasm32 | Isa::Wasm32Simd128 | Isa::Wasm32RelaxedSimd => Arch::Wasm32Simd128,
}
}
pub const fn is_arch(&self) -> bool {
matches!(self, Isa::Arm | Isa::Aarch64 | Isa::X86_64 | Isa::RiscV64 | Isa::Wasm32)
}
pub const fn of_arch(arch: Arch) -> Isa {
match arch {
Arch::Arm => Isa::Arm,
Arch::Aarch64 => Isa::Aarch64,
Arch::X86_64 => Isa::X86_64,
Arch::RiscV64 => Isa::RiscV64,
Arch::Wasm32Simd128 => Isa::Wasm32,
}
}
}
impl fmt::Display for Isa {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str(self.name())
}
}
#[derive(Copy, Clone, Default, PartialEq, Eq)]
pub struct IsaSet(u32);
impl IsaSet {
pub const fn empty() -> IsaSet {
IsaSet(0)
}
pub const fn of_arch(arch: Arch) -> IsaSet {
IsaSet::empty().with(Isa::of_arch(arch))
}
pub fn arch(self) -> Option<Arch> {
self.iter().find(|i| i.is_arch()).map(|i| i.arch())
}
pub const fn with(self, isa: Isa) -> IsaSet {
IsaSet(self.0 | 1 << isa as u32)
}
pub const fn without(self, isa: Isa) -> IsaSet {
IsaSet(self.0 & !(1 << isa as u32))
}
pub const fn has(self, isa: Isa) -> bool {
self.0 & (1 << isa as u32) != 0
}
pub fn fp16_arithmetic(self) -> bool {
self.iter().any(|i| i.fp16_arithmetic())
}
pub fn nickname(self) -> &'static str {
let Some(arch) = self.arch() else { return "none" };
match (arch, self.level()) {
(Arch::Arm, 0) => "vfp",
(Arch::Arm, _) => "neon",
(Arch::Aarch64, 0 | 1) => "neon",
(Arch::Aarch64, 2) => "fp16",
(Arch::Aarch64, 3) => "sve2",
(Arch::Aarch64, 4) => "sme",
(Arch::Aarch64, _) => "amx",
(Arch::X86_64, 0) => "sse2",
(Arch::X86_64, 1) => "avx",
(Arch::X86_64, 2) => "fma",
(Arch::X86_64, 3) => "avx512",
(Arch::X86_64, 4) => "fp16",
(Arch::X86_64, _) => "amx",
(Arch::RiscV64, 0) => "rv64",
(Arch::RiscV64, 1) => "rvv",
(Arch::RiscV64, 2) => "vlen256",
(Arch::RiscV64, _) => "zvfh",
(Arch::Wasm32Simd128, 0) => "simd",
(Arch::Wasm32Simd128, _) => "relaxed",
}
}
pub fn level(self) -> u8 {
self.iter().map(|i| i.level()).max().unwrap_or(0)
}
pub fn ladder(arch: Arch, level: u8) -> IsaSet {
let mut set = IsaSet::of_arch(arch);
for isa in Isa::ALL {
if isa.arch() == arch && isa.level() <= level {
set = set.with(isa);
}
}
set
}
pub fn every_ladder() -> impl Iterator<Item = IsaSet> {
Arch::ALL.into_iter().flat_map(|arch| (0..=MAX_LEVEL).map(move |l| IsaSet::ladder(arch, l)))
}
pub fn iter(self) -> impl Iterator<Item = Isa> {
Isa::ALL.into_iter().filter(move |i| self.has(*i))
}
}
impl fmt::Debug for IsaSet {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
if self.0 == 0 {
return f.write_str("-");
}
f.write_str(&self.iter().map(|i| i.name()).collect::<Vec<_>>().join(","))
}
}
#[derive(Copy, Clone, PartialEq, Eq, Hash)]
pub struct IsaReq {
pub needs: &'static [Isa],
}
impl IsaReq {
pub const ANY: IsaReq = IsaReq { needs: &[] };
pub const fn needing(self, needs: &'static [Isa]) -> IsaReq {
IsaReq { needs }
}
pub fn satisfied_by(&self, set: IsaSet) -> bool {
self.needs.iter().all(|i| set.has(*i))
}
pub fn level(&self) -> u8 {
self.needs.iter().map(|i| i.level()).max().unwrap_or(0)
}
}
pub const LEVEL_BOOST: isize = 10;
pub const MAX_LEVEL: u8 = 5;
pub const fn peer_of(mine: Isa, theirs: Isa) -> isize {
(theirs.level() as isize - mine.level() as isize) * LEVEL_BOOST
}
pub const NEVER_PREFERRED: isize = -(LEVEL_BOOST * MAX_LEVEL as isize) - 1;
impl fmt::Debug for IsaReq {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let needs = self.needs.iter().map(|i| i.name()).collect::<Vec<_>>().join("+");
if needs.is_empty() {
f.write_str("any")?;
} else {
f.write_str(&needs)?;
}
Ok(())
}
}
pub fn native() -> IsaSet {
static NATIVE: OnceLock<IsaSet> = OnceLock::new();
*NATIVE.get_or_init(|| {
let set = forced(probe());
log::debug!("ISA: {set:?}");
set
})
}
fn probe() -> IsaSet {
#[cfg(target_arch = "arm")]
return crate::arm32::isa_set();
#[cfg(target_arch = "aarch64")]
return crate::arm64::isa_set();
#[cfg(target_arch = "x86_64")]
return crate::x86_64::isa_set();
#[cfg(target_arch = "riscv64")]
return crate::riscv64::isa_set();
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
return crate::wasm::isa_set();
#[cfg(not(any(
target_arch = "arm",
target_arch = "aarch64",
target_arch = "x86_64",
target_arch = "riscv64",
all(target_arch = "wasm32", target_feature = "simd128")
)))]
IsaSet::empty()
}
impl std::str::FromStr for IsaSet {
type Err = tract_data::internal::TractError;
fn from_str(spec: &str) -> tract_data::internal::TractResult<IsaSet> {
let mut set = IsaSet::empty();
for token in spec.split(',').map(str::trim).filter(|t| !t.is_empty()) {
let Some(isa) = Isa::from_name(token) else {
tract_data::internal::bail!("{token:?} is no instruction set tract knows")
};
if let Some(arch) = set.arch()
&& isa.arch() != arch
{
tract_data::internal::bail!(
"{token} belongs to {}, and this machine is {arch}: a machine is one \
architecture",
isa.arch()
)
}
set = set.with(Isa::of_arch(isa.arch())).with(isa);
}
if set == IsaSet::empty() {
tract_data::internal::bail!("{spec:?} names no instruction set")
}
Ok(set)
}
}
pub(crate) fn forced(mut set: IsaSet) -> IsaSet {
let Some(spec) = crate::knobs::TRACT_CPU_ISA.get() else { return set };
for token in spec.split(',').map(str::trim).filter(|t| !t.is_empty()) {
let (add, name) = match token.split_at(1) {
("+", name) => (true, name),
("-", name) => (false, name),
_ => (true, token),
};
let Some(isa) = Isa::from_name(name) else {
log::warn!("TRACT_CPU_ISA: unknown feature {name:?}, ignored");
continue;
};
if let Some(arch) = set.arch() {
assert!(
isa.arch() == arch,
"TRACT_CPU_ISA: {name} belongs to {}, and this set is {arch} — a machine is one \
architecture, so the token cannot apply",
isa.arch()
);
}
set = if add { set.with(isa) } else { set.without(isa) };
}
set
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ladder_stays_within_the_bound() {
for isa in Isa::ALL {
assert!(
isa.level() <= MAX_LEVEL,
"{isa} is level {}, past MAX_LEVEL {MAX_LEVEL}",
isa.level()
);
}
}
#[test]
fn peer_of_cancels_the_steps_between() {
assert_eq!(peer_of(Isa::X86_64Fma, Isa::X86_64Avx512f), LEVEL_BOOST);
assert_eq!(peer_of(Isa::X86_64Avx, Isa::X86_64Avx512Vnni), 3 * LEVEL_BOOST);
assert_eq!(peer_of(Isa::X86_64Avx512f, Isa::X86_64AvxVnni), 0);
}
}