use std::path::Path;
use std::sync::{Arc, Mutex};
use ort::session::Session;
use crate::runtime::{
error::RuntimeError, factory::Runtime, session::RuntimeSession, tensor::Tensor,
};
use super::{factory::OrtExecutionProvider, tensor::value_to_tensor};
pub struct OrtRuntime {
intra_threads: usize,
provider: OrtExecutionProvider,
prepacked: Option<Arc<ort::session::builder::PrepackedWeights>>,
optimized_cache_dir: Option<std::path::PathBuf>,
}
impl OrtRuntime {
pub(crate) fn new(
intra_threads: usize,
provider: OrtExecutionProvider,
prepacked: Option<Arc<ort::session::builder::PrepackedWeights>>,
optimized_cache_dir: Option<std::path::PathBuf>,
) -> Self {
Self {
intra_threads,
provider,
prepacked,
optimized_cache_dir,
}
}
}
fn load_failed(path: &Path, e: impl std::fmt::Display) -> RuntimeError {
RuntimeError::LoadFailed {
path: path.into(),
message: e.to_string(),
}
}
fn cache_is_fresh(
cache_len: u64,
cache_mtime: std::time::SystemTime,
source_mtime: std::time::SystemTime,
) -> bool {
cache_len > 0 && cache_mtime >= source_mtime
}
fn optimized_cache_is_fresh(cache_path: &Path, source_path: &Path) -> bool {
let (Ok(cache), Ok(source)) = (
std::fs::metadata(cache_path),
std::fs::metadata(source_path),
) else {
return false;
};
match (cache.modified(), source.modified()) {
(Ok(cache_mtime), Ok(source_mtime)) => {
cache_is_fresh(cache.len(), cache_mtime, source_mtime)
}
_ => false,
}
}
fn optimized_cache_path(cache_dir: &Path, model_path: &Path) -> std::path::PathBuf {
let basename = crate::model::optimized_cache_basename(model_path)
.unwrap_or_else(|| "encoder_optimized.ort".into());
cache_dir.join(basename)
}
impl OrtRuntime {
fn session_builder(
&self,
model_path: &Path,
is_encoder: bool,
) -> Result<ort::session::builder::SessionBuilder, RuntimeError> {
let mut builder = Session::builder().map_err(|e| load_failed(model_path, e))?;
if let Some(prepacked) = self.prepacked.as_ref() {
builder = builder
.with_prepacked_weights(prepacked)
.map_err(|e| load_failed(model_path, e))?;
}
let eps = self.provider.execution_providers(model_path);
builder = builder
.with_execution_providers(&eps)
.map_err(|e| load_failed(model_path, e))?;
if self.provider.is_cpu() {
let intra_threads = if is_encoder {
self.intra_threads.max(1)
} else {
1
};
builder = builder
.with_intra_threads(intra_threads)
.map_err(|e| load_failed(model_path, e))?;
builder = builder
.with_inter_threads(1)
.map_err(|e| load_failed(model_path, e))?;
}
Ok(builder)
}
fn try_load_cached_encoder(&self, model_path: &Path) -> Option<Result<Session, RuntimeError>> {
if !self.provider.is_cpu() {
return None;
}
let cache_dir = self.optimized_cache_dir.as_ref()?;
let cache_path = optimized_cache_path(cache_dir, model_path);
if !optimized_cache_is_fresh(&cache_path, model_path) {
return None;
}
let result = self
.session_builder(model_path, true)
.and_then(|mut builder| {
builder = builder
.with_config_entry("session.load_model_format", "ORT")
.map_err(|e| load_failed(&cache_path, e))?;
builder = builder
.with_config_entry("session.use_memory_mapped_ort_model", "1")
.map_err(|e| load_failed(&cache_path, e))?;
builder = builder
.with_config_entry("session.use_ort_model_bytes_for_initializers", "1")
.map_err(|e| load_failed(&cache_path, e))?;
builder = builder
.with_prepacking(false)
.map_err(|e| load_failed(&cache_path, e))?;
builder
.commit_from_file(&cache_path)
.map_err(|e| load_failed(&cache_path, e))
});
match result {
Ok(session) => {
tracing::info!(
path = %cache_path.display(),
"encoder: loaded fresh optimized graph cache, skipped source model"
);
Some(Ok(session))
}
Err(e) => {
tracing::warn!(
path = %cache_path.display(),
error = %e,
"encoder: optimized graph cache failed to load; deleting it and falling back to the source model"
);
if let Err(rm) = std::fs::remove_file(&cache_path) {
tracing::warn!(
path = %cache_path.display(),
error = %rm,
"encoder: failed to delete broken optimized graph cache"
);
}
None
}
}
}
}
impl Runtime for OrtRuntime {
fn load_session(
&self,
model_path: &Path,
is_encoder: bool,
) -> Result<Box<dyn RuntimeSession>, RuntimeError> {
if is_encoder && let Some(result) = self.try_load_cached_encoder(model_path) {
let session = result?;
return Ok(Box::new(OrtSession {
session: Mutex::new(session),
}));
}
let mut builder = self.session_builder(model_path, is_encoder)?;
if self.provider.is_cpu()
&& is_encoder
&& let Some(cache_dir) = &self.optimized_cache_dir
{
std::fs::create_dir_all(cache_dir).map_err(|e| load_failed(model_path, e))?;
let cache_path = optimized_cache_path(cache_dir, model_path);
tracing::info!(
path = %cache_path.display(),
"encoder: loading source model and refreshing optimized graph cache"
);
builder = builder
.with_optimized_model_path(&cache_path)
.map_err(|e| load_failed(model_path, e))?;
builder = builder
.with_config_entry("session.save_model_format", "ORT")
.map_err(|e| load_failed(model_path, e))?;
}
let session = builder
.commit_from_file(model_path)
.map_err(|e| load_failed(model_path, e))?;
Ok(Box::new(OrtSession {
session: Mutex::new(session),
}))
}
}
pub struct OrtSession {
session: Mutex<Session>,
}
impl RuntimeSession for OrtSession {
fn run(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
let session_inputs: Vec<ort::session::SessionInputValue<'_>> = inputs
.iter()
.map(Tensor::as_ort_input)
.collect::<Result<_, _>>()?;
let mut session = self
.session
.lock()
.map_err(|_| RuntimeError::InferenceFailed("ort session mutex poisoned".into()))?;
let outputs = session
.run(&session_inputs[..])
.map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
outputs
.into_iter()
.map(|(_name, value)| value_to_tensor(value))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
use std::time::{Duration, SystemTime};
fn write_file(path: &Path, bytes: &[u8]) {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).unwrap();
}
let mut f = std::fs::File::create(path).unwrap();
f.write_all(bytes).unwrap();
}
#[test]
fn test_cache_is_fresh_nonempty_and_not_older_than_source() {
let now = SystemTime::now();
assert!(cache_is_fresh(100, now, now));
assert!(cache_is_fresh(100, now + Duration::from_secs(1), now));
}
#[test]
fn test_cache_is_fresh_rejects_stale_or_empty() {
let now = SystemTime::now();
assert!(!cache_is_fresh(100, now - Duration::from_secs(1), now));
assert!(!cache_is_fresh(0, now + Duration::from_secs(1), now));
}
#[test]
fn test_optimized_cache_is_fresh_with_real_files() {
let tmp = tempfile::tempdir().unwrap();
let dir = tmp.path();
let source = dir.join("v3_rnnt_encoder_int8.onnx");
let cache = dir.join("optimized_cache/v3_rnnt_encoder_int8_optimized.ort");
write_file(&source, b"source");
write_file(&cache, b"optimized");
assert!(optimized_cache_is_fresh(&cache, &source));
}
#[test]
fn test_optimized_cache_is_fresh_missing_or_empty_cache() {
let tmp = tempfile::tempdir().unwrap();
let dir = tmp.path();
let source = dir.join("encoder.onnx");
let cache = dir.join("optimized_cache/encoder_optimized.ort");
write_file(&source, b"source");
assert!(!optimized_cache_is_fresh(&cache, &source));
write_file(&cache, b"");
assert!(!optimized_cache_is_fresh(&cache, &source));
assert!(!optimized_cache_is_fresh(&cache, &dir.join("absent.onnx")));
}
#[test]
fn test_optimized_cache_path_matches_gc_keep_name() {
let cache_dir = Path::new("/models/optimized_cache");
let encoder = Path::new("/models/v3_rnnt_encoder_int8.onnx");
assert_eq!(
optimized_cache_path(cache_dir, encoder),
Path::new("/models/optimized_cache/v3_rnnt_encoder_int8_optimized.ort")
);
}
}