Add the MIGraphX rung for AMD GPUs
Benchmarks / CPU and I/O (per commit) (push) Failing after 6m20s
Benchmarks / Frame budget (on demand) (push) Skipped
Build and test / Desktop (Linux) (push) Failing after 45s
Build and test / Layer separation (push) Successful in 26s
Traceability / Requirement traces (push) Failing after 46s
🐳 Android image / Build and push (push) Successful in 1s
Build and test / android-image (push) Successful in 1s
🐳 Windows image / Build and push (push) Successful in 1s
Build and test / windows-image (push) Successful in 1s
Build and test / Android (aarch64) (push) Failing after 2m19s
Build and test / Windows (x86_64, cross) (push) Failing after 3m2s
Benchmarks / CPU and I/O (per commit) (push) Failing after 6m20s
Benchmarks / Frame budget (on demand) (push) Skipped
Build and test / Desktop (Linux) (push) Failing after 45s
Build and test / Layer separation (push) Successful in 26s
Traceability / Requirement traces (push) Failing after 46s
🐳 Android image / Build and push (push) Successful in 1s
Build and test / android-image (push) Successful in 1s
🐳 Windows image / Build and push (push) Successful in 1s
Build and test / windows-image (push) Successful in 1s
Build and test / Android (aarch64) (push) Failing after 2m19s
Build and test / Windows (x86_64, cross) (push) Failing after 3m2s
Measured on a Radeon RX 7900 XT against Arch's onnxruntime-rocm 1.29 (docs/inference.md §1.3): MIGraphX fp16 runs the detectors at 2.4–3.4 ms against 10–58 ms on the CPU provider, the inpainter at 8 ms against 514, with a 15–135 s compile per graph the first time and under a second from its cache after. A compiling rung on TensorRT's terms, wired the same way. The ROCm execution provider is gone (removed in ONNX Runtime 1.23), so the AMD ladder is MIGraphX then the CPU, with no non-compiling rung between. MIGraphX is registered through the runtime's generic key/value entry point rather than ort's builder: 1.29 reads the legacy options struct for its precision flags only, and the compiled-program cache directory (`migraphx_model_cache_dir`) only travels the generic way. The provider's cache key omits the precision, so f32 and fp16 programs get their own directories. The probe fingerprint now includes the provider libraries beside the runtime and the ROCm version, since a distribution's CPU and ROCm builds are the same file at the same path. `status().failed` reports only the rungs above the selection, so an AMD desktop's About line says why MIGraphX won rather than that the NVIDIA providers are not in the build. Two examples: `ep_probe` times each provider cold and from cache, and `ladder` drives `init` as the app does to watch the first-run sequence. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -3,8 +3,8 @@
|
||||
//!
|
||||
//! Consumers ask for a session by [`Role`] and get `ort`'s `Session` back;
|
||||
//! what built it — tract on one core, ONNX Runtime's CPU pool, a TensorRT
|
||||
//! engine, the Hexagon — is this crate's business and shows up in
|
||||
//! [`status`] for the settings row and nowhere else.
|
||||
//! engine, a MIGraphX program, the Hexagon — is this crate's business and
|
||||
//! shows up in [`status`] for the settings row and nowhere else.
|
||||
//!
|
||||
//! The shape follows §3 of the spec: `ort` links nothing (`alternative-backend`),
|
||||
//! and the first call hands it an API table from either a `libonnxruntime`
|
||||
@@ -60,7 +60,9 @@ pub enum Form {
|
||||
|
||||
/// A rung of the ladder (§2). Ordered: a user override names the highest rung
|
||||
/// the probe may take, and a compiling rung falls back to the one below it
|
||||
/// until its engine exists.
|
||||
/// until its engine exists. The order is within a vendor's ladder — a
|
||||
/// machine has NVIDIA rungs or an AMD rung, never both — so a ceiling is
|
||||
/// read as "no higher than this on whichever ladder the device has".
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
|
||||
pub enum Rung {
|
||||
/// ONNX Runtime's CPU provider, or tract when no runtime file was found.
|
||||
@@ -69,6 +71,11 @@ pub enum Rung {
|
||||
Cuda,
|
||||
/// NVIDIA, through a TensorRT engine compiled on this device. Desktop only.
|
||||
TensorRt,
|
||||
/// AMD, through a MIGraphX program compiled on this device. Desktop
|
||||
/// only. ONNX Runtime's ROCm provider, the CUDA provider's twin, was
|
||||
/// removed in ONNX Runtime 1.23, so there is no non-compiling AMD rung
|
||||
/// to fall back to: this one falls back to the CPU.
|
||||
MiGraphX,
|
||||
/// Qualcomm's Hexagon NPU through QNN, int8 models only. Android only.
|
||||
Hexagon,
|
||||
}
|
||||
@@ -79,6 +86,7 @@ impl Rung {
|
||||
Rung::Cpu => "CPU",
|
||||
Rung::Cuda => "CUDA",
|
||||
Rung::TensorRt => "TensorRT",
|
||||
Rung::MiGraphX => "MIGraphX",
|
||||
Rung::Hexagon => "Hexagon NPU",
|
||||
}
|
||||
}
|
||||
@@ -88,13 +96,13 @@ impl Rung {
|
||||
fn fallback(self) -> Rung {
|
||||
match self {
|
||||
Rung::TensorRt => Rung::Cuda,
|
||||
Rung::Hexagon | Rung::Cuda | Rung::Cpu => Rung::Cpu,
|
||||
Rung::MiGraphX | Rung::Hexagon | Rung::Cuda | Rung::Cpu => Rung::Cpu,
|
||||
}
|
||||
}
|
||||
|
||||
/// Whether a session on this rung needs an engine built first.
|
||||
fn compiles(self) -> bool {
|
||||
matches!(self, Rung::TensorRt | Rung::Hexagon)
|
||||
matches!(self, Rung::TensorRt | Rung::MiGraphX | Rung::Hexagon)
|
||||
}
|
||||
|
||||
/// The model form this rung wants for a role.
|
||||
@@ -164,7 +172,7 @@ impl Status {
|
||||
pub fn line(&self) -> String {
|
||||
let form = match self.rung {
|
||||
Rung::Hexagon => " · int8",
|
||||
Rung::TensorRt => " · fp16",
|
||||
Rung::TensorRt | Rung::MiGraphX => " · fp16",
|
||||
_ => "",
|
||||
};
|
||||
format!("{}{} · {}", self.rung.label(), form, self.runtime.label())
|
||||
@@ -269,8 +277,8 @@ fn acquire(role: Role, form: Form, bytes: &Arc<[u8]>, hash: u64) -> Result<Acqui
|
||||
return Ok(Acquired { entry });
|
||||
}
|
||||
|
||||
// Built outside the registry lock: a TensorRT engine load is long enough
|
||||
// that another role's acquire should not wait on it.
|
||||
// Built outside the registry lock: a TensorRT or MIGraphX engine load
|
||||
// is long enough that another role's acquire should not wait on it.
|
||||
let session = session::build(rung, role, bytes, &cfg)?;
|
||||
log::debug!("inference: {role:?} loaded on {}", rung.label());
|
||||
let entry = Arc::new(Loaded {
|
||||
@@ -401,7 +409,16 @@ pub fn status() -> Status {
|
||||
runtime: api::runtime(),
|
||||
rung,
|
||||
reason: s.cache.reason.clone(),
|
||||
failed: s.cache.failed.clone(),
|
||||
// Only what explains the selection: on an AMD machine the NVIDIA
|
||||
// rungs "not enabled in this build" say nothing about why MIGraphX
|
||||
// was taken. With the floor selected, everything tried is above it.
|
||||
failed: s
|
||||
.cache
|
||||
.failed
|
||||
.iter()
|
||||
.filter(|(r, _)| *r > rung)
|
||||
.cloned()
|
||||
.collect(),
|
||||
probing: s.probing,
|
||||
engines: if rung.compiles() {
|
||||
(s.cache.compiled.len(), s.wanted)
|
||||
@@ -614,8 +631,33 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_status_reports_only_the_rungs_above_the_selection() {
|
||||
let _serial = serial();
|
||||
let failed = vec![
|
||||
(Rung::TensorRt, "not enabled".to_string()),
|
||||
(Rung::Cuda, "not enabled".to_string()),
|
||||
];
|
||||
let before = state().lock().unwrap().cache.clone();
|
||||
state().lock().unwrap().cache = Cache {
|
||||
rung: Some(Rung::MiGraphX),
|
||||
failed: failed.clone(),
|
||||
..Cache::default()
|
||||
};
|
||||
// An AMD desktop: the NVIDIA rungs below MIGraphX are not the story.
|
||||
assert!(status().failed.is_empty());
|
||||
// An NVIDIA desktop on the CUDA provider: TensorRT's failure is.
|
||||
state().lock().unwrap().cache.rung = Some(Rung::Cuda);
|
||||
assert_eq!(status().failed, vec![failed[0].clone()]);
|
||||
// The floor: everything tried explains it.
|
||||
state().lock().unwrap().cache.rung = Some(Rung::Cpu);
|
||||
assert_eq!(status().failed.len(), 2);
|
||||
state().lock().unwrap().cache = before;
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_status_line_reads_as_the_floor_before_init() {
|
||||
let _serial = serial();
|
||||
let s = status();
|
||||
assert_eq!(s.rung, Rung::Cpu);
|
||||
assert!(s.line().starts_with("CPU"), "{}", s.line());
|
||||
|
||||
@@ -16,8 +16,11 @@ use crate::{api::Runtime, state, Cache, Config, Form, Role, Rung};
|
||||
fn ladder(ceiling: Option<Rung>) -> Vec<Rung> {
|
||||
#[cfg(target_os = "android")]
|
||||
let all = [Rung::Hexagon];
|
||||
// A desktop has one vendor's GPU; the other vendor's providers are
|
||||
// "not enabled in this build" or a library that fails to load, and
|
||||
// either answer arrives in milliseconds.
|
||||
#[cfg(not(target_os = "android"))]
|
||||
let all = [Rung::TensorRt, Rung::Cuda];
|
||||
let all = [Rung::TensorRt, Rung::Cuda, Rung::MiGraphX];
|
||||
all.into_iter()
|
||||
.filter(|r| ceiling.is_none_or(|c| *r <= c))
|
||||
.collect()
|
||||
@@ -206,15 +209,22 @@ fn first_line(s: &str) -> String {
|
||||
line[start..].chars().take(200).collect()
|
||||
}
|
||||
|
||||
/// Everything a change of which should re-probe: the runtime and where it
|
||||
/// came from, this crate, the platform, the driver or SoC, and the models.
|
||||
/// Everything a change of which should re-probe: the runtime, where it
|
||||
/// came from and which providers sit beside it, this crate, the platform,
|
||||
/// the driver or SoC, and the models.
|
||||
fn fingerprint(runtime: &Runtime, cfg: &Config) -> String {
|
||||
let mut parts = vec![
|
||||
format!("engine {}", env!("CARGO_PKG_VERSION")),
|
||||
format!("{} {}", std::env::consts::OS, std::env::consts::ARCH),
|
||||
match runtime {
|
||||
Runtime::Tract => "tract".to_string(),
|
||||
Runtime::OnnxRuntime { path, version } => format!("ort {version} {}", path.display()),
|
||||
Runtime::OnnxRuntime { path, version } => {
|
||||
format!(
|
||||
"ort {version} {} [{}]",
|
||||
path.display(),
|
||||
providers_beside(path)
|
||||
)
|
||||
}
|
||||
},
|
||||
device_identity(),
|
||||
];
|
||||
@@ -237,13 +247,41 @@ fn fingerprint(runtime: &Runtime, cfg: &Config) -> String {
|
||||
parts.join("\n")
|
||||
}
|
||||
|
||||
/// The `libonnxruntime_providers_*.so` files in the runtime's directory.
|
||||
/// A distribution's CPU-only and ROCm builds are the same version at the
|
||||
/// same path; the provider libraries beside them are what differs.
|
||||
fn providers_beside(runtime: &Path) -> String {
|
||||
let Some(dir) = runtime.parent() else {
|
||||
return String::new();
|
||||
};
|
||||
let mut names: Vec<String> = std::fs::read_dir(dir)
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter_map(|e| e.ok())
|
||||
.filter_map(|e| e.file_name().into_string().ok())
|
||||
.filter(|n| {
|
||||
n.starts_with("libonnxruntime_providers_") || n.starts_with("onnxruntime_providers_")
|
||||
})
|
||||
.collect();
|
||||
names.sort();
|
||||
names.join(" ")
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn device_identity() -> String {
|
||||
// The NVIDIA driver's version line; absent means no NVIDIA driver.
|
||||
std::fs::read_to_string("/proc/driver/nvidia/version")
|
||||
// The NVIDIA driver's version line, or the ROCm release the AMD stack
|
||||
// came from (`rocm-core` writes it; the kernel driver has no version
|
||||
// of its own). Absent means neither.
|
||||
if let Some(line) = std::fs::read_to_string("/proc/driver/nvidia/version")
|
||||
.ok()
|
||||
.and_then(|s| s.lines().next().map(str::to_string))
|
||||
.unwrap_or_else(|| "no nvidia driver".into())
|
||||
{
|
||||
return line;
|
||||
}
|
||||
if let Ok(rocm) = std::fs::read_to_string("/opt/rocm/.info/version") {
|
||||
return format!("rocm {}", rocm.trim());
|
||||
}
|
||||
"no nvidia driver, no rocm".into()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "android")]
|
||||
|
||||
@@ -83,10 +83,65 @@ fn providers(
|
||||
ep::CUDA::default().build(),
|
||||
])?)
|
||||
}
|
||||
Rung::MiGraphX => {
|
||||
// fp16 on the same terms as TensorRT (§7). MIGraphX compiles a
|
||||
// program per graph — 20–60 s here — and keeps it in the cache
|
||||
// directory, keyed on the graph, the GPU and its own version
|
||||
// but not the precision: hence one directory per precision.
|
||||
// The CPU takes any node it declines.
|
||||
let fp16 = role != Role::Embedder;
|
||||
let cache = cfg
|
||||
.cache_dir
|
||||
.join("migraphx")
|
||||
.join(if fp16 { "fp16" } else { "f32" });
|
||||
let _ = std::fs::create_dir_all(&cache);
|
||||
let mut b = b;
|
||||
migraphx(&mut b, fp16, &cache)?;
|
||||
Ok(b)
|
||||
}
|
||||
Rung::Hexagon => unreachable!("the Hexagon rung is not on a desktop ladder"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Register MIGraphX through ONNX Runtime's generic key/value entry point.
|
||||
///
|
||||
/// `ort`'s own builder (`ep::MIGraphX`) fills the legacy
|
||||
/// `OrtMIGraphXProviderOptions`, and 1.29 reads that struct for its
|
||||
/// precision flags and nothing else — the compiled-program cache directory
|
||||
/// is only a key in the generic map (`migraphx_model_cache_dir`), and
|
||||
/// without it every session is a full compile. Registration through the
|
||||
/// generic entry point needs no `ort` feature: it is one call on the API
|
||||
/// table, which is why the crate's `ort` dependency names no AMD feature.
|
||||
#[cfg(not(target_os = "android"))]
|
||||
fn migraphx(
|
||||
b: &mut ort::session::builder::SessionBuilder,
|
||||
fp16: bool,
|
||||
cache: &std::path::Path,
|
||||
) -> ort::Result<()> {
|
||||
use ort::AsPointer;
|
||||
use std::ffi::CString;
|
||||
let keys = [c"migraphx_fp16_enable", c"migraphx_model_cache_dir"];
|
||||
let values = [
|
||||
CString::new(if fp16 { "1" } else { "0" }).unwrap(),
|
||||
CString::new(cache.to_string_lossy().as_bytes())
|
||||
.map_err(|e| ort::Error::new(e.to_string()))?,
|
||||
];
|
||||
let key_ptrs: Vec<_> = keys.iter().map(|k| k.as_ptr()).collect();
|
||||
let value_ptrs: Vec<_> = values.iter().map(|v| v.as_ptr()).collect();
|
||||
// SAFETY: the documented C call over arrays that outlive it; the
|
||||
// runtime copies the strings into its own options map before returning.
|
||||
unsafe {
|
||||
let status = (ort::api().SessionOptionsAppendExecutionProvider)(
|
||||
b.ptr_mut(),
|
||||
c"MIGraphX".as_ptr(),
|
||||
key_ptrs.as_ptr(),
|
||||
value_ptrs.as_ptr(),
|
||||
keys.len(),
|
||||
);
|
||||
ort::Error::result_from_status(status)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "android")]
|
||||
fn providers(
|
||||
b: ort::session::builder::SessionBuilder,
|
||||
@@ -120,6 +175,8 @@ fn providers(
|
||||
.build()
|
||||
.error_on_failure()])?)
|
||||
}
|
||||
Rung::Cuda | Rung::TensorRt => unreachable!("no NVIDIA rung on Android"),
|
||||
Rung::Cuda | Rung::TensorRt | Rung::MiGraphX => {
|
||||
unreachable!("no desktop GPU rung on Android")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user