Skip to main content

lance_index/scalar/inverted/
tokenizer.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use lance_core::{Error, Result};
5use serde::{Deserialize, Serialize};
6use std::{env, path::PathBuf};
7
8#[cfg(feature = "tokenizer-jieba")]
9mod jieba;
10
11pub mod document_tokenizer;
12#[cfg(feature = "tokenizer-lindera")]
13mod lindera;
14
15#[cfg(feature = "tokenizer-jieba")]
16use jieba::JiebaTokenizerBuilder;
17
18#[cfg(feature = "tokenizer-lindera")]
19use lindera::LinderaTokenizerBuilder;
20
21use crate::pbold;
22use crate::scalar::inverted::tokenizer::document_tokenizer::{
23    JsonTokenizer, LanceTokenizer, TextTokenizer,
24};
25pub use lance_tokenizer::Language;
26use lance_tokenizer::{
27    AsciiFoldingFilter, LowerCaser, NgramTokenizer, RawTokenizer, RemoveLongFilter,
28    SimpleTokenizer, Stemmer, StopWordFilter, TextAnalyzer, TextAnalyzerBuilder,
29    WhitespaceTokenizer,
30};
31
32/// Tokenizer configs
33#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
34pub struct InvertedIndexParams {
35    /// lance tokenizer takes care of different data types, such as text, json, etc.
36    /// - 'text': parsing input documents into tokens
37    /// - 'json': parsing input json string into tokens
38    /// - none: auto type inference
39    pub(crate) lance_tokenizer: Option<String>,
40    /// base tokenizer:
41    /// - `simple`: splits tokens on whitespace and punctuation
42    /// - `whitespace`: splits tokens on whitespace
43    /// - `raw`: no tokenization
44    /// - `lindera/*`: Lindera tokenizer
45    /// - `jieba/*`: Jieba tokenizer
46    ///
47    /// `simple` is recommended for most cases and the default value
48    pub(crate) base_tokenizer: String,
49
50    /// language for stemming and stop words
51    /// this is only used when `stem` or `remove_stop_words` is true
52    pub(crate) language: Language,
53
54    /// If true, store the position of the term in the document
55    /// This can significantly increase the size of the index
56    /// If false, only store the frequency of the term in the document
57    /// Default is false
58    #[serde(default)]
59    pub(crate) with_position: bool,
60
61    /// maximum token length
62    /// - `None`: no limit
63    /// - `Some(n)`: remove tokens longer than `n`
64    pub(crate) max_token_length: Option<usize>,
65
66    /// whether lower case tokens
67    #[serde(default = "bool_true")]
68    pub(crate) lower_case: bool,
69
70    /// whether apply stemming
71    #[serde(default = "bool_true")]
72    pub(crate) stem: bool,
73
74    /// whether remove stop words
75    #[serde(default = "bool_true")]
76    pub(crate) remove_stop_words: bool,
77
78    /// use customized stop words.
79    /// - `None`: use built-in stop words based on language
80    /// - `Some(words)`: use customized stop words
81    pub(crate) custom_stop_words: Option<Vec<String>>,
82
83    /// ascii folding
84    #[serde(default = "bool_true")]
85    pub(crate) ascii_folding: bool,
86
87    /// min ngram length
88    #[serde(default = "default_min_ngram_length")]
89    pub(crate) min_ngram_length: u32,
90
91    /// max ngram length
92    #[serde(default = "default_max_ngram_length")]
93    pub(crate) max_ngram_length: u32,
94
95    /// whether prefix only
96    #[serde(default)]
97    pub(crate) prefix_only: bool,
98
99    /// Total memory limit in MiB for the build stage.
100    ///
101    /// This is split evenly across FTS workers at build time. By default Lance
102    /// uses roughly `num_cpus / 2` workers, unless `LANCE_FTS_NUM_SHARDS` is set.
103    /// If unset, each worker defaults to a 2 GiB build-time memory limit.
104    ///
105    /// This is a build-time only parameter and is not persisted with the index.
106    #[serde(
107        rename = "memory_limit",
108        skip_serializing,
109        default,
110        alias = "worker_memory_limit_mb"
111    )]
112    pub(crate) memory_limit_mb: Option<u64>,
113
114    /// Number of workers to use for FTS build.
115    ///
116    /// This is a build-time only parameter and is not persisted with the index.
117    /// By default Lance uses roughly `num_cpus / 2` workers.
118    /// The effective worker count is clamped to `[1, num_cpus - 2]`.
119    #[serde(rename = "num_workers", skip_serializing, default)]
120    pub(crate) num_workers: Option<usize>,
121}
122
123impl TryFrom<&InvertedIndexParams> for pbold::InvertedIndexDetails {
124    type Error = Error;
125
126    fn try_from(params: &InvertedIndexParams) -> Result<Self> {
127        Ok(Self {
128            base_tokenizer: Some(params.base_tokenizer.clone()),
129            language: serde_json::to_string(&params.language)?,
130            with_position: params.with_position,
131            max_token_length: params.max_token_length.map(|l| l as u32),
132            lower_case: params.lower_case,
133            stem: params.stem,
134            remove_stop_words: params.remove_stop_words,
135            ascii_folding: params.ascii_folding,
136            min_ngram_length: params.min_ngram_length,
137            max_ngram_length: params.max_ngram_length,
138            prefix_only: params.prefix_only,
139        })
140    }
141}
142
143impl TryFrom<&pbold::InvertedIndexDetails> for InvertedIndexParams {
144    type Error = Error;
145
146    fn try_from(details: &pbold::InvertedIndexDetails) -> Result<Self> {
147        let defaults = Self::default();
148        Ok(Self {
149            lance_tokenizer: defaults.lance_tokenizer,
150            base_tokenizer: details
151                .base_tokenizer
152                .as_ref()
153                .cloned()
154                .unwrap_or(defaults.base_tokenizer),
155            language: serde_json::from_str(details.language.as_str())?,
156            with_position: details.with_position,
157            max_token_length: details.max_token_length.map(|l| l as usize),
158            lower_case: details.lower_case,
159            stem: details.stem,
160            remove_stop_words: details.remove_stop_words,
161            custom_stop_words: defaults.custom_stop_words,
162            ascii_folding: details.ascii_folding,
163            min_ngram_length: details.min_ngram_length,
164            max_ngram_length: details.max_ngram_length,
165            prefix_only: details.prefix_only,
166            memory_limit_mb: defaults.memory_limit_mb,
167            num_workers: defaults.num_workers,
168        })
169    }
170}
171
172fn bool_true() -> bool {
173    true
174}
175
176fn default_min_ngram_length() -> u32 {
177    3
178}
179
180fn default_max_ngram_length() -> u32 {
181    3
182}
183
184impl Default for InvertedIndexParams {
185    fn default() -> Self {
186        Self::new("simple".to_owned(), Language::English)
187    }
188}
189
190impl InvertedIndexParams {
191    /// Create a new `InvertedIndexParams` with the given base tokenizer and language.
192    ///
193    /// The `base_tokenizer` can be one of the following:
194    /// - `simple`: splits tokens on whitespace and punctuation, default
195    /// - `whitespace`: splits tokens on whitespace
196    /// - `raw`: no tokenization
197    /// - `ngram`: N-Gram tokenizer
198    /// - `lindera/*`: Lindera tokenizer
199    /// - `jieba/*`: Jieba tokenizer
200    ///
201    /// The `language` is used for stemming and removing stop words,
202    /// this is not used for `lindera/*` and `jieba/*` tokenizers.
203    /// Default to `English`.
204    pub fn new(base_tokenizer: String, language: Language) -> Self {
205        Self {
206            lance_tokenizer: None,
207            base_tokenizer,
208            language,
209            with_position: false,
210            max_token_length: Some(40),
211            lower_case: true,
212            stem: true,
213            remove_stop_words: true,
214            custom_stop_words: None,
215            ascii_folding: true,
216            min_ngram_length: default_min_ngram_length(),
217            max_ngram_length: default_max_ngram_length(),
218            prefix_only: false,
219            memory_limit_mb: None,
220            num_workers: None,
221        }
222    }
223
224    pub fn lance_tokenizer(mut self, lance_tokenizer: String) -> Self {
225        self.lance_tokenizer = Some(lance_tokenizer);
226        self
227    }
228
229    pub fn base_tokenizer(mut self, base_tokenizer: String) -> Self {
230        self.base_tokenizer = base_tokenizer;
231        self
232    }
233
234    pub fn language(mut self, language: &str) -> Result<Self> {
235        // need to convert to valid JSON string
236        let language = serde_json::from_str(format!("\"{}\"", language).as_str())?;
237        self.language = language;
238        Ok(self)
239    }
240
241    /// Set whether to store the position of the term in the document.
242    /// This can significantly increase the size of the index.
243    /// If false, only store the frequency of the term in the document.
244    /// This doesn't work with `ngram` tokenizer.
245    /// Default to `false`.
246    pub fn with_position(mut self, with_position: bool) -> Self {
247        self.with_position = with_position;
248        self
249    }
250
251    /// Get whether positions are stored in this index.
252    pub fn has_positions(&self) -> bool {
253        self.with_position
254    }
255
256    pub fn max_token_length(mut self, max_token_length: Option<usize>) -> Self {
257        self.max_token_length = max_token_length;
258        self
259    }
260
261    pub fn lower_case(mut self, lower_case: bool) -> Self {
262        self.lower_case = lower_case;
263        self
264    }
265
266    pub fn stem(mut self, stem: bool) -> Self {
267        self.stem = stem;
268        self
269    }
270
271    pub fn remove_stop_words(mut self, remove_stop_words: bool) -> Self {
272        self.remove_stop_words = remove_stop_words;
273        self
274    }
275
276    pub fn custom_stop_words(mut self, custom_stop_words: Option<Vec<String>>) -> Self {
277        self.custom_stop_words = custom_stop_words;
278        self
279    }
280
281    pub fn ascii_folding(mut self, ascii_folding: bool) -> Self {
282        self.ascii_folding = ascii_folding;
283        self
284    }
285
286    /// Set the minimum N-Gram length, only works when `base_tokenizer` is `ngram`.
287    /// Must be greater than 0 and not greater than `max_ngram_length`.
288    /// Default to 3.
289    pub fn ngram_min_length(mut self, min_length: u32) -> Self {
290        self.min_ngram_length = min_length;
291        self
292    }
293
294    /// Set the maximum N-Gram length, only works when `base_tokenizer` is `ngram`.
295    /// Must be greater than 0 and not less than `min_ngram_length`.
296    /// Default to 3.
297    pub fn ngram_max_length(mut self, max_length: u32) -> Self {
298        self.max_ngram_length = max_length;
299        self
300    }
301
302    /// Set whether only prefix N-Gram is generated, only works when `base_tokenizer` is `ngram`.
303    /// Default to `false`.
304    pub fn ngram_prefix_only(mut self, prefix_only: bool) -> Self {
305        self.prefix_only = prefix_only;
306        self
307    }
308
309    pub fn memory_limit_mb(mut self, memory_limit_mb: u64) -> Self {
310        self.memory_limit_mb = Some(memory_limit_mb);
311        self
312    }
313
314    /// Set the number of workers to use for this build.
315    ///
316    /// By default Lance uses roughly `num_cpus / 2` workers.
317    /// The effective worker count is clamped to `[1, num_cpus - 2]`.
318    pub fn num_workers(mut self, num_workers: usize) -> Self {
319        self.num_workers = Some(num_workers);
320        self
321    }
322
323    /// Serialize params for the build/training path, including build-only fields.
324    pub fn to_training_json(&self) -> serde_json::Result<serde_json::Value> {
325        let mut value = serde_json::to_value(self)?;
326        let object = value
327            .as_object_mut()
328            .expect("inverted index params should serialize to a JSON object");
329        if let Some(memory_limit_mb) = self.memory_limit_mb {
330            object.insert(
331                "memory_limit".to_string(),
332                serde_json::Value::from(memory_limit_mb),
333            );
334        }
335        if let Some(num_workers) = self.num_workers {
336            object.insert(
337                "num_workers".to_string(),
338                serde_json::Value::from(num_workers),
339            );
340        }
341        Ok(value)
342    }
343
344    pub fn build(&self) -> Result<Box<dyn LanceTokenizer>> {
345        let mut builder = self.build_base_tokenizer()?;
346        if let Some(max_token_length) = self.max_token_length {
347            builder = builder.filter_dynamic(RemoveLongFilter::limit(max_token_length));
348        }
349        if self.lower_case {
350            builder = builder.filter_dynamic(LowerCaser);
351        }
352        if self.stem {
353            builder = builder.filter_dynamic(Stemmer::new(self.language));
354        }
355        if self.remove_stop_words {
356            let stop_word_filter = match &self.custom_stop_words {
357                Some(words) => StopWordFilter::remove(words.iter().cloned()),
358                None => StopWordFilter::new(self.language).ok_or_else(|| {
359                    Error::invalid_input(format!(
360                        "removing stop words for language {:?} is not supported yet",
361                        self.language
362                    ))
363                })?,
364            };
365            builder = builder.filter_dynamic(stop_word_filter);
366        }
367        if self.ascii_folding {
368            builder = builder.filter_dynamic(AsciiFoldingFilter);
369        }
370        let tokenizer = builder.build();
371
372        match self.lance_tokenizer {
373            Some(ref t) if t == "text" => Ok(Box::new(TextTokenizer::new(tokenizer))),
374            Some(ref t) if t == "json" => Ok(Box::new(JsonTokenizer::new(tokenizer))),
375            None => Ok(Box::new(TextTokenizer::new(tokenizer))),
376            _ => Err(Error::invalid_input(format!(
377                "unknown lance tokenizer {}",
378                self.lance_tokenizer.as_ref().unwrap()
379            ))),
380        }
381    }
382
383    fn build_base_tokenizer(&self) -> Result<TextAnalyzerBuilder> {
384        match self.base_tokenizer.as_str() {
385            "simple" => Ok(TextAnalyzer::builder(SimpleTokenizer::default()).dynamic()),
386            "whitespace" => Ok(TextAnalyzer::builder(WhitespaceTokenizer::default()).dynamic()),
387            "raw" => Ok(TextAnalyzer::builder(RawTokenizer::default()).dynamic()),
388            "ngram" => {
389                let tokenizer = NgramTokenizer::new(
390                    self.min_ngram_length as usize,
391                    self.max_ngram_length as usize,
392                    self.prefix_only,
393                )
394                .map_err(|e| Error::invalid_input(e.to_string()))?;
395                Ok(TextAnalyzer::builder(tokenizer).dynamic())
396            }
397            #[cfg(feature = "tokenizer-lindera")]
398            s if s.starts_with("lindera/") => {
399                let Some(home) = language_model_home() else {
400                    return Err(Error::invalid_input(format!(
401                        "unknown base tokenizer {}",
402                        self.base_tokenizer
403                    )));
404                };
405                lindera::LinderaBuilder::load(&home.join(s))?.build()
406            }
407            #[cfg(feature = "tokenizer-jieba")]
408            s if s.starts_with("jieba/") || s == "jieba" => {
409                let s = if s == "jieba" { "jieba/default" } else { s };
410                let Some(home) = language_model_home() else {
411                    return Err(Error::invalid_input(format!(
412                        "unknown base tokenizer {}",
413                        self.base_tokenizer
414                    )));
415                };
416                jieba::JiebaBuilder::load(&home.join(s))?.build()
417            }
418            _ => Err(Error::invalid_input(format!(
419                "unknown base tokenizer {}",
420                self.base_tokenizer
421            ))),
422        }
423    }
424}
425
426pub const LANCE_LANGUAGE_MODEL_HOME_ENV_KEY: &str = "LANCE_LANGUAGE_MODEL_HOME";
427
428pub const LANCE_LANGUAGE_MODEL_DEFAULT_DIRECTORY: &str = "lance/language_models";
429
430pub fn language_model_home() -> Option<PathBuf> {
431    match env::var(LANCE_LANGUAGE_MODEL_HOME_ENV_KEY) {
432        Ok(p) => Some(PathBuf::from(p)),
433        Err(_) => dirs::data_local_dir().map(|p| p.join(LANCE_LANGUAGE_MODEL_DEFAULT_DIRECTORY)),
434    }
435}
436
437#[cfg(test)]
438mod tests {
439    use super::InvertedIndexParams;
440
441    #[test]
442    fn test_build_only_fields_are_not_serialized() {
443        let params = InvertedIndexParams::default()
444            .memory_limit_mb(4096)
445            .num_workers(7);
446        let json = serde_json::to_value(&params).unwrap();
447        assert!(json.get("memory_limit").is_none());
448        assert!(json.get("num_workers").is_none());
449    }
450
451    #[test]
452    fn test_memory_limit_serde_accepts_legacy_worker_field_name() {
453        let mut json = serde_json::to_value(InvertedIndexParams::default()).unwrap();
454        let obj = json.as_object_mut().unwrap();
455        obj.remove("memory_limit");
456        obj.insert(
457            "worker_memory_limit_mb".to_string(),
458            serde_json::Value::from(2048),
459        );
460        let params: InvertedIndexParams = serde_json::from_value(json).unwrap();
461        assert_eq!(params.memory_limit_mb, Some(2048));
462    }
463
464    #[test]
465    fn test_build_only_fields_deserialize_from_public_names() {
466        let mut json = serde_json::to_value(InvertedIndexParams::default()).unwrap();
467        let obj = json.as_object_mut().unwrap();
468        obj.insert("memory_limit".to_string(), serde_json::Value::from(4096));
469        obj.insert("num_workers".to_string(), serde_json::Value::from(3));
470
471        let params: InvertedIndexParams = serde_json::from_value(json).unwrap();
472        assert_eq!(params.memory_limit_mb, Some(4096));
473        assert_eq!(params.num_workers, Some(3));
474    }
475
476    #[test]
477    fn test_training_json_serializes_build_only_fields() {
478        let params = InvertedIndexParams::default()
479            .memory_limit_mb(4096)
480            .num_workers(3);
481        let json = params.to_training_json().unwrap();
482        assert_eq!(
483            json.get("memory_limit"),
484            Some(&serde_json::Value::from(4096))
485        );
486        assert_eq!(json.get("num_workers"), Some(&serde_json::Value::from(3)));
487    }
488}