Skip to main content

runmat_runtime/builtins/strings/text_analytics/
ngrams.rs

1//! Bag-of-n-grams compatibility object for Text Analytics workflows.
2
3use 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}