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 match PARKED.lock() {
210 Ok(mut parked) => parked.extend(
211 embedders
212 .into_iter()
213 .map(|embedder| (key.to_string(), embedder)),
214 ),
215 Err(_) => embedders.into_iter().for_each(std::mem::forget),
216 }
217 }
218
219 fn open_embedders(plan: &Plan) -> Result<Vec<OnnxEmbedder>, Failed> {
226 let Some(dir) = install::model_dir(install::DEFAULT_MODEL) else {
227 return Err(Failed::misuse(format!(
228 "embed: no embedding model is installed. Run `inillucent setup-embeddings all` to \
229 download {} and the ONNX Runtime it needs",
230 install::DEFAULT_MODEL
231 )));
232 };
233 let manifest = ModelManifest::read(&dir).unwrap_or_else(|_| ModelManifest::nomic_v1_5());
234 let mut options = OnnxOptions::for_model_on(&manifest, plan.batch_size, plan.device);
235 options.intra_threads = plan.threads;
236 let key = session_key(plan);
237 let mut opened = take_parked(&key);
238 if opened.len() > plan.sessions {
239 park(&key, opened.split_off(plan.sessions));
240 }
241 while opened.len() < plan.sessions {
242 let embedder = OnnxEmbedder::open_model(&dir, &manifest.model_file, options.clone())
243 .map_err(|reason| Failed::misuse(format!("embed: {reason:#}")))?;
244 opened.push(embedder);
245 }
246 Ok(opened)
247 }
248
249 fn quoted(name: &str) -> String {
253 format!("\"{}\"", name.replace('"', "\"\""))
254 }
255
256 fn check_target(connection: &Connection<'_>, plan: &Plan) -> Result<(), Failed> {
261 let sql = format!(
262 "SELECT rowid, {}, {} FROM {} LIMIT 0",
263 quoted(&plan.text_column),
264 quoted(&plan.vector_column),
265 quoted(&plan.table)
266 );
267 connection
268 .query(&sql)
269 .map(|_| ())
270 .map_err(|error| Failed::from_engine(&error))
271 }
272
273 fn wanted(plan: &Plan) -> String {
280 match plan.all {
281 true => "1".to_string(),
282 false => format!("{} IS NULL", quoted(&plan.vector_column)),
283 }
284 }
285
286 fn count_rows(connection: &Connection<'_>, plan: &Plan) -> Result<usize, Failed> {
291 let sql = format!(
292 "SELECT count(*) FROM {} WHERE {}",
293 quoted(&plan.table),
294 wanted(plan)
295 );
296 let rows = connection
297 .query(&sql)
298 .map_err(|error| Failed::from_engine(&error))?;
299 let counted = rows
300 .first()
301 .and_then(|row| row.first())
302 .and_then(|value| match value {
303 OwnedDatum::Int(count) => Some(*count),
304 _ => None,
305 });
306 Ok(counted
307 .and_then(|count| usize::try_from(count).ok())
308 .unwrap_or(0))
309 }
310
311 fn embed_all(
319 connection: &Connection<'_>,
320 plan: &Plan,
321 embedders: &[OnnxEmbedder],
322 max_tokens: usize,
323 total: usize,
324 ) -> Result<Report, Failed> {
325 let mut report = Report::default();
326 let mut after = i64::MIN;
327 let started = Instant::now();
328 loop {
329 let slice = read_slice(connection, plan, after)?;
330 let Some(last) = slice.last() else {
331 break;
332 };
333 after = last.rowid;
334 let vectors = embed_slice(embedders, plan, &slice)?;
335 write_slice(connection, plan, &slice, &vectors, max_tokens, &mut report)?;
336 progress(&report, total, started.elapsed());
337 }
338 Ok(report)
339 }
340
341 fn read_slice(
347 connection: &Connection<'_>,
348 plan: &Plan,
349 after: i64,
350 ) -> Result<Vec<Row>, Failed> {
351 let sql = format!(
352 "SELECT rowid, {} FROM {} WHERE {} AND rowid > ?1 ORDER BY rowid LIMIT {}",
353 quoted(&plan.text_column),
354 quoted(&plan.table),
355 wanted(plan),
356 plan.commit_every
357 );
358 let mut statement = connection
359 .prepare(&sql)
360 .map_err(|error| Failed::from_engine(&error))?;
361 statement
362 .bind_integer(1, after)
363 .map_err(|error| Failed::from_engine(&error))?;
364 let mut rows = Vec::new();
365 while statement
366 .step()
367 .map_err(|error| Failed::from_engine(&error))?
368 {
369 let row = statement.row();
370 let Some(OwnedDatum::Int(rowid)) = row.first() else {
371 continue;
372 };
373 rows.push(Row {
374 rowid: *rowid,
375 text: row.get(1).and_then(text_of),
376 });
377 }
378 Ok(rows)
379 }
380
381 fn text_of(value: &OwnedDatum) -> Option<String> {
385 match value {
386 OwnedDatum::Text(bytes) | OwnedDatum::Blob(bytes) if !bytes.is_empty() => {
387 Some(String::from_utf8_lossy(bytes).into_owned())
388 }
389 OwnedDatum::Int(number) => Some(number.to_string()),
390 OwnedDatum::Real(number) => Some(number.to_string()),
391 _ => None,
392 }
393 }
394
395 fn embed_slice(
406 embedders: &[OnnxEmbedder],
407 plan: &Plan,
408 slice: &[Row],
409 ) -> Result<Vec<Option<(Vec<f32>, usize)>>, Failed> {
410 let mut order: Vec<usize> = (0..slice.len())
411 .filter(|at| slice.get(*at).is_some_and(|row| row.text.is_some()))
412 .collect();
413 order.sort_by_key(|at| {
414 slice
415 .get(*at)
416 .and_then(|row| row.text.as_ref())
417 .map_or(0, String::len)
418 });
419 let shares = deal(&order, embedders.len());
420 let results = run_shares(embedders, plan, slice, &shares)?;
421 let mut vectors: Vec<Option<(Vec<f32>, usize)>> = vec![None; slice.len()];
422 for (share, embedded) in shares.iter().zip(results) {
423 for ((at, vector), tokens) in share.iter().zip(embedded.vectors).zip(embedded.tokens) {
424 if let Some(slot) = vectors.get_mut(*at) {
425 *slot = Some((vector, tokens));
426 }
427 }
428 }
429 Ok(vectors)
430 }
431
432 fn deal(order: &[usize], sessions: usize) -> Vec<Vec<usize>> {
437 let mut shares: Vec<Vec<usize>> = vec![Vec::new(); sessions.max(1)];
438 for (turn, at) in order.iter().enumerate() {
439 if let Some(share) = shares.get_mut(turn % sessions.max(1)) {
440 share.push(*at);
441 }
442 }
443 shares
444 }
445
446 fn embed_share(
453 embedder: &OnnxEmbedder,
454 plan: &Plan,
455 slice: &[Row],
456 share: &[usize],
457 ) -> Result<Embedded, Failed> {
458 let texts: Vec<String> = share
459 .iter()
460 .filter_map(|at| slice.get(*at).and_then(|row| row.text.as_ref()))
461 .map(|text| format!("{}{text}", plan.prefix))
462 .collect();
463 embedder
464 .embed_prefixed_counted(&texts)
465 .map_err(|reason| Failed::said(Status::InvalidState, format!("embed: {reason:#}")))
466 }
467
468 fn run_shares(
475 embedders: &[OnnxEmbedder],
476 plan: &Plan,
477 slice: &[Row],
478 shares: &[Vec<usize>],
479 ) -> Result<Vec<Embedded>, Failed> {
480 if let ([embedder], [share]) = (embedders, shares) {
481 return Ok(vec![embed_share(embedder, plan, slice, share)?]);
482 }
483 let outcomes: Vec<Result<Embedded, String>> = std::thread::scope(|scope| {
484 let handles: Vec<_> = embedders
485 .iter()
486 .zip(shares)
487 .map(|(embedder, share)| {
488 let texts: Vec<String> = share
489 .iter()
490 .filter_map(|at| slice.get(*at).and_then(|row| row.text.as_ref()))
491 .map(|text| format!("{}{text}", plan.prefix))
492 .collect();
493 scope.spawn(move || {
494 embedder
495 .embed_prefixed_counted(&texts)
496 .map_err(|reason| format!("{reason:#}"))
497 })
498 })
499 .collect();
500 handles
501 .into_iter()
502 .map(|handle| {
503 handle.join().unwrap_or_else(|_| {
504 Err("an embedding thread stopped unexpectedly".to_string())
505 })
506 })
507 .collect()
508 });
509 outcomes
510 .into_iter()
511 .map(|outcome| {
512 outcome.map_err(|reason| {
513 Failed::said(Status::InvalidState, format!("embed: {reason}"))
514 })
515 })
516 .collect()
517 }
518
519 fn write_slice(
531 connection: &Connection<'_>,
532 plan: &Plan,
533 slice: &[Row],
534 vectors: &[Option<(Vec<f32>, usize)>],
535 max_tokens: usize,
536 report: &mut Report,
537 ) -> Result<(), Failed> {
538 let transaction = connection
539 .begin()
540 .map_err(|error| Failed::from_engine(&error))?;
541 let update = format!(
542 "UPDATE {} SET {} = ?1 WHERE rowid = ?2",
543 quoted(&plan.table),
544 quoted(&plan.vector_column)
545 );
546 let mut statement = transaction
547 .prepare(&update)
548 .map_err(|error| Failed::from_engine(&error))?;
549 for (row, vector) in slice.iter().zip(vectors) {
550 let Some((values, tokens)) = vector else {
551 report.skipped = report.skipped.saturating_add(1);
552 continue;
553 };
554 statement
555 .bind_blob(1, &as_bytes(values))
556 .and_then(|_| statement.bind_integer(2, row.rowid))
557 .and_then(|_| statement.step())
558 .map_err(|error| Failed::from_engine(&error))?;
559 statement.reset();
560 report.embedded = report.embedded.saturating_add(1);
561 if *tokens > max_tokens {
562 report.truncated = report.truncated.saturating_add(1);
563 if report.truncated_rowids.len() < MAX_NAMED_TRUNCATED {
564 report.truncated_rowids.push(row.rowid);
565 }
566 }
567 }
568 drop(statement);
569 transaction
570 .commit()
571 .map_err(|error| Failed::from_engine(&error))
572 }
573
574 fn as_bytes(values: &[f32]) -> Vec<u8> {
578 let mut bytes = Vec::with_capacity(values.len().saturating_mul(4));
579 for value in values {
580 bytes.extend_from_slice(&value.to_bits().to_le_bytes());
581 }
582 bytes
583 }
584
585 fn progress(report: &Report, total: usize, elapsed: std::time::Duration) {
591 if !std::io::stderr().is_terminal() {
592 return;
593 }
594 let done = report.embedded.saturating_add(report.skipped);
595 let rate = report.embedded as f64 / elapsed.as_secs_f64().max(0.001);
596 eprintln!("embed: {done} of {total} rows, {rate:.0} rows a second");
597 }
598
599 fn outcome(
606 plan: &Plan,
607 report: &Report,
608 elapsed: std::time::Duration,
609 loaded: std::time::Duration,
610 ) -> Outcome {
611 let seconds = elapsed.as_secs_f64();
612 let rate = report.embedded as f64 / seconds.max(0.001);
613 let mut text = format!(
614 "embedded {} rows in {seconds:.1} s, {rate:.0} rows a second, on {} with {} session{}. \
615 {} skipped because the text was NULL or empty. {} cut at the model's token limit",
616 report.embedded,
617 plan.device.label(),
618 plan.sessions,
619 if plan.sessions == 1 { "" } else { "s" },
620 report.skipped,
621 report.truncated
622 );
623 if !report.truncated_rowids.is_empty() {
624 let named: Vec<String> = report.truncated_rowids.iter().map(i64::to_string).collect();
625 text.push_str(&format!(". First cut rowids: {}", named.join(", ")));
626 }
627 let mut said = Outcome::said("embed", text);
628 said.changes = report.embedded as i64;
629 said.with("embedded", Json::Int(report.embedded as i64))
630 .with("skipped", Json::Int(report.skipped as i64))
631 .with("truncated", Json::Int(report.truncated as i64))
632 .with(
633 "truncated_rowids",
634 Json::Array(
635 report
636 .truncated_rowids
637 .iter()
638 .map(|id| Json::Int(*id))
639 .collect(),
640 ),
641 )
642 .with("seconds", Json::Real(seconds))
643 .with("rows_per_second", Json::Real(rate))
644 .with("load_seconds", Json::Real(loaded.as_secs_f64()))
645 .with("device", json::text(plan.device.label()))
646 .with("sessions", Json::Int(plan.sessions as i64))
647 }
648}