ironaccelerator_neuron/
backend.rs1use 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 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}