1use std::cell::Cell;
4use std::collections::{HashMap, HashSet};
5
6use runmat_builtins::{
7 Access, BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
8 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
9 CharArray, ClassDef, ObjectInstance, PropertyDef, ResolveContext, StringArray, Tensor, Type,
10 Value,
11};
12use runmat_macros::runtime_builtin;
13
14use crate::builtins::strings::common::is_missing_string;
15use crate::builtins::strings::core::compat::scalar_text;
16use crate::builtins::strings::text_analytics::documents::{
17 checked_count_len, documents_from_object, words_from_word_vector,
18 words_from_word_vector_preserving_missing, TOKENIZED_DOCUMENT_CLASS,
19};
20use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult};
21
22pub const BAG_OF_NGRAMS_CLASS: &str = "bagOfNgrams";
23
24thread_local! {
25 static BAG_OF_NGRAMS_CLASS_REGISTERED: Cell<bool> = const { Cell::new(false) };
26}
27
28const OUT_BAG: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
29 name: "bag",
30 ty: BuiltinParamType::Any,
31 arity: BuiltinParamArity::Required,
32 default: None,
33 description: "Bag-of-n-grams model object.",
34}];
35
36const IN_DOCUMENTS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
37 name: "documents",
38 ty: BuiltinParamType::Any,
39 arity: BuiltinParamArity::Required,
40 default: None,
41 description: "Tokenized documents or a single-document word vector.",
42}];
43
44const IN_DOCUMENTS_REST: [BuiltinParamDescriptor; 2] = [
45 BuiltinParamDescriptor {
46 name: "documents",
47 ty: BuiltinParamType::Any,
48 arity: BuiltinParamArity::Required,
49 default: None,
50 description: "Tokenized documents or a single-document word vector.",
51 },
52 BuiltinParamDescriptor {
53 name: "NameValue",
54 ty: BuiltinParamType::Any,
55 arity: BuiltinParamArity::Variadic,
56 default: None,
57 description: "Name-value option: NgramLengths.",
58 },
59];
60
61const IN_NGRAMS_COUNTS: [BuiltinParamDescriptor; 2] = [
62 BuiltinParamDescriptor {
63 name: "uniqueNgrams",
64 ty: BuiltinParamType::Any,
65 arity: BuiltinParamArity::Required,
66 default: None,
67 description: "Unique n-gram string matrix.",
68 },
69 BuiltinParamDescriptor {
70 name: "counts",
71 ty: BuiltinParamType::Any,
72 arity: BuiltinParamArity::Required,
73 default: None,
74 description: "N-gram counts per document.",
75 },
76];
77
78const IN_NGRAMS_COUNTS_REST: [BuiltinParamDescriptor; 3] = [
79 BuiltinParamDescriptor {
80 name: "uniqueNgrams",
81 ty: BuiltinParamType::Any,
82 arity: BuiltinParamArity::Required,
83 default: None,
84 description: "Unique n-gram string matrix.",
85 },
86 BuiltinParamDescriptor {
87 name: "counts",
88 ty: BuiltinParamType::Any,
89 arity: BuiltinParamArity::Required,
90 default: None,
91 description: "N-gram counts per document.",
92 },
93 BuiltinParamDescriptor {
94 name: "NameValue",
95 ty: BuiltinParamType::Any,
96 arity: BuiltinParamArity::Variadic,
97 default: None,
98 description: "Name-value option: NgramLengths.",
99 },
100];
101
102const ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
103 code: "RM.TEXT_ANALYTICS_NGRAMS.INVALID_INPUT",
104 identifier: Some("RunMat:bagOfNgrams:InvalidInput"),
105 when: "Inputs do not match a supported bagOfNgrams form.",
106 message: "bagOfNgrams: invalid input",
107};
108
109const ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_INVALID_INPUT];
110
111pub const BAG_OF_NGRAMS_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
112 signatures: &[
113 BuiltinSignatureDescriptor {
114 label: "bag = bagOfNgrams",
115 inputs: &[],
116 outputs: &OUT_BAG,
117 },
118 BuiltinSignatureDescriptor {
119 label: "bag = bagOfNgrams(documents)",
120 inputs: &IN_DOCUMENTS,
121 outputs: &OUT_BAG,
122 },
123 BuiltinSignatureDescriptor {
124 label: "bag = bagOfNgrams(___, 'NgramLengths', lengths)",
125 inputs: &IN_DOCUMENTS_REST,
126 outputs: &OUT_BAG,
127 },
128 BuiltinSignatureDescriptor {
129 label: "bag = bagOfNgrams(uniqueNgrams, counts)",
130 inputs: &IN_NGRAMS_COUNTS,
131 outputs: &OUT_BAG,
132 },
133 BuiltinSignatureDescriptor {
134 label: "bag = bagOfNgrams(uniqueNgrams, counts, 'NgramLengths', lengths)",
135 inputs: &IN_NGRAMS_COUNTS_REST,
136 outputs: &OUT_BAG,
137 },
138 ],
139 output_mode: BuiltinOutputMode::Fixed,
140 completion_policy: BuiltinCompletionPolicy::Public,
141 errors: &ERRORS,
142};
143
144fn any_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
145 Type::Unknown
146}
147
148fn ngrams_error(message: impl Into<String>) -> crate::RuntimeError {
149 let mut builder = build_runtime_error(message).with_builtin("bagOfNgrams");
150 if let Some(identifier) = ERROR_INVALID_INPUT.identifier {
151 builder = builder.with_identifier(identifier);
152 }
153 builder.build()
154}
155
156fn ensure_bag_of_ngrams_class_registered() {
157 BAG_OF_NGRAMS_CLASS_REGISTERED.with(|registered| {
158 if registered.get() {
159 return;
160 }
161 let mut properties = HashMap::new();
162 for name in [
163 "Counts",
164 "Ngrams",
165 "NgramLengths",
166 "Vocabulary",
167 "NumNgrams",
168 "NumDocuments",
169 ] {
170 properties.insert(name.to_string(), property_def(name));
171 }
172 runmat_builtins::register_class(ClassDef {
173 name: BAG_OF_NGRAMS_CLASS.to_string(),
174 parent: None,
175 properties,
176 methods: HashMap::new(),
177 });
178 registered.set(true);
179 });
180}
181
182fn property_def(name: &str) -> PropertyDef {
183 PropertyDef {
184 name: name.to_string(),
185 is_static: false,
186 is_constant: false,
187 is_dependent: false,
188 get_access: Access::Public,
189 set_access: Access::Public,
190 default_value: None,
191 }
192}
193
194#[runtime_builtin(
195 name = "bagOfNgrams",
196 category = "strings/text_analytics",
197 summary = "Create bag-of-n-grams model objects.",
198 keywords = "bagOfNgrams,text analytics,n-grams,word counts",
199 accel = "sink",
200 type_resolver(any_type),
201 descriptor(crate::builtins::strings::text_analytics::ngrams::BAG_OF_NGRAMS_DESCRIPTOR),
202 builtin_path = "crate::builtins::strings::text_analytics::ngrams"
203)]
204async fn bag_of_ngrams_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
205 let gathered = gather_args(args).await?;
206 let parsed = parse_args(gathered)?;
207 match parsed.source {
208 NgramSource::Empty => bag_object(Vec::new(), parsed.lengths, Vec::new(), 0),
209 NgramSource::Documents(documents) => bag_from_documents(documents, parsed.lengths),
210 NgramSource::Unique {
211 ngrams,
212 counts,
213 requested_lengths,
214 } => bag_from_unique_ngrams(ngrams, counts, requested_lengths),
215 }
216}
217
218async fn gather_args(args: Vec<Value>) -> BuiltinResult<Vec<Value>> {
219 let mut out = Vec::with_capacity(args.len());
220 for arg in args {
221 out.push(
222 gather_if_needed_async(&arg).await.map_err(|err| {
223 ngrams_error(format!("bagOfNgrams: failed to gather input: {err}"))
224 })?,
225 );
226 }
227 Ok(out)
228}
229
230struct ParsedArgs {
231 source: NgramSource,
232 lengths: Vec<usize>,
233}
234
235enum NgramSource {
236 Empty,
237 Documents(Vec<Vec<String>>),
238 Unique {
239 ngrams: Vec<Vec<String>>,
240 counts: Tensor,
241 requested_lengths: Option<Vec<usize>>,
242 },
243}
244
245fn parse_args(args: Vec<Value>) -> BuiltinResult<ParsedArgs> {
246 if args.is_empty() {
247 return Ok(ParsedArgs {
248 source: NgramSource::Empty,
249 lengths: vec![2],
250 });
251 }
252
253 let first_is_option = is_option_name(&args[0], "NgramLengths");
254 if first_is_option {
255 let lengths = parse_options(&args, 0)?.unwrap_or_else(|| vec![2]);
256 return Ok(ParsedArgs {
257 source: NgramSource::Empty,
258 lengths,
259 });
260 }
261
262 if args.len() >= 2 && !is_option_name(&args[1], "NgramLengths") {
263 let counts = match &args[1] {
264 Value::Tensor(tensor) => tensor.clone(),
265 other => {
266 return Err(ngrams_error(format!(
267 "bagOfNgrams: counts must be a numeric matrix, got {other:?}"
268 )))
269 }
270 };
271 let lengths = parse_options(&args, 2)?;
272 return Ok(ParsedArgs {
273 lengths: lengths.clone().unwrap_or_else(|| vec![2]),
274 source: NgramSource::Unique {
275 ngrams: unique_ngrams_from_value(&args[0], counts.cols)?,
276 counts,
277 requested_lengths: lengths,
278 },
279 });
280 }
281
282 let lengths = parse_options(&args, 1)?.unwrap_or_else(|| vec![2]);
283 Ok(ParsedArgs {
284 source: NgramSource::Documents(documents_from_value(&args[0])?),
285 lengths,
286 })
287}
288
289fn is_option_name(value: &Value, expected: &str) -> bool {
290 scalar_text(value, "bagOfNgrams")
291 .map(|text| text.eq_ignore_ascii_case(expected))
292 .unwrap_or(false)
293}
294
295fn parse_options(args: &[Value], start: usize) -> BuiltinResult<Option<Vec<usize>>> {
296 if start >= args.len() {
297 return Ok(None);
298 }
299 if !(args.len() - start).is_multiple_of(2) {
300 return Err(ngrams_error(
301 "bagOfNgrams: name-value options must appear in pairs",
302 ));
303 }
304 let mut lengths = None;
305 let mut idx = start;
306 while idx < args.len() {
307 let name =
308 scalar_text(&args[idx], "bagOfNgrams").map_err(|err| ngrams_error(err.to_string()))?;
309 match name.to_ascii_lowercase().as_str() {
310 "ngramlengths" => {
311 lengths = Some(parse_lengths(&args[idx + 1])?);
312 }
313 other => {
314 return Err(ngrams_error(format!(
315 "bagOfNgrams: unsupported option '{other}'"
316 )));
317 }
318 }
319 idx += 2;
320 }
321 Ok(lengths)
322}
323
324fn parse_lengths(value: &Value) -> BuiltinResult<Vec<usize>> {
325 let raw = match value {
326 Value::Num(n) => vec![*n],
327 Value::Tensor(tensor) if !tensor.data.is_empty() => tensor.data.clone(),
328 other => {
329 return Err(ngrams_error(format!(
330 "bagOfNgrams: NgramLengths must be a positive integer scalar or vector, got {other:?}"
331 )))
332 }
333 };
334 let mut lengths = Vec::with_capacity(raw.len());
335 let mut seen = HashSet::new();
336 for n in raw {
337 if !n.is_finite() || n <= 0.0 || n.fract() != 0.0 {
338 return Err(ngrams_error(format!(
339 "bagOfNgrams: NgramLengths must contain positive integers, got {n}"
340 )));
341 }
342 let len = n as usize;
343 if seen.insert(len) {
344 lengths.push(len);
345 }
346 }
347 Ok(lengths)
348}
349
350fn documents_from_value(value: &Value) -> BuiltinResult<Vec<Vec<String>>> {
351 match value {
352 Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => {
353 documents_from_object(object, "bagOfNgrams")
354 }
355 Value::Object(object) => Err(ngrams_error(format!(
356 "bagOfNgrams: expected tokenizedDocument object, got {}",
357 object.class_name
358 ))),
359 other => {
360 validate_row_word_vector(other)?;
361 Ok(vec![words_from_word_vector(other, "bagOfNgrams")?])
362 }
363 }
364}
365
366fn validate_row_word_vector(value: &Value) -> BuiltinResult<()> {
367 match value {
368 Value::String(_) | Value::Num(_) => Ok(()),
369 Value::StringArray(array) if array.rows <= 1 => Ok(()),
370 Value::CharArray(CharArray { rows, .. }) if *rows <= 1 => Ok(()),
371 Value::Cell(cell) if cell.rows <= 1 => Ok(()),
372 Value::StringArray(array) => Err(ngrams_error(format!(
373 "bagOfNgrams: non-tokenized documents input must be a row word vector; got string array with shape {}x{}",
374 array.rows, array.cols
375 ))),
376 Value::CharArray(CharArray { rows, cols, .. }) => Err(ngrams_error(format!(
377 "bagOfNgrams: non-tokenized documents input must be a row word vector; got char array with shape {rows}x{cols}"
378 ))),
379 Value::Cell(cell) => Err(ngrams_error(format!(
380 "bagOfNgrams: non-tokenized documents input must be a row word vector; got cell array with shape {}x{}",
381 cell.rows, cell.cols
382 ))),
383 _ => Ok(()),
384 }
385}
386
387fn bag_from_documents(documents: Vec<Vec<String>>, lengths: Vec<usize>) -> BuiltinResult<Value> {
388 let mut ngrams = Vec::new();
389 let mut positions: HashMap<Vec<String>, usize> = HashMap::new();
390 let rows = documents.len();
391 let mut counts = Vec::<f64>::new();
392
393 for (doc_idx, document) in documents.iter().enumerate() {
394 for &length in &lengths {
395 if length > document.len() {
396 continue;
397 }
398 for start in 0..=document.len() - length {
399 let ngram = document[start..start + length].to_vec();
400 let col = if let Some(col) = positions.get(&ngram) {
401 *col
402 } else {
403 let col = ngrams.len();
404 positions.insert(ngram.clone(), col);
405 ngrams.push(ngram);
406 counts.resize(checked_count_len(rows, ngrams.len(), "bagOfNgrams")?, 0.0);
407 col
408 };
409 counts[doc_idx + col * rows] += 1.0;
410 }
411 }
412 }
413
414 bag_object(ngrams, lengths, counts, rows)
415}
416
417fn bag_from_unique_ngrams(
418 raw_ngrams: Vec<Vec<String>>,
419 counts: Tensor,
420 requested_lengths: Option<Vec<usize>>,
421) -> BuiltinResult<Value> {
422 if counts.cols != raw_ngrams.len() {
423 return Err(ngrams_error(format!(
424 "bagOfNgrams: counts columns ({}) must match uniqueNgrams rows ({})",
425 counts.cols,
426 raw_ngrams.len()
427 )));
428 }
429 if counts
430 .data
431 .iter()
432 .any(|value| !value.is_finite() || *value < 0.0 || value.fract() != 0.0)
433 {
434 return Err(ngrams_error(
435 "bagOfNgrams: counts must be nonnegative integers",
436 ));
437 }
438
439 let requested = requested_lengths
440 .as_ref()
441 .map(|lengths| lengths.iter().copied().collect::<HashSet<_>>());
442 let mut seen = HashSet::new();
443 let mut keep_cols = Vec::new();
444 let mut ngrams = Vec::new();
445 for (col, ngram) in raw_ngrams.iter().enumerate() {
446 if ngram.iter().any(|word| is_missing_string(word)) {
447 continue;
448 }
449 if ngram.is_empty() {
450 return Err(ngrams_error(
451 "bagOfNgrams: each n-gram must contain at least one word",
452 ));
453 }
454 if requested
455 .as_ref()
456 .is_some_and(|lengths| !lengths.contains(&ngram.len()))
457 {
458 continue;
459 }
460 if !seen.insert(ngram.clone()) {
461 return Err(ngrams_error(format!(
462 "bagOfNgrams: uniqueNgrams contains duplicate n-gram '{}'",
463 ngram.join(" ")
464 )));
465 }
466 keep_cols.push(col);
467 ngrams.push(ngram.clone());
468 }
469
470 let mut filtered_counts = Vec::with_capacity(checked_count_len(
471 counts.rows,
472 keep_cols.len(),
473 "bagOfNgrams",
474 )?);
475 for col in keep_cols {
476 for row in 0..counts.rows {
477 filtered_counts.push(counts.data[row + col * counts.rows]);
478 }
479 }
480 let lengths = requested_lengths.unwrap_or_else(|| infer_ngram_lengths(&ngrams));
481 bag_object(ngrams, lengths, filtered_counts, counts.rows)
482}
483
484fn infer_ngram_lengths(ngrams: &[Vec<String>]) -> Vec<usize> {
485 let mut lengths = Vec::new();
486 let mut seen = HashSet::new();
487 for ngram in ngrams {
488 if seen.insert(ngram.len()) {
489 lengths.push(ngram.len());
490 }
491 }
492 if lengths.is_empty() {
493 lengths.push(2);
494 }
495 lengths
496}
497
498fn unique_ngrams_from_value(
499 value: &Value,
500 expected_rows: usize,
501) -> BuiltinResult<Vec<Vec<String>>> {
502 match value {
503 Value::StringArray(array) => {
504 if array.rows != expected_rows {
505 return Err(ngrams_error(format!(
506 "bagOfNgrams: uniqueNgrams rows ({}) must match counts columns ({expected_rows})",
507 array.rows
508 )));
509 }
510 let mut out = Vec::with_capacity(array.rows);
511 for row in 0..array.rows {
512 let mut ngram = Vec::new();
513 let mut row_has_missing = false;
514 for col in 0..array.cols {
515 let word = array.data[row + col * array.rows].clone();
516 if is_missing_string(&word) {
517 row_has_missing = true;
518 break;
519 }
520 if !word.is_empty() {
521 ngram.push(word);
522 }
523 }
524 if row_has_missing {
525 out.push(vec!["<missing>".to_string()]);
526 } else {
527 out.push(ngram);
528 }
529 }
530 Ok(out)
531 }
532 Value::Cell(cell) => {
533 if cell.rows != expected_rows {
534 return Err(ngrams_error(format!(
535 "bagOfNgrams: uniqueNgrams rows ({}) must match counts columns ({expected_rows})",
536 cell.rows
537 )));
538 }
539 let mut out = Vec::with_capacity(cell.rows);
540 for row in 0..cell.rows {
541 let mut ngram = Vec::new();
542 let mut row_has_missing = false;
543 for col in 0..cell.cols {
544 let idx = row + col * cell.rows;
545 let word = scalar_text(&cell.data[idx], "bagOfNgrams")
546 .map_err(|err| ngrams_error(err.to_string()))?;
547 if is_missing_string(&word) {
548 row_has_missing = true;
549 break;
550 }
551 if !word.is_empty() {
552 ngram.push(word);
553 }
554 }
555 if row_has_missing {
556 out.push(vec!["<missing>".to_string()]);
557 } else {
558 out.push(ngram);
559 }
560 }
561 Ok(out)
562 }
563 other => {
564 let words = words_from_word_vector_preserving_missing(other, "bagOfNgrams")?;
565 if expected_rows != 1 {
566 return Err(ngrams_error(format!(
567 "bagOfNgrams: uniqueNgrams rows (1) must match counts columns ({expected_rows})"
568 )));
569 }
570 Ok(vec![words
571 .into_iter()
572 .filter(|word| !word.is_empty())
573 .collect()])
574 }
575 }
576}
577
578fn bag_object(
579 ngrams: Vec<Vec<String>>,
580 lengths: Vec<usize>,
581 counts: Vec<f64>,
582 rows: usize,
583) -> BuiltinResult<Value> {
584 ensure_bag_of_ngrams_class_registered();
585 let cols = ngrams.len();
586 let expected = checked_count_len(rows, cols, "bagOfNgrams")?;
587 if counts.len() != expected {
588 return Err(ngrams_error(format!(
589 "bagOfNgrams: count storage has {} values but expected {} for a {}x{} model",
590 counts.len(),
591 expected,
592 rows,
593 cols
594 )));
595 }
596 let max_len = ngrams.iter().map(Vec::len).max().unwrap_or(0);
597 let mut object = ObjectInstance::new(BAG_OF_NGRAMS_CLASS.to_string());
598 object.properties.insert(
599 "Ngrams".to_string(),
600 Value::StringArray(ngram_array(&ngrams, max_len)?),
601 );
602 object.properties.insert(
603 "Counts".to_string(),
604 Value::Tensor(Tensor::new(counts, vec![rows, cols]).map_err(|err| ngrams_error(err))?),
605 );
606 object.properties.insert(
607 "NgramLengths".to_string(),
608 Value::Tensor(
609 Tensor::new(
610 lengths.iter().map(|length| *length as f64).collect(),
611 vec![1, lengths.len()],
612 )
613 .map_err(|err| ngrams_error(err))?,
614 ),
615 );
616 object.properties.insert(
617 "Vocabulary".to_string(),
618 Value::StringArray(vocabulary_array(&ngrams)?),
619 );
620 object
621 .properties
622 .insert("NumNgrams".to_string(), Value::Num(cols as f64));
623 object
624 .properties
625 .insert("NumDocuments".to_string(), Value::Num(rows as f64));
626 Ok(Value::Object(object))
627}
628
629pub(in crate::builtins::strings::text_analytics) fn ngrams_from_bag(
630 object: &ObjectInstance,
631 fn_name: &str,
632) -> BuiltinResult<Vec<Vec<String>>> {
633 match object.properties.get("Ngrams") {
634 Some(Value::StringArray(array)) => {
635 let mut ngrams = Vec::with_capacity(array.rows);
636 let mut seen = HashSet::new();
637 for row in 0..array.rows {
638 let mut ngram = Vec::new();
639 for col in 0..array.cols {
640 let word = &array.data[row + col * array.rows];
641 if !word.is_empty() && !is_missing_string(word) {
642 ngram.push(word.clone());
643 }
644 }
645 if ngram.is_empty() {
646 return Err(ngrams_error(format!(
647 "{fn_name}: bagOfNgrams object contains an empty n-gram"
648 )));
649 }
650 if !seen.insert(ngram.clone()) {
651 return Err(ngrams_error(format!(
652 "{fn_name}: bagOfNgrams object contains duplicate n-gram '{}'",
653 ngram.join(" ")
654 )));
655 }
656 ngrams.push(ngram);
657 }
658 Ok(ngrams)
659 }
660 Some(other) => Err(ngrams_error(format!(
661 "{fn_name}: bagOfNgrams Ngrams property must be a string array, got {other:?}"
662 ))),
663 None => Err(ngrams_error(format!(
664 "{fn_name}: bagOfNgrams object missing Ngrams property"
665 ))),
666 }
667}
668
669fn ngram_array(ngrams: &[Vec<String>], max_len: usize) -> BuiltinResult<StringArray> {
670 let rows = ngrams.len();
671 let mut data = Vec::with_capacity(rows * max_len);
672 for col in 0..max_len {
673 for ngram in ngrams {
674 data.push(ngram.get(col).cloned().unwrap_or_default());
675 }
676 }
677 StringArray::new(data, vec![rows, max_len]).map_err(|err| ngrams_error(err))
678}
679
680fn vocabulary_array(ngrams: &[Vec<String>]) -> BuiltinResult<StringArray> {
681 let mut seen = HashSet::new();
682 let mut words = Vec::new();
683 for word in ngrams.iter().flatten() {
684 if seen.insert(word.clone()) {
685 words.push(word.clone());
686 }
687 }
688 StringArray::new(words.clone(), vec![1, words.len()]).map_err(|err| ngrams_error(err))
689}
690
691#[cfg(test)]
692mod tests {
693 use super::*;
694 use crate::builtins::strings::text_analytics::documents::TOKENIZED_DOCUMENT_CLASS;
695 use runmat_builtins::CellArray;
696
697 fn run(args: Vec<Value>) -> BuiltinResult<Value> {
698 futures::executor::block_on(bag_of_ngrams_builtin(args))
699 }
700
701 fn object(value: Value) -> ObjectInstance {
702 let Value::Object(object) = value else {
703 panic!("expected object");
704 };
705 object
706 }
707
708 fn tokenized(documents: Vec<Vec<&str>>) -> Value {
709 let values = documents
710 .into_iter()
711 .map(|doc| {
712 let len = doc.len();
713 Value::StringArray(
714 StringArray::new(
715 doc.into_iter().map(str::to_string).collect::<Vec<_>>(),
716 vec![1, len],
717 )
718 .unwrap(),
719 )
720 })
721 .collect::<Vec<_>>();
722 let rows = values.len();
723 let mut object = ObjectInstance::new(TOKENIZED_DOCUMENT_CLASS.to_string());
724 object.properties.insert(
725 "Documents".to_string(),
726 Value::Cell(CellArray::new(values, rows, 1).unwrap()),
727 );
728 Value::Object(object)
729 }
730
731 fn string_array_property(object: &ObjectInstance, name: &str) -> StringArray {
732 let Some(Value::StringArray(array)) = object.properties.get(name) else {
733 panic!("expected string array property {name}");
734 };
735 array.clone()
736 }
737
738 fn tensor_property(object: &ObjectInstance, name: &str) -> Tensor {
739 let Some(Value::Tensor(tensor)) = object.properties.get(name) else {
740 panic!("expected tensor property {name}");
741 };
742 tensor.clone()
743 }
744
745 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
746 #[test]
747 fn counts_default_bigrams_from_tokenized_documents() {
748 let bag = object(
749 run(vec![tokenized(vec![
750 vec!["a", "b", "a"],
751 vec!["a", "b", "c"],
752 ])])
753 .expect("bag"),
754 );
755 assert_eq!(bag.class_name, BAG_OF_NGRAMS_CLASS);
756 assert_eq!(bag.properties.get("NumDocuments"), Some(&Value::Num(2.0)));
757 assert_eq!(bag.properties.get("NumNgrams"), Some(&Value::Num(3.0)));
758
759 let ngrams = string_array_property(&bag, "Ngrams");
760 assert_eq!(ngrams.shape, vec![3, 2]);
761 assert_eq!(
762 ngrams.data,
763 vec!["a", "b", "b", "b", "a", "c"]
764 .into_iter()
765 .map(str::to_string)
766 .collect::<Vec<_>>()
767 );
768 let counts = tensor_property(&bag, "Counts");
769 assert_eq!(counts.shape, vec![2, 3]);
770 assert_eq!(counts.data, vec![1.0, 1.0, 1.0, 0.0, 0.0, 1.0]);
771 }
772
773 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
774 #[test]
775 fn accepts_ngram_lengths_vector() {
776 let lengths = Tensor::new(vec![1.0, 3.0], vec![1, 2]).unwrap();
777 let bag = object(
778 run(vec![
779 Value::StringArray(
780 StringArray::new(vec!["a".into(), "b".into(), "c".into()], vec![1, 3]).unwrap(),
781 ),
782 Value::String("NgramLengths".to_string()),
783 Value::Tensor(lengths),
784 ])
785 .expect("bag"),
786 );
787 let ngram_lengths = tensor_property(&bag, "NgramLengths");
788 assert_eq!(ngram_lengths.data, vec![1.0, 3.0]);
789 assert_eq!(bag.properties.get("NumNgrams"), Some(&Value::Num(4.0)));
790 }
791
792 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
793 #[test]
794 fn accepts_unique_ngram_matrix_and_counts() {
795 let ngrams = StringArray::new(
796 vec!["a".into(), "b".into(), "b".into(), "c".into()],
797 vec![2, 2],
798 )
799 .unwrap();
800 let counts = Tensor::new(vec![2.0, 0.0, 1.0, 3.0], vec![2, 2]).unwrap();
801 let bag =
802 object(run(vec![Value::StringArray(ngrams), Value::Tensor(counts)]).expect("bag"));
803 assert_eq!(bag.properties.get("NumDocuments"), Some(&Value::Num(2.0)));
804 assert_eq!(bag.properties.get("NumNgrams"), Some(&Value::Num(2.0)));
805 assert_eq!(tensor_property(&bag, "NgramLengths").data, vec![2.0]);
806 }
807
808 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
809 #[test]
810 fn infers_and_filters_unique_ngram_lengths() {
811 let ngrams = StringArray::new(
812 vec![
813 "a".into(),
814 "b".into(),
815 "d".into(),
816 "".into(),
817 "c".into(),
818 "e".into(),
819 ],
820 vec![3, 2],
821 )
822 .unwrap();
823 let counts = Tensor::new(vec![2.0, 0.0, 4.0], vec![1, 3]).unwrap();
824 let bag = object(
825 run(vec![
826 Value::StringArray(ngrams.clone()),
827 Value::Tensor(counts.clone()),
828 ])
829 .expect("bag"),
830 );
831 assert_eq!(tensor_property(&bag, "NgramLengths").data, vec![1.0, 2.0]);
832
833 let filtered = object(
834 run(vec![
835 Value::StringArray(ngrams),
836 Value::Tensor(counts),
837 Value::String("NgramLengths".to_string()),
838 Value::Num(2.0),
839 ])
840 .expect("bag"),
841 );
842 assert_eq!(filtered.properties.get("NumNgrams"), Some(&Value::Num(2.0)));
843 assert_eq!(tensor_property(&filtered, "NgramLengths").data, vec![2.0]);
844 assert_eq!(tensor_property(&filtered, "Counts").data, vec![0.0, 4.0]);
845 }
846
847 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
848 #[test]
849 fn accepts_cell_unique_ngram_matrix_in_column_major_order() {
850 let cells = vec![
851 Value::String("a".into()),
852 Value::String("c".into()),
853 Value::String("e".into()),
854 Value::String("b".into()),
855 Value::String("d".into()),
856 Value::String("f".into()),
857 ];
858 let ngrams = CellArray::new(cells, 3, 2).unwrap();
859 let counts = Tensor::new(vec![2.0, 3.0, 4.0], vec![1, 3]).unwrap();
860 let bag = object(run(vec![Value::Cell(ngrams), Value::Tensor(counts)]).expect("bag"));
861 let ngrams = string_array_property(&bag, "Ngrams");
862 assert_eq!(
863 ngrams.data,
864 vec!["a", "c", "e", "b", "d", "f"]
865 .into_iter()
866 .map(str::to_string)
867 .collect::<Vec<_>>()
868 );
869 }
870
871 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
872 #[test]
873 fn rejects_non_tokenized_column_word_vectors() {
874 let err = run(vec![Value::StringArray(
875 StringArray::new(vec!["a".into(), "b".into()], vec![2, 1]).unwrap(),
876 )])
877 .expect_err("expected column word vector rejection");
878 assert!(
879 err.to_string().contains("row word vector"),
880 "unexpected error: {err}"
881 );
882 }
883
884 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
885 #[test]
886 fn reports_bag_of_ngrams_identifier() {
887 let err = run(vec![
888 Value::String("a b c".to_string()),
889 Value::String("NgramLengths".to_string()),
890 Value::Num(0.0),
891 ])
892 .expect_err("expected bad length rejection");
893 assert_eq!(
894 err.identifier.as_deref(),
895 Some("RunMat:bagOfNgrams:InvalidInput")
896 );
897 }
898
899 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
900 #[test]
901 fn drops_missing_unique_ngram_and_count_column() {
902 let ngrams = StringArray::new(
903 vec!["a".into(), "<missing>".into(), "b".into(), "ignored".into()],
904 vec![2, 2],
905 )
906 .unwrap();
907 let counts = Tensor::new(vec![4.0, 9.0, 5.0, 9.0], vec![2, 2]).unwrap();
908 let bag =
909 object(run(vec![Value::StringArray(ngrams), Value::Tensor(counts)]).expect("bag"));
910 assert_eq!(bag.properties.get("NumNgrams"), Some(&Value::Num(1.0)));
911 let counts = tensor_property(&bag, "Counts");
912 assert_eq!(counts.shape, vec![2, 1]);
913 assert_eq!(counts.data, vec![4.0, 9.0]);
914 }
915
916 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
917 #[test]
918 fn rejects_duplicate_unique_ngrams_and_bad_lengths() {
919 let ngrams = StringArray::new(
920 vec!["a".into(), "a".into(), "b".into(), "b".into()],
921 vec![2, 2],
922 )
923 .unwrap();
924 let counts = Tensor::new(vec![1.0, 1.0], vec![1, 2]).unwrap();
925 let err = run(vec![Value::StringArray(ngrams), Value::Tensor(counts)])
926 .expect_err("expected duplicate ngram rejection");
927 assert!(err.to_string().contains("duplicate"));
928
929 let err = run(vec![
930 Value::String("a b c".to_string()),
931 Value::String("NgramLengths".to_string()),
932 Value::Num(0.0),
933 ])
934 .expect_err("expected bad length rejection");
935 assert!(err.to_string().contains("positive integers"));
936 }
937}