use std::path::Path;
use std::sync::Mutex;
use fathomdb_embedder_api::{Embedder, EmbedderError, EmbedderIdentity, Vector};
use ort::execution_providers::{
CPUExecutionProvider, CUDAExecutionProvider, DirectMLExecutionProvider, ExecutionProvider,
ExecutionProviderDispatch, OpenVINOExecutionProvider, ROCmExecutionProvider,
};
use ort::session::Session;
use ort::value::Tensor;
use tokenizers::{Tokenizer, TruncationParams};
use crate::device::{parse_device_request, DeviceRequest};
pub const ORT_BGE_EMBEDDER_NAME: &str = "fathomdb-bge-small-en-v1.5-onnx";
pub const ORT_BGE_EMBEDDER_DIM: u32 = 384;
const HF_REVISION: &str = "5c38ec7c405ec4b44b94cc5a9bb96e735b38267a";
const MAX_SEQUENCE_TOKENS: usize = 512;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OrtPooling {
Mean,
Cls,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OrtProvider {
Cpu,
Cuda(i32),
Rocm(i32),
DirectMl(i32),
OpenVino,
}
impl OrtProvider {
fn label(self) -> &'static str {
match self {
OrtProvider::Cpu => "cpu",
OrtProvider::Cuda(_) => "cuda",
OrtProvider::Rocm(_) => "rocm",
OrtProvider::DirectMl(_) => "directml",
OrtProvider::OpenVino => "openvino",
}
}
}
pub(crate) fn map_device_request(req: &DeviceRequest) -> (OrtProvider, Option<String>) {
match req {
DeviceRequest::Cpu => (OrtProvider::Cpu, None),
DeviceRequest::Cuda(idx) => (OrtProvider::Cuda(*idx as i32), None),
DeviceRequest::Metal => (
OrtProvider::Cpu,
Some(
"FATHOMDB_EMBED_DEVICE=metal is a candle backend; the ONNX path has no Metal \
execution provider (use rocm|directml|openvino for cross-vendor GPUs, or \
candle's embed-metal build); using CPU"
.to_string(),
),
),
DeviceRequest::Unknown(name) => map_extended_provider(name),
}
}
fn map_extended_provider(raw: &str) -> (OrtProvider, Option<String>) {
let (head, idx) = match raw.split_once(':') {
Some((h, i)) => (h, i.parse::<i32>().unwrap_or(0)),
None => (raw, 0),
};
match head {
"rocm" => (OrtProvider::Rocm(idx), None),
"directml" | "dml" => (OrtProvider::DirectMl(idx), None),
"openvino" | "ovep" => (OrtProvider::OpenVino, None),
other => (
OrtProvider::Cpu,
Some(format!(
"FATHOMDB_EMBED_DEVICE={other} is not a recognized ONNX execution provider \
(expected cpu|cuda|cuda:N|rocm|rocm:N|directml|openvino); using CPU"
)),
),
}
}
#[allow(clippy::print_stderr)] fn emit_onnx_warning(msg: &str) {
eprintln!("fathomdb-embedder(onnx): {msg}");
}
fn resolve_provider_from_env() -> OrtProvider {
let raw = std::env::var("FATHOMDB_EMBED_DEVICE").unwrap_or_default();
let (provider, warn) = map_device_request(&parse_device_request(&raw));
if let Some(msg) = warn {
emit_onnx_warning(&msg);
}
provider
}
fn resolve_effective_provider_with(
requested: OrtProvider,
is_available: impl Fn(OrtProvider) -> bool,
) -> (OrtProvider, Option<String>) {
if matches!(requested, OrtProvider::Cpu) {
return (OrtProvider::Cpu, None);
}
if is_available(requested) {
(requested, None)
} else {
(
OrtProvider::Cpu,
Some(format!(
"requested ONNX execution provider {requested:?} is unavailable in this ONNX \
Runtime build/runtime (ort is built default-features=false, so its own fallback \
log is compiled out); falling back to CPU"
)),
)
}
}
fn resolve_effective_provider(requested: OrtProvider) -> (OrtProvider, Option<String>) {
resolve_effective_provider_with(requested, provider_is_available)
}
fn provider_is_available(provider: OrtProvider) -> bool {
match provider {
OrtProvider::Cpu => true,
OrtProvider::Cuda(_) => CUDAExecutionProvider::default().is_available().unwrap_or(false),
OrtProvider::Rocm(_) => ROCmExecutionProvider::default().is_available().unwrap_or(false),
OrtProvider::DirectMl(_) => {
DirectMLExecutionProvider::default().is_available().unwrap_or(false)
}
OrtProvider::OpenVino => {
OpenVINOExecutionProvider::default().is_available().unwrap_or(false)
}
}
}
fn provider_dispatch(provider: OrtProvider) -> ExecutionProviderDispatch {
match provider {
OrtProvider::Cpu => CPUExecutionProvider::default().build(),
OrtProvider::Cuda(idx) => {
CUDAExecutionProvider::default().with_device_id(idx).build().error_on_failure()
}
OrtProvider::Rocm(idx) => {
ROCmExecutionProvider::default().with_device_id(idx).build().error_on_failure()
}
OrtProvider::DirectMl(idx) => {
DirectMLExecutionProvider::default().with_device_id(idx).build().error_on_failure()
}
OrtProvider::OpenVino => OpenVINOExecutionProvider::default().build().error_on_failure(),
}
}
fn build_session_with_fallback<S, E, F>(
effective: OrtProvider,
build: F,
) -> Result<(S, Option<String>), E>
where
F: Fn(OrtProvider) -> Result<S, E>,
E: std::fmt::Display,
{
match build(effective) {
Ok(session) => Ok((session, None)),
Err(e) if matches!(effective, OrtProvider::Cpu) => Err(e),
Err(e) => {
let warn = format!(
"requested ONNX execution provider {effective:?} was reported available but \
FAILED during ONNX Runtime session creation ({e}); falling back to CPU"
);
let session = build(OrtProvider::Cpu)?;
Ok((session, Some(warn)))
}
}
}
fn sha256_file(path: &Path) -> Result<[u8; 32], EmbedderError> {
use sha2::{Digest, Sha256};
let mut file = std::fs::File::open(path).map_err(|e| err("asset digest open", e))?;
let mut hasher = Sha256::new();
let mut buf = [0_u8; 64 * 1024];
loop {
let n =
std::io::Read::read(&mut file, &mut buf).map_err(|e| err("asset digest read", e))?;
if n == 0 {
break;
}
hasher.update(&buf[..n]);
}
Ok(hasher.finalize().into())
}
fn compose_asset_revision(model_digest: &[u8; 32], tokenizer_digest: &[u8; 32]) -> String {
use sha2::{Digest, Sha256};
let mut hasher = Sha256::new();
hasher.update(model_digest);
hasher.update(tokenizer_digest);
let combined = hasher.finalize();
let mut hex = String::with_capacity(12);
for byte in combined.iter().take(6) {
use std::fmt::Write as _;
let _ = write!(hex, "{byte:02x}");
}
format!("{HF_REVISION}+onnx-{hex}")
}
fn derive_asset_revision(
model_path: &Path,
tokenizer_path: &Path,
) -> Result<String, EmbedderError> {
let model_digest = sha256_file(model_path)?;
let tokenizer_digest = sha256_file(tokenizer_path)?;
Ok(compose_asset_revision(&model_digest, &tokenizer_digest))
}
pub struct OrtBgeEmbedder {
identity: EmbedderIdentity,
tokenizer: Tokenizer,
session: Mutex<Session>,
pooling: OrtPooling,
effective_provider: OrtProvider,
}
fn err(context: &str, e: impl std::fmt::Display) -> EmbedderError {
EmbedderError::Failed { message: format!("ort_bge {context}: {e}") }
}
impl OrtBgeEmbedder {
pub fn from_files(model_path: &Path, tokenizer_path: &Path) -> Result<Self, EmbedderError> {
Self::from_files_with_provider(model_path, tokenizer_path, resolve_provider_from_env())
}
pub fn from_env() -> Result<Self, EmbedderError> {
let model = std::env::var("FATHOMDB_ONNX_MODEL_PATH")
.map_err(|_| err("from_env", "FATHOMDB_ONNX_MODEL_PATH is unset"))?;
let tok = std::env::var("FATHOMDB_ONNX_TOKENIZER_PATH")
.map_err(|_| err("from_env", "FATHOMDB_ONNX_TOKENIZER_PATH is unset"))?;
Self::from_files(Path::new(&model), Path::new(&tok))
}
fn from_files_with_provider(
model_path: &Path,
tokenizer_path: &Path,
provider: OrtProvider,
) -> Result<Self, EmbedderError> {
let mut tokenizer =
Tokenizer::from_file(tokenizer_path).map_err(|e| err("tokenizer load", e))?;
tokenizer
.with_truncation(Some(TruncationParams {
max_length: MAX_SEQUENCE_TOKENS,
..Default::default()
}))
.map_err(|e| err("tokenizer truncation", e))?;
let (effective, avail_warn) = resolve_effective_provider(provider);
if let Some(msg) = avail_warn {
emit_onnx_warning(&msg);
}
let (session, build_warn) = build_session_with_fallback(effective, |p| {
Session::builder()?
.with_execution_providers([provider_dispatch(p)])?
.commit_from_file(model_path)
})
.map_err(|e| err("session build", e))?;
let effective_provider = if build_warn.is_some() { OrtProvider::Cpu } else { effective };
if let Some(msg) = build_warn {
emit_onnx_warning(&msg);
}
let revision = derive_asset_revision(model_path, tokenizer_path)?;
let identity = EmbedderIdentity::new(ORT_BGE_EMBEDDER_NAME, revision, ORT_BGE_EMBEDDER_DIM);
Ok(Self {
identity,
tokenizer,
session: Mutex::new(session),
pooling: OrtPooling::Cls,
effective_provider,
})
}
#[must_use]
pub fn with_pooling(mut self, pooling: OrtPooling) -> Self {
self.pooling = pooling;
self
}
#[must_use]
pub fn effective_provider(&self) -> &'static str {
self.effective_provider.label()
}
}
impl Embedder for OrtBgeEmbedder {
fn identity(&self) -> EmbedderIdentity {
self.identity.clone()
}
fn embed(&self, input: &str) -> Result<Vector, EmbedderError> {
let encoding = self.tokenizer.encode(input, true).map_err(|e| err("tokenize", e))?;
let ids: Vec<i64> = encoding.get_ids().iter().map(|&x| i64::from(x)).collect();
let mask: Vec<i64> = encoding.get_attention_mask().iter().map(|&x| i64::from(x)).collect();
let len = ids.len();
let token_type: Vec<i64> = vec![0; len];
let shape = vec![1_i64, len as i64];
let ids_t = Tensor::from_array((shape.clone(), ids)).map_err(|e| err("input_ids", e))?;
let mask_t =
Tensor::from_array((shape.clone(), mask)).map_err(|e| err("attention_mask", e))?;
let tt_t = Tensor::from_array((shape, token_type)).map_err(|e| err("token_type_ids", e))?;
let mut session = self.session.lock().map_err(|_| err("session lock", "poisoned"))?;
let outputs = session
.run(ort::inputs![
"input_ids" => ids_t,
"attention_mask" => mask_t,
"token_type_ids" => tt_t,
])
.map_err(|e| err("forward", e))?;
let (out_shape, data) =
outputs[0].try_extract_tensor::<f32>().map_err(|e| err("extract", e))?;
let dims: Vec<usize> = out_shape.iter().map(|&d| d as usize).collect();
if dims.len() != 3 {
return Err(err("output shape", format!("expected rank-3 (1,L,H), got {dims:?}")));
}
let (seq_len, hidden) = (dims[1], dims[2]);
if hidden != ORT_BGE_EMBEDDER_DIM as usize {
return Err(err(
"output dim",
format!("expected hidden {ORT_BGE_EMBEDDER_DIM}, got {hidden}"),
));
}
let mut pooled = vec![0.0_f32; hidden];
match self.pooling {
OrtPooling::Cls => {
pooled.copy_from_slice(&data[0..hidden]);
}
OrtPooling::Mean => {
for pos in 0..seq_len {
let base = pos * hidden;
for (j, slot) in pooled.iter_mut().enumerate() {
*slot += data[base + j];
}
}
let denom = seq_len.max(1) as f32;
for slot in &mut pooled {
*slot /= denom;
}
}
}
let norm = pooled.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
for slot in &mut pooled {
*slot /= norm;
}
Ok(pooled)
}
}
#[cfg(test)]
mod tests {
use super::{map_device_request, OrtProvider};
use crate::device::parse_device_request;
fn resolve(raw: &str) -> (OrtProvider, Option<String>) {
map_device_request(&parse_device_request(raw))
}
#[test]
fn cpu_and_unset_map_to_cpu_no_warning() {
for raw in ["", "cpu", " CPU "] {
let (p, warn) = resolve(raw);
assert_eq!(p, OrtProvider::Cpu, "{raw:?}");
assert!(warn.is_none(), "{raw:?} should not warn");
}
}
#[test]
fn cuda_maps_to_cuda_provider_with_index() {
assert_eq!(resolve("cuda").0, OrtProvider::Cuda(0));
assert_eq!(resolve("cuda:1").0, OrtProvider::Cuda(1));
assert_eq!(resolve("cuda:2").0, OrtProvider::Cuda(2));
assert!(resolve("cuda:1").1.is_none());
}
#[test]
fn rocm_maps_to_rocm_provider() {
assert_eq!(resolve("rocm").0, OrtProvider::Rocm(0));
assert_eq!(resolve("rocm:1").0, OrtProvider::Rocm(1));
assert!(resolve("rocm").1.is_none());
assert!(resolve("ROCm").1.is_none()); }
#[test]
fn directml_maps_to_directml_provider() {
assert_eq!(resolve("directml").0, OrtProvider::DirectMl(0));
assert_eq!(resolve("dml").0, OrtProvider::DirectMl(0));
assert_eq!(resolve("directml:1").0, OrtProvider::DirectMl(1));
assert!(resolve("directml").1.is_none());
}
#[test]
fn openvino_maps_to_openvino_provider() {
assert_eq!(resolve("openvino").0, OrtProvider::OpenVino);
assert_eq!(resolve("ovep").0, OrtProvider::OpenVino);
assert!(resolve("openvino").1.is_none());
}
#[test]
fn metal_falls_back_to_cpu_loudly() {
let (p, warn) = resolve("metal");
assert_eq!(p, OrtProvider::Cpu);
assert!(warn.is_some(), "metal must warn on CPU fallback");
}
#[test]
fn unrecognized_device_falls_back_to_cpu_loudly() {
for raw in ["vulkan", "tpu", "gpu"] {
let (p, warn) = resolve(raw);
assert_eq!(p, OrtProvider::Cpu, "{raw:?}");
assert!(warn.is_some(), "{raw:?} must warn on CPU fallback");
}
}
#[test]
fn requested_gpu_unavailable_falls_back_to_cpu_loudly() {
for requested in [
OrtProvider::Cuda(0),
OrtProvider::Rocm(1),
OrtProvider::DirectMl(0),
OrtProvider::OpenVino,
] {
let (eff, warn) = super::resolve_effective_provider_with(requested, |_| false);
assert_eq!(eff, OrtProvider::Cpu, "{requested:?} must downgrade to CPU");
let msg = warn.expect("unavailable GPU provider must emit a warning");
assert!(
msg.contains(&format!("{requested:?}")),
"warning must name the requested provider {requested:?}, got {msg:?}"
);
assert!(
msg.to_lowercase().contains("cpu"),
"warning must state the CPU fallback, got {msg:?}"
);
}
}
#[test]
fn requested_available_provider_is_honored_without_warning() {
let (eff, warn) = super::resolve_effective_provider_with(OrtProvider::Cuda(2), |_| true);
assert_eq!(eff, OrtProvider::Cuda(2));
assert!(warn.is_none(), "an available provider must not warn");
}
#[test]
fn cpu_request_never_probes_and_never_warns() {
let (eff, warn) = super::resolve_effective_provider_with(OrtProvider::Cpu, |_| {
panic!("CPU request must not probe provider availability")
});
assert_eq!(eff, OrtProvider::Cpu);
assert!(warn.is_none());
}
#[test]
fn session_build_gpu_fails_retries_cpu_with_warning() {
use std::cell::RefCell;
for requested in [OrtProvider::Cuda(0), OrtProvider::Rocm(1), OrtProvider::OpenVino] {
let attempts = RefCell::new(Vec::new());
let build = |p: OrtProvider| -> Result<&'static str, String> {
attempts.borrow_mut().push(p);
match p {
OrtProvider::Cpu => Ok("cpu-session"),
other => Err(format!("no runtime lib for {other:?}")),
}
};
let (session, warn) = super::build_session_with_fallback(requested, build)
.expect("must retry CPU and succeed");
assert_eq!(session, "cpu-session", "{requested:?} must yield the CPU session");
let msg = warn.expect("a runtime GPU-build failure must emit a warning");
assert!(
msg.contains(&format!("{requested:?}")),
"warning must name the requested provider {requested:?}, got {msg:?}"
);
assert!(
msg.to_lowercase().contains("cpu"),
"warning must state the CPU fallback, got {msg:?}"
);
assert_eq!(
*attempts.borrow(),
vec![requested, OrtProvider::Cpu],
"must attempt the GPU provider first, then CPU"
);
}
}
#[test]
fn session_build_gpu_and_cpu_both_fail_errors() {
let build =
|p: OrtProvider| -> Result<&'static str, String> { Err(format!("hard fail {p:?}")) };
let res = super::build_session_with_fallback(OrtProvider::Cuda(0), build);
assert!(res.is_err(), "both GPU and CPU failing must be an error");
}
#[test]
fn session_build_cpu_failure_is_error_without_retry() {
use std::cell::Cell;
let calls = Cell::new(0);
let build = |_p: OrtProvider| -> Result<&'static str, String> {
calls.set(calls.get() + 1);
Err("cpu build failed".to_string())
};
let res = super::build_session_with_fallback(OrtProvider::Cpu, build);
assert!(res.is_err(), "a CPU build failure must be an error");
assert_eq!(calls.get(), 1, "a CPU failure must NOT retry");
}
#[test]
fn session_build_success_no_warning_single_attempt() {
use std::cell::Cell;
let calls = Cell::new(0);
let build = |_p: OrtProvider| -> Result<&'static str, String> {
calls.set(calls.get() + 1);
Ok("session")
};
let (session, warn) = super::build_session_with_fallback(OrtProvider::Cuda(2), build)
.expect("build succeeds");
assert_eq!(session, "session");
assert!(warn.is_none(), "a successful build must not warn");
assert_eq!(calls.get(), 1, "a successful build must attempt exactly once");
}
#[test]
fn non_cpu_dispatch_errors_on_failure_cpu_does_not() {
for p in [
OrtProvider::Cuda(0),
OrtProvider::Rocm(1),
OrtProvider::DirectMl(0),
OrtProvider::OpenVino,
] {
let dbg = format!("{:?}", super::provider_dispatch(p));
assert!(
dbg.contains("error_on_failure: true"),
"non-CPU provider {p:?} must set error_on_failure, got {dbg:?}"
);
}
let cpu = format!("{:?}", super::provider_dispatch(OrtProvider::Cpu));
assert!(
cpu.contains("error_on_failure: false"),
"CPU provider must keep the silent (fallback-target) default, got {cpu:?}"
);
}
mod asset_revision {
use std::io::Write as _;
use std::path::Path;
use super::super::{compose_asset_revision, derive_asset_revision, HF_REVISION};
fn write_temp(name: &str, bytes: &[u8]) -> std::path::PathBuf {
let dir = std::env::temp_dir();
let path = dir.join(format!("fathomdb-ort-rev-{}-{name}", std::process::id()));
let mut f = std::fs::File::create(&path).expect("create temp asset");
f.write_all(bytes).expect("write temp asset");
f.flush().expect("flush temp asset");
path
}
fn revision_of(model_bytes: &[u8], tok_bytes: &[u8], tag: &str) -> String {
let m = write_temp(&format!("model-{tag}.onnx"), model_bytes);
let t = write_temp(&format!("tok-{tag}.json"), tok_bytes);
let rev = derive_asset_revision(Path::new(&m), Path::new(&t)).expect("derive revision");
let _ = std::fs::remove_file(&m);
let _ = std::fs::remove_file(&t);
rev
}
#[test]
fn distinct_assets_yield_distinct_revisions() {
let a = revision_of(b"model-bytes-AAAA", b"tokenizer-bytes-AAAA", "a");
let b = revision_of(b"model-bytes-BBBB", b"tokenizer-bytes-AAAA", "b");
assert_ne!(a, b, "a different MODEL must change the identity revision");
let c = revision_of(b"model-bytes-AAAA", b"tokenizer-bytes-CCCC", "c");
assert_ne!(a, c, "a different TOKENIZER must change the identity revision");
assert_ne!(b, c, "distinct model+tokenizer pairs must differ");
}
#[test]
fn identical_bytes_yield_stable_revision() {
let first = revision_of(b"same-model-bytes", b"same-tokenizer-bytes", "stable1");
let second = revision_of(b"same-model-bytes", b"same-tokenizer-bytes", "stable2");
assert_eq!(first, second, "identical asset bytes must yield a stable revision");
}
#[test]
fn revision_has_pinned_base_and_onnx_digest_shape() {
let rev = revision_of(b"m", b"t", "shape");
let prefix = format!("{HF_REVISION}+onnx-");
let hex = rev.strip_prefix(&prefix).expect("revision must start with pinned base");
assert_eq!(hex.len(), 12, "digest suffix must be 12 hex chars, got {hex:?}");
assert!(
hex.chars().all(|c| c.is_ascii_hexdigit()),
"digest suffix must be lower-hex, got {hex:?}"
);
}
#[test]
fn compose_is_pure_and_order_sensitive() {
let m = [1_u8; 32];
let t = [2_u8; 32];
assert_eq!(compose_asset_revision(&m, &t), compose_asset_revision(&m, &t));
assert_ne!(
compose_asset_revision(&m, &t),
compose_asset_revision(&t, &m),
"the two asset roles must not be interchangeable"
);
}
}
#[test]
fn ort_bge_embeds_384_dim_finite_deterministic_vector() {
use std::path::Path;
use fathomdb_embedder_api::Embedder;
use super::{OrtBgeEmbedder, OrtProvider};
let (Ok(_dylib), Ok(model), Ok(tok)) = (
std::env::var("ORT_DYLIB_PATH"),
std::env::var("FATHOMDB_ONNX_MODEL_PATH"),
std::env::var("FATHOMDB_ONNX_TOKENIZER_PATH"),
) else {
eprintln!(
"SKIP ort_bge_embeds_384_dim_finite_deterministic_vector: set ORT_DYLIB_PATH + \
FATHOMDB_ONNX_MODEL_PATH + FATHOMDB_ONNX_TOKENIZER_PATH to run the real-vector \
R-ONNX-1 test (see dev/tools/onnx/README.md)"
);
return;
};
let embedder = OrtBgeEmbedder::from_files_with_provider(
Path::new(&model),
Path::new(&tok),
OrtProvider::Cpu,
)
.expect("open OrtBgeEmbedder on the CPU EP from the provisioned asset");
let id = embedder.identity();
assert_eq!(id.name, super::ORT_BGE_EMBEDDER_NAME, "distinct ONNX identity name");
assert_eq!(id.dimension, super::ORT_BGE_EMBEDDER_DIM, "384-dim identity");
assert!(
id.revision.contains("+onnx-"),
"revision must carry the asset digest, got {:?}",
id.revision
);
let v1 = embedder.embed("the quick brown fox").expect("embed");
let v2 = embedder.embed("the quick brown fox").expect("embed");
assert_eq!(v1.len(), super::ORT_BGE_EMBEDDER_DIM as usize);
assert!(v1.iter().all(|x| x.is_finite()), "all components finite");
assert_eq!(v1, v2, "deterministic for identical input");
let norm = v1.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-3, "L2-normalized (got {norm})");
}
}