Skip to main content

inillucent_cli/
bulk_embed.rs

1//! `inillucent embed`: fill a vector column for every row of a table, on the processor or a graphics card.
2//!
3//! Invariant: **a row is either embedded by the same model and the same text that
4//! `embed(prefix || text)` would use, or it is reported.** The command reads the rows whose
5//! vector column is `NULL`, embeds each text with the prefix put in front, and writes the vector
6//! in the layout `embed()` returns. A row whose text is `NULL` or empty is skipped and counted. A
7//! row cut at the model's token limit is embedded from the tokens that fit and is counted and
8//! named, so nothing is silently shortened. The device is the one that was asked for: a card that
9//! will not start is an error that names the installer command and never a run on the processor.
10//!
11//! ## Why it is a command and not `UPDATE t SET v = embed(text)`
12//!
13//! `embed()` embeds one row per call on the processor, which for the 558,429 chunks of the study
14//! would take about 12 hours. This command sorts a slice of rows by length, groups them under a
15//! memory ceiling, runs several sessions at once, and commits every `--commit-every` rows, so a
16//! stopped run resumes at the first row still `NULL`.
17//!
18//! The model is opened here from `inillucent_core` because the SQL engine sits below the
19//! retrieval engine and cannot reach it; this crate already links both.
20
21use crate::command::{Arguments, Context, Failed, Outcome};
22
23/// The rows written in one transaction when `--commit-every` is not given.
24pub const DEFAULT_COMMIT_EVERY: i64 = 1024;
25
26/// How many rowids of truncated rows the report names.
27pub const MAX_NAMED_TRUNCATED: usize = 20;
28
29/// `embed`: embeds a text column into a vector column.
30///
31/// @param context - the open database
32/// @param arguments - what was asked for
33#[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/// `embed`: embeds a text column into a vector column.
42///
43/// @param context - the open database
44/// @param arguments - what was asked for
45#[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    /// What one run was asked to do.
67    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    /// What one run did.
81    #[derive(Default)]
82    struct Report {
83        embedded: usize,
84        skipped: usize,
85        truncated: usize,
86        truncated_rowids: Vec<i64>,
87    }
88
89    /// One row read from the table.
90    struct Row {
91        rowid: i64,
92        text: Option<String>,
93    }
94
95    /// Runs the command.
96    ///
97    /// @param context - the open database
98    /// @param arguments - what was asked for
99    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        // The table and the columns are checked before a model is opened, so a misspelled name
103        // is refused at once and does not cost the load of the model.
104        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    /// Reads and checks the arguments, refusing a value that cannot be used before any model is opened.
120    ///
121    /// The device and the thread count default to the machine's settings: the environment
122    /// variable, then what `setup-embeddings` recorded, then the processor.
123    ///
124    /// @param arguments - what was asked for
125    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    /// Reads a parameter that must be a whole number of at least 1.
153    ///
154    /// @param arguments - what was asked for
155    /// @param name - the parameter
156    /// @param default - the value when it is not given
157    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    /// The sessions an earlier call left behind, so a long lived process does not open them again.
168    ///
169    /// **A finished run parks its sessions here and never drops them.** Dropping an ONNX Runtime
170    /// session and then exiting the process ended it with `STATUS_STACK_BUFFER_OVERRUN` in about one
171    /// run in six on this machine, after the command had printed its result and committed every row,
172    /// so a script saw a failure for a run that had worked. A session that is never dropped does not
173    /// do it: 0 failures in 30 runs, against 5 in 30. `embed()` has the same property, because its
174    /// embedder is a static that is never dropped. Parked sessions are reused by the next call with
175    /// the same settings, so the MCP server, which calls this command many times in one process,
176    /// holds one set of sessions and not one per call.
177    static PARKED: std::sync::Mutex<Vec<(String, OnnxEmbedder)>> =
178        std::sync::Mutex::new(Vec::new());
179
180    /// Returns the key two runs share sessions under: the device, the thread count and the batch size.
181    ///
182    /// @param plan - what was asked for
183    fn session_key(plan: &Plan) -> String {
184        format!(
185            "{}|{:?}|{}",
186            plan.device.label(),
187            plan.threads,
188            plan.batch_size
189        )
190    }
191
192    /// Takes every parked session that was opened under a key.
193    ///
194    /// @param key - the settings the sessions were opened with
195    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    /// Parks sessions for the next call, and forgets them when the lock is not available.
205    ///
206    /// @param key - the settings the sessions were opened with
207    /// @param embedders - the sessions to keep
208    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    /// Opens one session per `--sessions` on the device, from the installed model.
220    ///
221    /// A device that will not start fails here, with the message that names
222    /// `inillucent setup-embeddings runtime --gpu`. There is no fallback to the processor.
223    ///
224    /// @param plan - what was asked for
225    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    /// Quotes a table or column name for SQL, doubling any quote inside it.
250    ///
251    /// @param name - the name as the caller typed it
252    fn quoted(name: &str) -> String {
253        format!("\"{}\"", name.replace('"', "\"\""))
254    }
255
256    /// Checks that the table and both columns exist, before any row is read.
257    ///
258    /// @param connection - the open database
259    /// @param plan - what was asked for
260    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    /// The condition that picks the rows to embed.
274    ///
275    /// Every row with `--all`, and otherwise the rows whose vector column is `NULL`, which is
276    /// what makes a stopped run resume where it ended.
277    ///
278    /// @param plan - what was asked for
279    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    /// Counts the rows the run will visit, for the progress line.
287    ///
288    /// @param connection - the open database
289    /// @param plan - what was asked for
290    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    /// Embeds every wanted row, one transaction at a time.
312    ///
313    /// @param connection - the open database
314    /// @param plan - what was asked for
315    /// @param embedders - the open sessions
316    /// @param max_tokens - the model's token limit
317    /// @param total - how many rows the run will visit
318    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    /// Reads the next `--commit-every` wanted rows after a rowid, in rowid order.
342    ///
343    /// @param connection - the open database
344    /// @param plan - what was asked for
345    /// @param after - the last rowid already visited
346    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    /// Returns a cell as the text to embed, or `None` for NULL and for an empty value.
382    ///
383    /// @param value - the text column's value
384    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    /// Embeds the rows of a slice that have text, spreading them over the sessions.
396    ///
397    /// The rows are sorted by length and dealt out one at a time, so every session gets a mix of
398    /// short and long texts and none waits for another. Each session sorts and groups its own
399    /// share by token count under the memory ceiling.
400    ///
401    /// @param embedders - the open sessions
402    /// @param plan - what was asked for
403    /// @param slice - the rows read
404    /// @returns one entry per row of the slice: its vector and token count, or `None` for a row with no text
405    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    /// Deals a sorted list of row positions out to a number of sessions, one at a time.
433    ///
434    /// @param order - the positions, sorted by text length
435    /// @param sessions - how many sessions share them
436    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    /// Embeds one session's share of a slice.
447    ///
448    /// @param embedder - the session
449    /// @param plan - what was asked for
450    /// @param slice - the rows read
451    /// @param share - the row positions this session embeds
452    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    /// Runs each session's share on its own thread and returns the results in share order.
469    ///
470    /// @param embedders - the open sessions
471    /// @param plan - what was asked for
472    /// @param slice - the rows read
473    /// @param shares - the row positions each session embeds
474    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    /// Writes one slice's vectors in one transaction, and adds what happened to the report.
520    ///
521    /// A crash loses at most this transaction, and the rows in it are still `NULL`, so a rerun
522    /// embeds them again.
523    ///
524    /// @param connection - the open database
525    /// @param plan - what was asked for
526    /// @param slice - the rows read
527    /// @param vectors - one entry per row of the slice
528    /// @param max_tokens - the model's token limit
529    /// @param report - the totals so far
530    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    /// Returns a vector as the bytes a `VECTOR(n)` column holds: little endian 32 bit floats.
575    ///
576    /// @param values - the vector
577    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    /// Prints one progress line to standard error when a person is watching.
586    ///
587    /// @param report - the totals so far
588    /// @param total - how many rows the run will visit
589    /// @param elapsed - the time since the first row
590    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    /// Builds what the command prints and returns as JSON.
600    ///
601    /// @param plan - what was asked for
602    /// @param report - what the run did
603    /// @param elapsed - the time from the first row to the last
604    /// @param loaded - the time it took to open the sessions
605    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}