pub mod tokens;
#[cfg(feature = "pretrained-embed")]
pub mod token_vocab;
pub use tokens::EMBEDDING_DIM;
#[cfg(feature = "pretrained-embed")]
pub use token_vocab::{PRETRAINED_DIM, embed_code_pretrained};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Backend {
Hashing,
#[allow(dead_code)]
Pretrained,
}
impl Backend {
pub const fn dims(self) -> usize {
match self {
Self::Hashing => EMBEDDING_DIM,
#[cfg(feature = "pretrained-embed")]
Self::Pretrained => PRETRAINED_DIM,
#[cfg(not(feature = "pretrained-embed"))]
Self::Pretrained => EMBEDDING_DIM, }
}
pub const fn provider_name(self) -> &'static str {
match self {
Self::Hashing => "hashing-trick (256d, static)",
#[cfg(feature = "pretrained-embed")]
Self::Pretrained => "nomic-embed-code (768d, pretrained)",
#[cfg(not(feature = "pretrained-embed"))]
Self::Pretrained => "hashing-trick (256d, pretrained not compiled)",
}
}
}
static ACTIVE_BACKEND: std::sync::OnceLock<Backend> = std::sync::OnceLock::new();
pub fn resolve_backend(config_embedding: &str) -> Backend {
if let Some(&b) = ACTIVE_BACKEND.get() {
return b;
}
let backend = match config_embedding {
"hashing" => Backend::Hashing,
"pretrained" => {
#[cfg(feature = "pretrained-embed")]
{
Backend::Pretrained
}
#[cfg(not(feature = "pretrained-embed"))]
{
tracing::warn!(
"brain.embedding=pretrained but cora was not compiled with --features pretrained-embed; \
falling back to hashing-trick 256d"
);
Backend::Hashing
}
}
_ => {
#[cfg(feature = "pretrained-embed")]
{
Backend::Pretrained
}
#[cfg(not(feature = "pretrained-embed"))]
{
Backend::Hashing
}
}
};
let _ = ACTIVE_BACKEND.set(backend);
tracing::debug!(
config = config_embedding,
backend = ?backend,
"resolved embedding backend"
);
backend
}
pub fn active_dims() -> usize {
ACTIVE_BACKEND.get().map(|b| b.dims()).unwrap_or_else(|| {
#[cfg(feature = "pretrained-embed")]
{
PRETRAINED_DIM
}
#[cfg(not(feature = "pretrained-embed"))]
{
EMBEDDING_DIM
}
})
}
pub fn active_provider_name() -> &'static str {
ACTIVE_BACKEND
.get()
.map(|b| b.provider_name())
.unwrap_or_else(|| {
#[cfg(feature = "pretrained-embed")]
{
"nomic-embed-code (768d, pretrained)"
}
#[cfg(not(feature = "pretrained-embed"))]
{
"hashing-trick (256d, static)"
}
})
}
pub fn embed_code_dispatch(code: &str) -> Vec<f32> {
let backend = ACTIVE_BACKEND.get().copied().unwrap_or_else(|| {
resolve_backend("auto")
});
match backend {
Backend::Hashing => {
let embedding = tokens::embed_code(code);
embedding.as_slice().iter().map(|&v| v as f32).collect()
}
#[cfg(feature = "pretrained-embed")]
Backend::Pretrained => embed_code_pretrained(code),
#[cfg(not(feature = "pretrained-embed"))]
Backend::Pretrained => {
let embedding = tokens::embed_code(code);
embedding.as_slice().iter().map(|&v| v as f32).collect()
}
}
}
#[allow(dead_code)]
pub const fn has_pretrained() -> bool {
cfg!(feature = "pretrained-embed")
}