use crate::command::{Arguments, Context, Failed, Outcome};
pub const DEFAULT_COMMIT_EVERY: i64 = 1024;
pub const MAX_NAMED_TRUNCATED: usize = 20;
#[cfg(not(feature = "embed"))]
pub fn embed_table(_context: &mut Context, _arguments: &Arguments) -> Result<Outcome, Failed> {
Err(Failed::unsupported(
"embed",
"embed: this build has no embedding support compiled in",
))
}
#[cfg(feature = "embed")]
pub fn embed_table(context: &mut Context, arguments: &Arguments) -> Result<Outcome, Failed> {
real::embed_table(context, arguments)
}
#[cfg(feature = "embed")]
mod real {
use std::io::IsTerminal;
use std::time::Instant;
use inillucent_core::embed_onnx::{Device, Embedded, OnnxEmbedder, OnnxOptions};
use inillucent_core::install;
use inillucent_core::model::ModelManifest;
use inillucent_driver::Status;
use inillucent_engine::connect::Connection;
use inillucent_engine::OwnedDatum;
use super::{DEFAULT_COMMIT_EVERY, MAX_NAMED_TRUNCATED};
use crate::command::{Arguments, Context, Failed, Outcome};
use crate::json::{self, Json};
struct Plan {
table: String,
text_column: String,
vector_column: String,
prefix: String,
device: Device,
threads: Option<usize>,
sessions: usize,
commit_every: usize,
batch_size: usize,
all: bool,
}
#[derive(Default)]
struct Report {
embedded: usize,
skipped: usize,
truncated: usize,
truncated_rowids: Vec<i64>,
}
struct Row {
rowid: i64,
text: Option<String>,
}
pub fn embed_table(context: &mut Context, arguments: &Arguments) -> Result<Outcome, Failed> {
let plan = read_plan(arguments)?;
let connection = context.shell().connection();
check_target(&connection, &plan)?;
let total = count_rows(&connection, &plan)?;
let started = Instant::now();
let embedders = open_embedders(&plan)?;
let loaded = started.elapsed();
let max_tokens = embedders
.first()
.map(|embedder| embedder.options().max_tokens)
.unwrap_or(0);
let began = Instant::now();
let report = embed_all(&connection, &plan, &embedders, max_tokens, total);
park(&session_key(&plan), embedders);
Ok(outcome(&plan, &report?, began.elapsed(), loaded))
}
fn read_plan(arguments: &Arguments) -> Result<Plan, Failed> {
let device_text = match arguments.text("device") {
Some(text) => install::parse_device(text).map_err(Failed::misuse)?,
None => install::configured_device().value,
};
let device =
Device::parse(&device_text).map_err(|reason| Failed::misuse(format!("{reason:#}")))?;
let threads = match arguments.integer("threads") {
Some(count) => {
Some(install::parse_threads(&count.to_string()).map_err(Failed::misuse)?)
}
None => install::configured_threads().value,
};
Ok(Plan {
table: arguments.required_text("table")?.to_string(),
text_column: arguments.required_text("text")?.to_string(),
vector_column: arguments.required_text("vector")?.to_string(),
prefix: arguments.text("prefix").unwrap_or("").to_string(),
device,
threads,
sessions: positive(arguments, "sessions", 1)?,
commit_every: positive(arguments, "commit-every", DEFAULT_COMMIT_EVERY)?,
batch_size: positive(arguments, "batch-size", 16)?,
all: arguments.flag("all"),
})
}
fn positive(arguments: &Arguments, name: &str, default: i64) -> Result<usize, Failed> {
let value = arguments.integer(name).unwrap_or(default);
match usize::try_from(value) {
Ok(count) if count >= 1 => Ok(count),
_ => Err(Failed::misuse(format!(
"embed: {name} must be a whole number of at least 1, not {value}"
))),
}
}
static PARKED: std::sync::Mutex<Vec<(String, OnnxEmbedder)>> =
std::sync::Mutex::new(Vec::new());
fn session_key(plan: &Plan) -> String {
format!(
"{}|{:?}|{}",
plan.device.label(),
plan.threads,
plan.batch_size
)
}
fn take_parked(key: &str) -> Vec<OnnxEmbedder> {
let Ok(mut parked) = PARKED.lock() else {
return Vec::new();
};
let (mine, others): (Vec<_>, Vec<_>) = parked.drain(..).partition(|(held, _)| held == key);
*parked = others;
mine.into_iter().map(|(_, embedder)| embedder).collect()
}
fn park(key: &str, embedders: Vec<OnnxEmbedder>) {
match PARKED.lock() {
Ok(mut parked) => parked.extend(
embedders
.into_iter()
.map(|embedder| (key.to_string(), embedder)),
),
Err(_) => embedders.into_iter().for_each(std::mem::forget),
}
}
fn open_embedders(plan: &Plan) -> Result<Vec<OnnxEmbedder>, Failed> {
let Some(dir) = install::model_dir(install::DEFAULT_MODEL) else {
return Err(Failed::misuse(format!(
"embed: no embedding model is installed. Run `inillucent setup-embeddings all` to \
download {} and the ONNX Runtime it needs",
install::DEFAULT_MODEL
)));
};
let manifest = ModelManifest::read(&dir).unwrap_or_else(|_| ModelManifest::nomic_v1_5());
let mut options = OnnxOptions::for_model_on(&manifest, plan.batch_size, plan.device);
options.intra_threads = plan.threads;
let key = session_key(plan);
let mut opened = take_parked(&key);
if opened.len() > plan.sessions {
park(&key, opened.split_off(plan.sessions));
}
while opened.len() < plan.sessions {
let embedder = OnnxEmbedder::open_model(&dir, &manifest.model_file, options.clone())
.map_err(|reason| Failed::misuse(format!("embed: {reason:#}")))?;
opened.push(embedder);
}
Ok(opened)
}
fn quoted(name: &str) -> String {
format!("\"{}\"", name.replace('"', "\"\""))
}
fn check_target(connection: &Connection<'_>, plan: &Plan) -> Result<(), Failed> {
let sql = format!(
"SELECT rowid, {}, {} FROM {} LIMIT 0",
quoted(&plan.text_column),
quoted(&plan.vector_column),
quoted(&plan.table)
);
connection
.query(&sql)
.map(|_| ())
.map_err(|error| Failed::from_engine(&error))
}
fn wanted(plan: &Plan) -> String {
match plan.all {
true => "1".to_string(),
false => format!("{} IS NULL", quoted(&plan.vector_column)),
}
}
fn count_rows(connection: &Connection<'_>, plan: &Plan) -> Result<usize, Failed> {
let sql = format!(
"SELECT count(*) FROM {} WHERE {}",
quoted(&plan.table),
wanted(plan)
);
let rows = connection
.query(&sql)
.map_err(|error| Failed::from_engine(&error))?;
let counted = rows
.first()
.and_then(|row| row.first())
.and_then(|value| match value {
OwnedDatum::Int(count) => Some(*count),
_ => None,
});
Ok(counted
.and_then(|count| usize::try_from(count).ok())
.unwrap_or(0))
}
fn embed_all(
connection: &Connection<'_>,
plan: &Plan,
embedders: &[OnnxEmbedder],
max_tokens: usize,
total: usize,
) -> Result<Report, Failed> {
let mut report = Report::default();
let mut after = i64::MIN;
let started = Instant::now();
loop {
let slice = read_slice(connection, plan, after)?;
let Some(last) = slice.last() else {
break;
};
after = last.rowid;
let vectors = embed_slice(embedders, plan, &slice)?;
write_slice(connection, plan, &slice, &vectors, max_tokens, &mut report)?;
progress(&report, total, started.elapsed());
}
Ok(report)
}
fn read_slice(
connection: &Connection<'_>,
plan: &Plan,
after: i64,
) -> Result<Vec<Row>, Failed> {
let sql = format!(
"SELECT rowid, {} FROM {} WHERE {} AND rowid > ?1 ORDER BY rowid LIMIT {}",
quoted(&plan.text_column),
quoted(&plan.table),
wanted(plan),
plan.commit_every
);
let mut statement = connection
.prepare(&sql)
.map_err(|error| Failed::from_engine(&error))?;
statement
.bind_integer(1, after)
.map_err(|error| Failed::from_engine(&error))?;
let mut rows = Vec::new();
while statement
.step()
.map_err(|error| Failed::from_engine(&error))?
{
let row = statement.row();
let Some(OwnedDatum::Int(rowid)) = row.first() else {
continue;
};
rows.push(Row {
rowid: *rowid,
text: row.get(1).and_then(text_of),
});
}
Ok(rows)
}
fn text_of(value: &OwnedDatum) -> Option<String> {
match value {
OwnedDatum::Text(bytes) | OwnedDatum::Blob(bytes) if !bytes.is_empty() => {
Some(String::from_utf8_lossy(bytes).into_owned())
}
OwnedDatum::Int(number) => Some(number.to_string()),
OwnedDatum::Real(number) => Some(number.to_string()),
_ => None,
}
}
fn embed_slice(
embedders: &[OnnxEmbedder],
plan: &Plan,
slice: &[Row],
) -> Result<Vec<Option<(Vec<f32>, usize)>>, Failed> {
let mut order: Vec<usize> = (0..slice.len())
.filter(|at| slice.get(*at).is_some_and(|row| row.text.is_some()))
.collect();
order.sort_by_key(|at| {
slice
.get(*at)
.and_then(|row| row.text.as_ref())
.map_or(0, String::len)
});
let shares = deal(&order, embedders.len());
let results = run_shares(embedders, plan, slice, &shares)?;
let mut vectors: Vec<Option<(Vec<f32>, usize)>> = vec![None; slice.len()];
for (share, embedded) in shares.iter().zip(results) {
for ((at, vector), tokens) in share.iter().zip(embedded.vectors).zip(embedded.tokens) {
if let Some(slot) = vectors.get_mut(*at) {
*slot = Some((vector, tokens));
}
}
}
Ok(vectors)
}
fn deal(order: &[usize], sessions: usize) -> Vec<Vec<usize>> {
let mut shares: Vec<Vec<usize>> = vec![Vec::new(); sessions.max(1)];
for (turn, at) in order.iter().enumerate() {
if let Some(share) = shares.get_mut(turn % sessions.max(1)) {
share.push(*at);
}
}
shares
}
fn embed_share(
embedder: &OnnxEmbedder,
plan: &Plan,
slice: &[Row],
share: &[usize],
) -> Result<Embedded, Failed> {
let texts: Vec<String> = share
.iter()
.filter_map(|at| slice.get(*at).and_then(|row| row.text.as_ref()))
.map(|text| format!("{}{text}", plan.prefix))
.collect();
embedder
.embed_prefixed_counted(&texts)
.map_err(|reason| Failed::said(Status::InvalidState, format!("embed: {reason:#}")))
}
fn run_shares(
embedders: &[OnnxEmbedder],
plan: &Plan,
slice: &[Row],
shares: &[Vec<usize>],
) -> Result<Vec<Embedded>, Failed> {
if let ([embedder], [share]) = (embedders, shares) {
return Ok(vec![embed_share(embedder, plan, slice, share)?]);
}
let outcomes: Vec<Result<Embedded, String>> = std::thread::scope(|scope| {
let handles: Vec<_> = embedders
.iter()
.zip(shares)
.map(|(embedder, share)| {
let texts: Vec<String> = share
.iter()
.filter_map(|at| slice.get(*at).and_then(|row| row.text.as_ref()))
.map(|text| format!("{}{text}", plan.prefix))
.collect();
scope.spawn(move || {
embedder
.embed_prefixed_counted(&texts)
.map_err(|reason| format!("{reason:#}"))
})
})
.collect();
handles
.into_iter()
.map(|handle| {
handle.join().unwrap_or_else(|_| {
Err("an embedding thread stopped unexpectedly".to_string())
})
})
.collect()
});
outcomes
.into_iter()
.map(|outcome| {
outcome.map_err(|reason| {
Failed::said(Status::InvalidState, format!("embed: {reason}"))
})
})
.collect()
}
fn write_slice(
connection: &Connection<'_>,
plan: &Plan,
slice: &[Row],
vectors: &[Option<(Vec<f32>, usize)>],
max_tokens: usize,
report: &mut Report,
) -> Result<(), Failed> {
let transaction = connection
.begin()
.map_err(|error| Failed::from_engine(&error))?;
let update = format!(
"UPDATE {} SET {} = ?1 WHERE rowid = ?2",
quoted(&plan.table),
quoted(&plan.vector_column)
);
let mut statement = transaction
.prepare(&update)
.map_err(|error| Failed::from_engine(&error))?;
for (row, vector) in slice.iter().zip(vectors) {
let Some((values, tokens)) = vector else {
report.skipped = report.skipped.saturating_add(1);
continue;
};
statement
.bind_blob(1, &as_bytes(values))
.and_then(|_| statement.bind_integer(2, row.rowid))
.and_then(|_| statement.step())
.map_err(|error| Failed::from_engine(&error))?;
statement.reset();
report.embedded = report.embedded.saturating_add(1);
if *tokens > max_tokens {
report.truncated = report.truncated.saturating_add(1);
if report.truncated_rowids.len() < MAX_NAMED_TRUNCATED {
report.truncated_rowids.push(row.rowid);
}
}
}
drop(statement);
transaction
.commit()
.map_err(|error| Failed::from_engine(&error))
}
fn as_bytes(values: &[f32]) -> Vec<u8> {
let mut bytes = Vec::with_capacity(values.len().saturating_mul(4));
for value in values {
bytes.extend_from_slice(&value.to_bits().to_le_bytes());
}
bytes
}
fn progress(report: &Report, total: usize, elapsed: std::time::Duration) {
if !std::io::stderr().is_terminal() {
return;
}
let done = report.embedded.saturating_add(report.skipped);
let rate = report.embedded as f64 / elapsed.as_secs_f64().max(0.001);
eprintln!("embed: {done} of {total} rows, {rate:.0} rows a second");
}
fn outcome(
plan: &Plan,
report: &Report,
elapsed: std::time::Duration,
loaded: std::time::Duration,
) -> Outcome {
let seconds = elapsed.as_secs_f64();
let rate = report.embedded as f64 / seconds.max(0.001);
let mut text = format!(
"embedded {} rows in {seconds:.1} s, {rate:.0} rows a second, on {} with {} session{}. \
{} skipped because the text was NULL or empty. {} cut at the model's token limit",
report.embedded,
plan.device.label(),
plan.sessions,
if plan.sessions == 1 { "" } else { "s" },
report.skipped,
report.truncated
);
if !report.truncated_rowids.is_empty() {
let named: Vec<String> = report.truncated_rowids.iter().map(i64::to_string).collect();
text.push_str(&format!(". First cut rowids: {}", named.join(", ")));
}
let mut said = Outcome::said("embed", text);
said.changes = report.embedded as i64;
said.with("embedded", Json::Int(report.embedded as i64))
.with("skipped", Json::Int(report.skipped as i64))
.with("truncated", Json::Int(report.truncated as i64))
.with(
"truncated_rowids",
Json::Array(
report
.truncated_rowids
.iter()
.map(|id| Json::Int(*id))
.collect(),
),
)
.with("seconds", Json::Real(seconds))
.with("rows_per_second", Json::Real(rate))
.with("load_seconds", Json::Real(loaded.as_secs_f64()))
.with("device", json::text(plan.device.label()))
.with("sessions", Json::Int(plan.sessions as i64))
}
}