//! One session builder per rung (docs/dev/inference.md §2, §7, §9). use ort::session::Session; use crate::{Config, Role, Rung}; /// Build a session for `bytes` on `rung`. /// /// Not strict about the CPU: `session.disable_cpu_ep_fallback` was tried as /// the probe's proof that a provider took the graph, and it refuses the /// Hexagon over the ten quantise/dequantise nodes at the graph's edges that /// QNN declines by policy and that cost microseconds. The probe's proof is /// its clock instead (§4): a provider that hands real work to the CPU is /// slower than the CPU floor and rejected by the same measurement. pub fn build(rung: Rung, role: Role, bytes: &[u8], cfg: &Config) -> ort::Result { // No optimisation level named. ONNX Runtime's default is already its // fullest, and on tract any level but "disabled" means `into_optimized`, // whose optimiser divides by zero inside yolo26n-seg (tract-data // `stack_tensors`) — a panic across the C API, which is an abort. The // app never asked tract for that and does not start now. let mut b = Session::builder()?.with_intra_threads(threads(cfg))?; // A Hexagon session loads the compiled context when there is one and // compiles it from the model when there is not; the engine thread is // what makes the second case rare (§6). let context = (rung == Rung::Hexagon).then(|| crate::engines::context_path(cfg, bytes)); let ready = context.as_ref().is_some_and(|p| p.is_file()); b = providers( b, rung, role, cfg, if ready { None } else { context.as_deref() }, )?; match (ready, context) { (true, Some(path)) => b.commit_from_file(path), _ => b.commit_from_memory(bytes), } } /// The intra-op pool: what the config says, else the cores less two for /// the compositor and the decoder (§9). tract ignores it. fn threads(cfg: &Config) -> usize { if cfg.threads > 0 { return cfg.threads; } std::thread::available_parallelism() .map(|n| n.get().saturating_sub(2).max(1)) .unwrap_or(1) } #[cfg(not(target_os = "android"))] fn providers( b: ort::session::builder::SessionBuilder, rung: Rung, role: Role, cfg: &Config, _generate_context: Option<&std::path::Path>, ) -> ort::Result { use ort::ep; match rung { Rung::Cpu => Ok(b), Rung::Cuda => { Ok(b.with_execution_providers([ep::CUDA::default().build().error_on_failure()])?) } Rung::TensorRt => { let cache = cfg.cache_dir.join("tensorrt"); let _ = std::fs::create_dir_all(&cache); let cache = cache.to_string_lossy().into_owned(); // fp16 for everything but the embedder, whose comparability // across devices is worth more than its 0.2 ms (§7). The // workspace cap keeps the develop view's tiles on the card // (NFR-RES-2). CUDA behind it takes any node TensorRT declines. Ok(b.with_execution_providers([ ep::TensorRT::default() .with_fp16(role != Role::Embedder) .with_engine_cache(true) .with_engine_cache_path(&cache) .with_timing_cache(true) .with_timing_cache_path(&cache) .with_max_workspace_size(512 << 20) .build() .error_on_failure(), 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, rung: Rung, _role: Role, _cfg: &Config, generate_context: Option<&std::path::Path>, ) -> ort::Result { use ort::ep; match rung { Rung::Cpu => Ok(b), Rung::Hexagon => { // The HTP compiles the graph once per device (0.8–1.7 s here). // With `ep.context_enable` ONNX Runtime writes the compiled // context beside the probe cache; the next session loads that // file as its model and skips the compile (§5). let mut b = b; if let Some(ctx) = generate_context { let _ = std::fs::create_dir_all(ctx.parent().unwrap()); b = b .with_config_entry("ep.context_enable", "1")? .with_config_entry("ep.context_file_path", ctx.to_string_lossy())? .with_config_entry("ep.context_embed_mode", "0")?; } // Quantise/dequantise at the graph's edges stay on the NPU too, // so a strict build is a whole-graph build. Ok(b.with_execution_providers([ep::QNN::default() .with_backend_path("libQnnHtp.so") .with_performance_mode(ep::qnn::PerformanceMode::Burst) .with_offload_graph_io_quantization(false) .build() .error_on_failure()])?) } Rung::Cuda | Rung::TensorRt | Rung::MiGraphX => { unreachable!("no desktop GPU rung on Android") } } }