pub fn parse_device(name: &str) -> anyhow::Result<rlx::Device> {
Ok(match name.to_lowercase().as_str() {
"cpu" => rlx::Device::Cpu,
"metal" | "mps" => rlx::Device::Metal,
"mlx" => rlx::Device::Mlx,
"ane" | "coreml" | "neural-engine" => rlx::Device::Ane,
"cuda" => rlx::Device::Cuda,
"rocm" | "hip" => rlx::Device::Rocm,
"gpu" | "wgpu" => rlx::Device::Gpu,
"vulkan" => rlx::Device::Vulkan,
other => {
anyhow::bail!("unknown RLX device '{other}' (try: cpu, metal, mlx, cuda, wgpu, ane)")
}
})
}
pub fn default_rlx_device() -> rlx::Device {
for d in [
rlx::Device::Cuda,
rlx::Device::Metal,
rlx::Device::Mlx,
rlx::Device::Ane,
rlx::Device::Gpu,
rlx::Device::Cpu,
] {
if require_available(d) {
return d;
}
}
rlx::Device::Cpu
}
pub fn resolve_rlx_device(name: &str) -> anyhow::Result<rlx::Device> {
let device = if name.eq_ignore_ascii_case("auto") {
default_rlx_device()
} else {
parse_device(name)?
};
if !require_available(device) {
anyhow::bail!(
"RLX device '{}' is not available in this build (enabled: {:?})",
device_label(device),
available_device_labels()
);
}
Ok(device)
}
pub fn parse_engine(engine: &str) -> Option<rlx::Device> {
match engine.trim().to_ascii_lowercase().as_str() {
"rlx" | "rlx-cpu" => Some(rlx::Device::Cpu),
"rlx-metal" | "rlx-mps" => Some(rlx::Device::Metal),
"rlx-cuda" => Some(rlx::Device::Cuda),
"rlx-mlx" => Some(rlx::Device::Mlx),
"rlx-wgpu" | "rlx-gpu" => Some(rlx::Device::Gpu),
"rlx-ane" | "rlx-coreml" => Some(rlx::Device::Ane),
"rlx-rocm" => Some(rlx::Device::Rocm),
_ => None,
}
}
pub fn device_label(device: rlx::Device) -> &'static str {
match device {
rlx::Device::Cpu => "cpu",
rlx::Device::Metal => "metal",
rlx::Device::Mlx => "mlx",
rlx::Device::Ane => "ane",
rlx::Device::Cuda => "cuda",
rlx::Device::Rocm => "rocm",
rlx::Device::Gpu => "wgpu",
rlx::Device::Vulkan => "vulkan",
rlx::Device::Tpu => "tpu",
_ => "unknown",
}
}
pub fn require_available(device: rlx::Device) -> bool {
rlx::runtime::is_available(device)
}
pub fn available_device_labels() -> Vec<&'static str> {
let mut out = Vec::new();
for d in rlx::runtime::available_devices() {
let label = device_label(d);
if label != "unknown" && !out.contains(&label) {
out.push(label);
}
}
if out.is_empty() {
out.push("cpu");
}
out
}