//! One session builder per rung (docs/inference.md §2, §7, §9). use ort::session::builder::GraphOptimizationLevel; use ort::session::Session; use crate::{Config, Role, Rung}; /// Build a session for `bytes` on `rung`. /// /// `strict` is the probe's flag: with it, a provider that would hand any /// node to the CPU fails the build instead, so "the session built" means /// "the provider took the graph" and not "the provider registered" (§4). pub fn build( rung: Rung, role: Role, bytes: &[u8], cfg: &Config, strict: bool, ) -> ort::Result { let mut b = Session::builder()? .with_optimization_level(GraphOptimizationLevel::Level3)? .with_intra_threads(threads(cfg))?; if strict { b = b.with_config_entry("session.disable_cpu_ep_fallback", "1")?; } // 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::Hexagon => unreachable!("the Hexagon rung is not on a desktop ladder"), } } #[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 => unreachable!("no NVIDIA rung on Android"), } }