pub const DEFAULT_MAX_SEQ_LEN: usize = 128;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const DEFAULT_ONNX_INFERENCE_BATCH_SIZE: usize = 1;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const DEFAULT_DYNAMIC_ONNX_INFERENCE_BATCH_SIZE: usize = 32;
pub const DEFAULT_MIGRAPHX_INFERENCE_BATCH_SIZE: usize = 8;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
const MAX_ONNX_INFERENCE_BATCH_SIZE: usize = 256;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const ONNX_INFERENCE_BATCH_SIZE_ENV: &str = "LEINDEX_ONNX_INFERENCE_BATCH_SIZE";
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
const MIN_ONNX_SEQUENCE_LEN: usize = 8;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const MAX_ONNX_SEQUENCE_LEN: usize = 512;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const ONNX_SEQUENCE_LEN_ENV: &str = "LEINDEX_ONNX_SEQUENCE_LEN";
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const ONNX_LOG_SHAPES_ENV: &str = "LEINDEX_ONNX_LOG_SHAPES";
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const MIGRAPHX_FP16_ENV: &str = "LEINDEX_MIGRAPHX_FP16";
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const MIGRAPHX_EXHAUSTIVE_TUNE_ENV: &str = "LEINDEX_MIGRAPHX_EXHAUSTIVE_TUNE";
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) const MIGRAPHX_MODEL_CACHE_PATH_ENV: &str = "ORT_MIGRAPHX_MODEL_CACHE_PATH";
pub const DEFAULT_MIN_AVAILABLE_MB: u64 = 2048;
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub fn configured_onnx_inference_batch_size(model_name: &str, provider: &str) -> usize {
std::env::var(ONNX_INFERENCE_BATCH_SIZE_ENV)
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&v| v > 0)
.map(|v| v.min(MAX_ONNX_INFERENCE_BATCH_SIZE))
.unwrap_or_else(|| {
if provider.eq_ignore_ascii_case("migraphx") || provider.eq_ignore_ascii_case("rocm") {
DEFAULT_MIGRAPHX_INFERENCE_BATCH_SIZE
} else if model_name.ends_with("-dynamic") {
DEFAULT_DYNAMIC_ONNX_INFERENCE_BATCH_SIZE
} else {
DEFAULT_ONNX_INFERENCE_BATCH_SIZE
}
})
}
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub fn configured_onnx_sequence_len() -> usize {
std::env::var(ONNX_SEQUENCE_LEN_ENV)
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&v| v >= MIN_ONNX_SEQUENCE_LEN)
.map(|v| v.min(MAX_ONNX_SEQUENCE_LEN))
.unwrap_or(DEFAULT_MAX_SEQ_LEN)
}
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) fn env_flag(name: &str) -> bool {
std::env::var(name)
.ok()
.map(|value| {
matches!(
value.to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
)
})
.unwrap_or(false)
}
pub(crate) fn default_ort_threads() -> usize {
std::thread::available_parallelism()
.map(|n| {
let n = n.get();
((n.saturating_mul(3) / 4).max(2)).min(n)
})
.unwrap_or(2)
}
#[cfg(target_os = "linux")]
pub(crate) fn process_rss_kib() -> Option<u64> {
let statm = std::fs::read_to_string("/proc/self/statm").ok()?;
let resident_pages: u64 = statm.split_whitespace().nth(1)?.parse().ok()?;
let page_size = unsafe { libc::sysconf(libc::_SC_PAGESIZE) };
if page_size <= 0 {
return None;
}
Some(resident_pages.saturating_mul(page_size as u64) / 1024)
}
#[cfg(not(target_os = "linux"))]
pub(crate) fn process_rss_kib() -> Option<u64> {
None
}
#[cfg(target_os = "linux")]
pub(crate) fn mem_available_kib() -> Option<u64> {
let contents = std::fs::read_to_string("/proc/meminfo").ok()?;
let line = contents
.lines()
.find(|line| line.starts_with("MemAvailable:"))?;
line.split_whitespace().nth(1)?.parse().ok()
}
#[cfg(not(target_os = "linux"))]
pub(crate) fn mem_available_kib() -> Option<u64> {
None
}
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) fn build_position_ids(batch_size: usize, sequence_len: usize) -> Vec<i64> {
(0..batch_size)
.flat_map(|_| (0..sequence_len).map(|position| position as i64))
.collect()
}
#[cfg_attr(not(feature = "onnx"), allow(dead_code))]
pub(crate) fn prune_migraphx_cache(dir: &std::path::Path, keep: usize) {
let Ok(entries) = std::fs::read_dir(dir) else {
return;
};
let mut mxr: Vec<(std::path::PathBuf, std::time::SystemTime)> = entries
.filter_map(|entry| entry.ok())
.map(|entry| entry.path())
.filter(|path| path.extension().is_some_and(|ext| ext == "mxr"))
.filter_map(|path| {
std::fs::metadata(&path)
.and_then(|metadata| metadata.modified())
.ok()
.map(|mtime| (path, mtime))
})
.collect();
mxr.sort_by_key(|item| std::cmp::Reverse(item.1)); for (path, _) in mxr.into_iter().skip(keep) {
match std::fs::remove_file(&path) {
Ok(()) => tracing::debug!("pruned stale MIGraphX cache file: {}", path.display()),
Err(error) => {
tracing::warn!(
"failed to prune stale MIGraphX cache file {}: {}",
path.display(),
error
)
}
}
}
}
pub(crate) fn unix_now_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|duration| duration.as_millis().min(u64::MAX as u128) as u64)
.unwrap_or(0)
}