pub mod acquire;
pub mod diag;
pub mod error;
pub mod jina;
pub mod policy;
pub mod probe;
pub mod providers;
pub use acquire::{fetch_tokenizer_file, parse_owner_name};
pub use diag::{report, OrtReport};
pub use jina::{JinaV5, TokenizerHandle, MAX_LENGTH, MODEL_MAX_TOKENS, NATIVE_DIM};
pub use policy::{DeviceReq, SessionPolicy};
pub use probe::cuda_available;
pub use providers::{apply_providers, map_provider, ProviderMapping};
#[cfg(feature = "extension-module")]
use pyo3::prelude::*;
#[cfg(feature = "extension-module")]
use std::path::PathBuf;
#[cfg(feature = "extension-module")]
#[pyclass(name = "JinaV5")]
struct PyJinaV5 {
inner: JinaV5,
}
#[cfg(feature = "extension-module")]
#[pymethods]
impl PyJinaV5 {
#[staticmethod]
#[pyo3(signature = (model_id, revision=None, cache_dir=None, truncate_dim=512, device="auto", max_length=None))]
fn open(
model_id: &str,
revision: Option<String>,
cache_dir: Option<String>,
truncate_dim: usize,
device: &str,
max_length: Option<usize>,
) -> PyResult<Self> {
let inner = JinaV5::open(
model_id,
revision.as_deref(),
cache_dir.map(PathBuf::from),
truncate_dim,
DeviceReq::parse(device)
.map_err(|e| pyo3::exceptions::PyValueError::new_err(e.to_string()))?,
max_length,
)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))?;
Ok(Self { inner })
}
#[staticmethod]
#[pyo3(signature = (onnx_path, tokenizer_path, truncate_dim=512, device="auto", max_length=None))]
fn open_files(
onnx_path: &str,
tokenizer_path: &str,
truncate_dim: usize,
device: &str,
max_length: Option<usize>,
) -> PyResult<Self> {
let inner = JinaV5::open_files(
std::path::Path::new(onnx_path),
std::path::Path::new(tokenizer_path),
truncate_dim,
DeviceReq::parse(device)
.map_err(|e| pyo3::exceptions::PyValueError::new_err(e.to_string()))?,
max_length,
)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))?;
Ok(Self { inner })
}
#[pyo3(signature = (text, task="Document"))]
fn encode(&self, py: Python<'_>, text: &str, task: &str) -> PyResult<Vec<f32>> {
py.detach(|| {
self.inner
.encode_one(text, task)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))
})
}
#[pyo3(signature = (texts, task="Document"))]
fn encode_batch(
&self,
py: Python<'_>,
texts: Vec<String>,
task: &str,
) -> PyResult<Vec<Vec<f32>>> {
py.detach(|| {
self.inner
.encode_many(&texts, task)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))
})
}
#[getter]
fn dim(&self) -> usize {
self.inner.dim()
}
#[getter]
fn model_id(&self) -> &str {
self.inner.model_id()
}
#[getter]
fn used_cuda(&self) -> bool {
self.inner.used_cuda()
}
#[getter]
fn max_length(&self) -> usize {
self.inner.max_len()
}
fn count_tokens(&self, text: &str) -> PyResult<usize> {
self.inner
.count_tokens(text)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))
}
}
#[cfg(feature = "extension-module")]
#[pyclass(name = "JinaTokenizer")]
struct PyJinaTokenizer {
inner: TokenizerHandle,
}
#[cfg(feature = "extension-module")]
#[pymethods]
impl PyJinaTokenizer {
#[staticmethod]
#[pyo3(signature = (model_id, revision=None, cache_dir=None))]
fn open(
model_id: &str,
revision: Option<String>,
cache_dir: Option<String>,
) -> PyResult<Self> {
let inner = TokenizerHandle::open(
model_id,
revision.as_deref(),
cache_dir.map(PathBuf::from),
)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))?;
Ok(Self { inner })
}
#[staticmethod]
#[pyo3(signature = (tokenizer_path))]
fn open_files(tokenizer_path: &str) -> PyResult<Self> {
let inner = TokenizerHandle::open_files(std::path::Path::new(tokenizer_path))
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))?;
Ok(Self { inner })
}
fn count_tokens(&self, text: &str) -> PyResult<usize> {
self.inner
.count_tokens(text)
.map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(format!("{e:#}")))
}
}
#[cfg(feature = "extension-module")]
#[pymodule]
fn embroider(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<PyJinaV5>()?;
m.add_class::<PyJinaTokenizer>()?;
m.add("NATIVE_DIM", NATIVE_DIM)?;
m.add("MAX_LENGTH", MAX_LENGTH)?;
m.add("MODEL_MAX_TOKENS", MODEL_MAX_TOKENS)?;
Ok(())
}