Skip to main content

ironaccelerator_neuron/
backend.rs

1//! AWS Neuron `Backend` impl. One descriptor per NeuronCore.
2
3use ironaccelerator_core::{
4    Backend, BackendKind, Capability, CapabilityFlags, ComputeTier, DeviceDescriptor, DeviceId,
5    Result, Vendor,
6};
7
8use crate::drv::NeuronGen;
9
10pub struct NeuronBackend;
11pub static NEURON_BACKEND: NeuronBackend = NeuronBackend;
12
13impl Backend for NeuronBackend {
14    fn kind(&self) -> BackendKind {
15        BackendKind::Neuron
16    }
17
18    fn is_available(&self) -> bool {
19        crate::drv::is_available() && crate::drv::total_cores() > 0
20    }
21
22    fn enumerate(&self) -> Result<Vec<DeviceDescriptor>> {
23        let cores = crate::drv::total_cores();
24        if cores == 0 {
25            return Ok(Vec::new());
26        }
27        let gen = crate::drv::detect_generation();
28        let flags = capability_flags(gen);
29        let arch = arch_string(gen);
30        let hbm = hbm_bytes_per_core(gen);
31        Ok((0..cores)
32            .map(|ord| DeviceDescriptor {
33                id: DeviceId {
34                    backend: BackendKind::Neuron,
35                    ordinal: ord,
36                },
37                vendor: Vendor::Aws,
38                name: format!("AWS NeuronCore ({})", arch),
39                arch: arch.to_string(),
40                total_memory_bytes: hbm,
41                multiprocessor_count: 0,
42                clock_khz: 0,
43                capability: Capability {
44                    flags,
45                    tier: ComputeTier::Datacenter,
46                    fp16_tflops: None,
47                    fp8_tflops: None,
48                    mem_bandwidth_gbs: None,
49                },
50            })
51            .collect())
52    }
53
54    fn capabilities(&self, device: u32) -> Result<CapabilityFlags> {
55        // All NeuronCores on an instance are the same generation; validate the
56        // ordinal so this agrees with `enumerate` on what exists.
57        if !self.enumerate()?.iter().any(|d| d.id.ordinal == device) {
58            return Err(ironaccelerator_core::Error::InvalidArgument(
59                "neuron core ordinal out of range",
60            ));
61        }
62        Ok(capability_flags(crate::drv::detect_generation()))
63    }
64}
65
66fn capability_flags(gen: NeuronGen) -> CapabilityFlags {
67    let base = CapabilityFlags::FP32
68        | CapabilityFlags::BF16
69        | CapabilityFlags::INT8
70        | CapabilityFlags::TENSOR_CORES
71        | CapabilityFlags::HBM
72        | CapabilityFlags::MULTI_STREAM;
73    match gen {
74        NeuronGen::Trn2 => {
75            base | CapabilityFlags::FP8_E4M3 | CapabilityFlags::FP8_E5M2 | CapabilityFlags::INT4
76        }
77        NeuronGen::Trn1 => base | CapabilityFlags::FP8_E4M3 | CapabilityFlags::FP8_E5M2,
78        NeuronGen::Inf1 => {
79            CapabilityFlags::FP32
80                | CapabilityFlags::BF16
81                | CapabilityFlags::INT8
82                | CapabilityFlags::TENSOR_CORES
83                | CapabilityFlags::MULTI_STREAM
84        }
85        NeuronGen::Unknown => base,
86    }
87}
88
89fn arch_string(gen: NeuronGen) -> &'static str {
90    match gen {
91        NeuronGen::Inf1 => "neuron-v1-inferentia",
92        NeuronGen::Trn1 => "neuron-v2-trainium",
93        NeuronGen::Trn2 => "neuron-v3-trainium2",
94        NeuronGen::Unknown => "neuron",
95    }
96}
97
98fn hbm_bytes_per_core(gen: NeuronGen) -> u64 {
99    const GIB: u64 = 1024 * 1024 * 1024;
100    match gen {
101        NeuronGen::Inf1 => 8 * GIB,
102        NeuronGen::Trn1 => 16 * GIB,
103        NeuronGen::Trn2 => 24 * GIB,
104        NeuronGen::Unknown => 0,
105    }
106}