lance_index/scalar/inverted/
tokenizer.rs1use 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
34pub struct InvertedIndexParams {
35 pub(crate) lance_tokenizer: Option<String>,
40 pub(crate) base_tokenizer: String,
49
50 pub(crate) language: Language,
53
54 #[serde(default)]
59 pub(crate) with_position: bool,
60
61 pub(crate) max_token_length: Option<usize>,
65
66 #[serde(default = "bool_true")]
68 pub(crate) lower_case: bool,
69
70 #[serde(default = "bool_true")]
72 pub(crate) stem: bool,
73
74 #[serde(default = "bool_true")]
76 pub(crate) remove_stop_words: bool,
77
78 pub(crate) custom_stop_words: Option<Vec<String>>,
82
83 #[serde(default = "bool_true")]
85 pub(crate) ascii_folding: bool,
86
87 #[serde(default = "default_min_ngram_length")]
89 pub(crate) min_ngram_length: u32,
90
91 #[serde(default = "default_max_ngram_length")]
93 pub(crate) max_ngram_length: u32,
94
95 #[serde(default)]
97 pub(crate) prefix_only: bool,
98
99 #[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 #[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(¶ms.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 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 let language = serde_json::from_str(format!("\"{}\"", language).as_str())?;
237 self.language = language;
238 Ok(self)
239 }
240
241 pub fn with_position(mut self, with_position: bool) -> Self {
247 self.with_position = with_position;
248 self
249 }
250
251 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 pub fn ngram_min_length(mut self, min_length: u32) -> Self {
290 self.min_ngram_length = min_length;
291 self
292 }
293
294 pub fn ngram_max_length(mut self, max_length: u32) -> Self {
298 self.max_ngram_length = max_length;
299 self
300 }
301
302 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 pub fn num_workers(mut self, num_workers: usize) -> Self {
319 self.num_workers = Some(num_workers);
320 self
321 }
322
323 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(¶ms).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}