1use crate::command::{Arguments, Context, Failed, Outcome};
22
23pub const DEFAULT_COMMIT_EVERY: i64 = 1024;
25
26pub const MAX_NAMED_TRUNCATED: usize = 20;
28
29#[cfg(not(feature = "embed"))]
34pub fn embed_table(_context: &mut Context, _arguments: &Arguments) -> Result<Outcome, Failed> {
35 Err(Failed::unsupported(
36 "embed",
37 "embed: this build has no embedding support compiled in",
38 ))
39}
40
41#[cfg(feature = "embed")]
46pub fn embed_table(context: &mut Context, arguments: &Arguments) -> Result<Outcome, Failed> {
47 real::embed_table(context, arguments)
48}
49
50#[cfg(feature = "embed")]
51mod real {
52 use std::io::IsTerminal;
53 use std::time::Instant;
54
55 use inillucent_core::embed_onnx::{Device, Embedded, OnnxEmbedder, OnnxOptions};
56 use inillucent_core::install;
57 use inillucent_core::model::ModelManifest;
58 use inillucent_driver::Status;
59 use inillucent_engine::connect::Connection;
60 use inillucent_engine::OwnedDatum;
61
62 use super::{DEFAULT_COMMIT_EVERY, MAX_NAMED_TRUNCATED};
63 use crate::command::{Arguments, Context, Failed, Outcome};
64 use crate::json::{self, Json};
65
66 struct Plan {
68 table: String,
69 text_column: String,
70 vector_column: String,
71 prefix: String,
72 device: Device,
73 threads: Option<usize>,
74 sessions: usize,
75 commit_every: usize,
76 batch_size: usize,
77 all: bool,
78 }
79
80 #[derive(Default)]
82 struct Report {
83 embedded: usize,
84 skipped: usize,
85 truncated: usize,
86 truncated_rowids: Vec<i64>,
87 }
88
89 struct Row {
91 rowid: i64,
92 text: Option<String>,
93 }
94
95 pub fn embed_table(context: &mut Context, arguments: &Arguments) -> Result<Outcome, Failed> {
100 let plan = read_plan(arguments)?;
101 let connection = context.shell().connection();
102 check_target(&connection, &plan)?;
105 let total = count_rows(&connection, &plan)?;
106 let started = Instant::now();
107 let embedders = open_embedders(&plan)?;
108 let loaded = started.elapsed();
109 let max_tokens = embedders
110 .first()
111 .map(|embedder| embedder.options().max_tokens)
112 .unwrap_or(0);
113 let began = Instant::now();
114 let report = embed_all(&connection, &plan, &embedders, max_tokens, total);
115 park(&session_key(&plan), embedders);
116 Ok(outcome(&plan, &report?, began.elapsed(), loaded))
117 }
118
119 fn read_plan(arguments: &Arguments) -> Result<Plan, Failed> {
126 let device_text = match arguments.text("device") {
127 Some(text) => install::parse_device(text).map_err(Failed::misuse)?,
128 None => install::configured_device().value,
129 };
130 let device =
131 Device::parse(&device_text).map_err(|reason| Failed::misuse(format!("{reason:#}")))?;
132 let threads = match arguments.integer("threads") {
133 Some(count) => {
134 Some(install::parse_threads(&count.to_string()).map_err(Failed::misuse)?)
135 }
136 None => install::configured_threads().value,
137 };
138 Ok(Plan {
139 table: arguments.required_text("table")?.to_string(),
140 text_column: arguments.required_text("text")?.to_string(),
141 vector_column: arguments.required_text("vector")?.to_string(),
142 prefix: arguments.text("prefix").unwrap_or("").to_string(),
143 device,
144 threads,
145 sessions: positive(arguments, "sessions", 1)?,
146 commit_every: positive(arguments, "commit-every", DEFAULT_COMMIT_EVERY)?,
147 batch_size: positive(arguments, "batch-size", 16)?,
148 all: arguments.flag("all"),
149 })
150 }
151
152 fn positive(arguments: &Arguments, name: &str, default: i64) -> Result<usize, Failed> {
158 let value = arguments.integer(name).unwrap_or(default);
159 match usize::try_from(value) {
160 Ok(count) if count >= 1 => Ok(count),
161 _ => Err(Failed::misuse(format!(
162 "embed: {name} must be a whole number of at least 1, not {value}"
163 ))),
164 }
165 }
166
167 static PARKED: std::sync::Mutex<Vec<(String, OnnxEmbedder)>> =
178 std::sync::Mutex::new(Vec::new());
179
180 fn session_key(plan: &Plan) -> String {
184 format!(
185 "{}|{:?}|{}",
186 plan.device.label(),
187 plan.threads,
188 plan.batch_size
189 )
190 }
191
192 fn take_parked(key: &str) -> Vec<OnnxEmbedder> {
196 let Ok(mut parked) = PARKED.lock() else {
197 return Vec::new();
198 };
199 let (mine, others): (Vec<_>, Vec<_>) = parked.drain(..).partition(|(held, _)| held == key);
200 *parked = others;
201 mine.into_iter().map(|(_, embedder)| embedder).collect()
202 }
203
204 fn park(key: &str, embedders: Vec<OnnxEmbedder>) {
209 static EXIT_REGISTERED: std::sync::OnceLock<()> = std::sync::OnceLock::new();
214 EXIT_REGISTERED.get_or_init(|| {
215 inillucent_core::exit_hook::register(Box::new(|| {
216 if let Ok(mut parked) = PARKED.try_lock() {
217 parked.clear();
218 }
219 }));
220 });
221 match PARKED.lock() {
222 Ok(mut parked) => parked.extend(
223 embedders
224 .into_iter()
225 .map(|embedder| (key.to_string(), embedder)),
226 ),
227 Err(_) => embedders.into_iter().for_each(std::mem::forget),
228 }
229 }
230
231 fn open_embedders(plan: &Plan) -> Result<Vec<OnnxEmbedder>, Failed> {
238 let Some(dir) = install::model_dir(install::DEFAULT_MODEL) else {
239 return Err(Failed::misuse(format!(
240 "embed: no embedding model is installed. Run `inillucent setup-embeddings all` to \
241 download {} and the ONNX Runtime it needs",
242 install::DEFAULT_MODEL
243 )));
244 };
245 let manifest = ModelManifest::read(&dir).unwrap_or_else(|_| ModelManifest::nomic_v1_5());
246 let mut options = OnnxOptions::for_model_on(&manifest, plan.batch_size, plan.device);
247 options.intra_threads = plan.threads;
248 let key = session_key(plan);
249 let mut opened = take_parked(&key);
250 if opened.len() > plan.sessions {
251 park(&key, opened.split_off(plan.sessions));
252 }
253 while opened.len() < plan.sessions {
254 let embedder = OnnxEmbedder::open_model(&dir, &manifest.model_file, options.clone())
255 .map_err(|reason| Failed::misuse(format!("embed: {reason:#}")))?;
256 opened.push(embedder);
257 }
258 Ok(opened)
259 }
260
261 fn quoted(name: &str) -> String {
265 format!("\"{}\"", name.replace('"', "\"\""))
266 }
267
268 fn check_target(connection: &Connection<'_>, plan: &Plan) -> Result<(), Failed> {
273 let sql = format!(
274 "SELECT rowid, {}, {} FROM {} LIMIT 0",
275 quoted(&plan.text_column),
276 quoted(&plan.vector_column),
277 quoted(&plan.table)
278 );
279 connection
280 .query(&sql)
281 .map(|_| ())
282 .map_err(|error| Failed::from_engine(&error))
283 }
284
285 fn wanted(plan: &Plan) -> String {
292 match plan.all {
293 true => "1".to_string(),
294 false => format!("{} IS NULL", quoted(&plan.vector_column)),
295 }
296 }
297
298 fn count_rows(connection: &Connection<'_>, plan: &Plan) -> Result<usize, Failed> {
303 let sql = format!(
304 "SELECT count(*) FROM {} WHERE {}",
305 quoted(&plan.table),
306 wanted(plan)
307 );
308 let rows = connection
309 .query(&sql)
310 .map_err(|error| Failed::from_engine(&error))?;
311 let counted = rows
312 .first()
313 .and_then(|row| row.first())
314 .and_then(|value| match value {
315 OwnedDatum::Int(count) => Some(*count),
316 _ => None,
317 });
318 Ok(counted
319 .and_then(|count| usize::try_from(count).ok())
320 .unwrap_or(0))
321 }
322
323 fn embed_all(
331 connection: &Connection<'_>,
332 plan: &Plan,
333 embedders: &[OnnxEmbedder],
334 max_tokens: usize,
335 total: usize,
336 ) -> Result<Report, Failed> {
337 let mut report = Report::default();
338 let mut after = i64::MIN;
339 let started = Instant::now();
340 loop {
341 let slice = read_slice(connection, plan, after)?;
342 let Some(last) = slice.last() else {
343 break;
344 };
345 after = last.rowid;
346 let vectors = embed_slice(embedders, plan, &slice)?;
347 write_slice(connection, plan, &slice, &vectors, max_tokens, &mut report)?;
348 progress(&report, total, started.elapsed());
349 }
350 Ok(report)
351 }
352
353 fn read_slice(
359 connection: &Connection<'_>,
360 plan: &Plan,
361 after: i64,
362 ) -> Result<Vec<Row>, Failed> {
363 let sql = format!(
364 "SELECT rowid, {} FROM {} WHERE {} AND rowid > ?1 ORDER BY rowid LIMIT {}",
365 quoted(&plan.text_column),
366 quoted(&plan.table),
367 wanted(plan),
368 plan.commit_every
369 );
370 let mut statement = connection
371 .prepare(&sql)
372 .map_err(|error| Failed::from_engine(&error))?;
373 statement
374 .bind_integer(1, after)
375 .map_err(|error| Failed::from_engine(&error))?;
376 let mut rows = Vec::new();
377 while statement
378 .step()
379 .map_err(|error| Failed::from_engine(&error))?
380 {
381 let row = statement.row();
382 let Some(OwnedDatum::Int(rowid)) = row.first() else {
383 continue;
384 };
385 rows.push(Row {
386 rowid: *rowid,
387 text: row.get(1).and_then(text_of),
388 });
389 }
390 Ok(rows)
391 }
392
393 fn text_of(value: &OwnedDatum) -> Option<String> {
397 match value {
398 OwnedDatum::Text(bytes) | OwnedDatum::Blob(bytes) if !bytes.is_empty() => {
399 Some(String::from_utf8_lossy(bytes).into_owned())
400 }
401 OwnedDatum::Int(number) => Some(number.to_string()),
402 OwnedDatum::Real(number) => Some(number.to_string()),
403 _ => None,
404 }
405 }
406
407 fn embed_slice(
418 embedders: &[OnnxEmbedder],
419 plan: &Plan,
420 slice: &[Row],
421 ) -> Result<Vec<Option<(Vec<f32>, usize)>>, Failed> {
422 let mut order: Vec<usize> = (0..slice.len())
423 .filter(|at| slice.get(*at).is_some_and(|row| row.text.is_some()))
424 .collect();
425 order.sort_by_key(|at| {
426 slice
427 .get(*at)
428 .and_then(|row| row.text.as_ref())
429 .map_or(0, String::len)
430 });
431 let shares = deal(&order, embedders.len());
432 let results = run_shares(embedders, plan, slice, &shares)?;
433 let mut vectors: Vec<Option<(Vec<f32>, usize)>> = vec![None; slice.len()];
434 for (share, embedded) in shares.iter().zip(results) {
435 for ((at, vector), tokens) in share.iter().zip(embedded.vectors).zip(embedded.tokens) {
436 if let Some(slot) = vectors.get_mut(*at) {
437 *slot = Some((vector, tokens));
438 }
439 }
440 }
441 Ok(vectors)
442 }
443
444 fn deal(order: &[usize], sessions: usize) -> Vec<Vec<usize>> {
449 let mut shares: Vec<Vec<usize>> = vec![Vec::new(); sessions.max(1)];
450 for (turn, at) in order.iter().enumerate() {
451 if let Some(share) = shares.get_mut(turn % sessions.max(1)) {
452 share.push(*at);
453 }
454 }
455 shares
456 }
457
458 fn embed_share(
465 embedder: &OnnxEmbedder,
466 plan: &Plan,
467 slice: &[Row],
468 share: &[usize],
469 ) -> Result<Embedded, Failed> {
470 let texts: Vec<String> = share
471 .iter()
472 .filter_map(|at| slice.get(*at).and_then(|row| row.text.as_ref()))
473 .map(|text| format!("{}{text}", plan.prefix))
474 .collect();
475 embedder
476 .embed_prefixed_counted(&texts)
477 .map_err(|reason| Failed::said(Status::InvalidState, format!("embed: {reason:#}")))
478 }
479
480 fn run_shares(
487 embedders: &[OnnxEmbedder],
488 plan: &Plan,
489 slice: &[Row],
490 shares: &[Vec<usize>],
491 ) -> Result<Vec<Embedded>, Failed> {
492 if let ([embedder], [share]) = (embedders, shares) {
493 return Ok(vec![embed_share(embedder, plan, slice, share)?]);
494 }
495 let outcomes: Vec<Result<Embedded, String>> = std::thread::scope(|scope| {
496 let handles: Vec<_> = embedders
497 .iter()
498 .zip(shares)
499 .map(|(embedder, share)| {
500 let texts: Vec<String> = share
501 .iter()
502 .filter_map(|at| slice.get(*at).and_then(|row| row.text.as_ref()))
503 .map(|text| format!("{}{text}", plan.prefix))
504 .collect();
505 scope.spawn(move || {
506 embedder
507 .embed_prefixed_counted(&texts)
508 .map_err(|reason| format!("{reason:#}"))
509 })
510 })
511 .collect();
512 handles
513 .into_iter()
514 .map(|handle| {
515 handle.join().unwrap_or_else(|_| {
516 Err("an embedding thread stopped unexpectedly".to_string())
517 })
518 })
519 .collect()
520 });
521 outcomes
522 .into_iter()
523 .map(|outcome| {
524 outcome.map_err(|reason| {
525 Failed::said(Status::InvalidState, format!("embed: {reason}"))
526 })
527 })
528 .collect()
529 }
530
531 fn write_slice(
543 connection: &Connection<'_>,
544 plan: &Plan,
545 slice: &[Row],
546 vectors: &[Option<(Vec<f32>, usize)>],
547 max_tokens: usize,
548 report: &mut Report,
549 ) -> Result<(), Failed> {
550 let transaction = connection
551 .begin()
552 .map_err(|error| Failed::from_engine(&error))?;
553 let update = format!(
554 "UPDATE {} SET {} = ?1 WHERE rowid = ?2",
555 quoted(&plan.table),
556 quoted(&plan.vector_column)
557 );
558 let mut statement = transaction
559 .prepare(&update)
560 .map_err(|error| Failed::from_engine(&error))?;
561 for (row, vector) in slice.iter().zip(vectors) {
562 let Some((values, tokens)) = vector else {
563 report.skipped = report.skipped.saturating_add(1);
564 continue;
565 };
566 statement
567 .bind_blob(1, &as_bytes(values))
568 .and_then(|_| statement.bind_integer(2, row.rowid))
569 .and_then(|_| statement.step())
570 .map_err(|error| Failed::from_engine(&error))?;
571 statement.reset();
572 report.embedded = report.embedded.saturating_add(1);
573 if *tokens > max_tokens {
574 report.truncated = report.truncated.saturating_add(1);
575 if report.truncated_rowids.len() < MAX_NAMED_TRUNCATED {
576 report.truncated_rowids.push(row.rowid);
577 }
578 }
579 }
580 drop(statement);
581 transaction
582 .commit()
583 .map_err(|error| Failed::from_engine(&error))
584 }
585
586 fn as_bytes(values: &[f32]) -> Vec<u8> {
590 let mut bytes = Vec::with_capacity(values.len().saturating_mul(4));
591 for value in values {
592 bytes.extend_from_slice(&value.to_bits().to_le_bytes());
593 }
594 bytes
595 }
596
597 fn progress(report: &Report, total: usize, elapsed: std::time::Duration) {
603 if !std::io::stderr().is_terminal() {
604 return;
605 }
606 let done = report.embedded.saturating_add(report.skipped);
607 let rate = report.embedded as f64 / elapsed.as_secs_f64().max(0.001);
608 eprintln!("embed: {done} of {total} rows, {rate:.0} rows a second");
609 }
610
611 fn outcome(
618 plan: &Plan,
619 report: &Report,
620 elapsed: std::time::Duration,
621 loaded: std::time::Duration,
622 ) -> Outcome {
623 let seconds = elapsed.as_secs_f64();
624 let rate = report.embedded as f64 / seconds.max(0.001);
625 let mut text = format!(
626 "embedded {} rows in {seconds:.1} s, {rate:.0} rows a second, on {} with {} session{}. \
627 {} skipped because the text was NULL or empty. {} cut at the model's token limit",
628 report.embedded,
629 plan.device.label(),
630 plan.sessions,
631 if plan.sessions == 1 { "" } else { "s" },
632 report.skipped,
633 report.truncated
634 );
635 if !report.truncated_rowids.is_empty() {
636 let named: Vec<String> = report.truncated_rowids.iter().map(i64::to_string).collect();
637 text.push_str(&format!(". First cut rowids: {}", named.join(", ")));
638 }
639 let mut said = Outcome::said("embed", text);
640 said.changes = report.embedded as i64;
641 said.with("embedded", Json::Int(report.embedded as i64))
642 .with("skipped", Json::Int(report.skipped as i64))
643 .with("truncated", Json::Int(report.truncated as i64))
644 .with(
645 "truncated_rowids",
646 Json::Array(
647 report
648 .truncated_rowids
649 .iter()
650 .map(|id| Json::Int(*id))
651 .collect(),
652 ),
653 )
654 .with("seconds", Json::Real(seconds))
655 .with("rows_per_second", Json::Real(rate))
656 .with("load_seconds", Json::Real(loaded.as_secs_f64()))
657 .with("device", json::text(plan.device.label()))
658 .with("sessions", Json::Int(plan.sessions as i64))
659 }
660}