Skip to main content

runmat_runtime/builtins/strings/text_analytics/
encode.rs

1//! Count-matrix encoding for Text Analytics bag models.
2
3use std::collections::{BTreeMap, HashMap, HashSet};
4
5use runmat_builtins::{
6    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinExtensionDescriptor,
7    BuiltinExtensionMode, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
8    BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
9    BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
10    BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule, BuiltinOutputMode,
11    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
12    ResolveContext, Type,
13};
14use runmat_macros::runtime_builtin;
15use runmat_value::{CellArray, ObjectInstance, SparseTensor, Value};
16
17use crate::builtins::common::spec::{
18    BroadcastSemantics, BuiltinGpuSpec, ConstantStrategy, GpuOpKind, ReductionNaN, ResidencyPolicy,
19};
20use crate::builtins::strings::core::compat::scalar_text;
21use crate::builtins::strings::text_analytics::documents::{
22    documents_from_object, vocabulary_from_bag, words_from_word_vector, BAG_OF_WORDS_CLASS,
23    TOKENIZED_DOCUMENT_CLASS,
24};
25use crate::builtins::strings::text_analytics::ngrams::{ngrams_from_bag, BAG_OF_NGRAMS_CLASS};
26use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult};
27
28#[runmat_macros::register_gpu_spec(
29    builtin_path = "crate::builtins::strings::text_analytics::encode"
30)]
31pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
32    name: "encode",
33    op_kind: GpuOpKind::Custom("text-analytics-encode"),
34    supported_precisions: &[],
35    broadcast: BroadcastSemantics::None,
36    provider_hooks: &[],
37    constant_strategy: ConstantStrategy::InlineLiteral,
38    residency: ResidencyPolicy::NewHandle,
39    nan_mode: ReductionNaN::Include,
40    two_pass_threshold: None,
41    workgroup_size: None,
42    accepts_nan_mode: false,
43    notes: "The builtin owns resident arguments so object/text rejection and ForceCellOutput compatibility gates run before provider access; admitted scalar controls gather explicitly and count outputs remain host sparse values.",
44};
45
46const OUT_COUNTS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
47    name: "counts",
48    ty: BuiltinParamType::Any,
49    arity: BuiltinParamArity::Required,
50    default: None,
51    description: "Sparse word or n-gram count matrix.",
52}];
53
54const IN_BAG_INPUT_REST: [BuiltinParamDescriptor; 3] = [
55    BuiltinParamDescriptor {
56        name: "bag",
57        ty: BuiltinParamType::Any,
58        arity: BuiltinParamArity::Required,
59        default: None,
60        description: "bagOfWords or bagOfNgrams model.",
61    },
62    BuiltinParamDescriptor {
63        name: "documentsOrWords",
64        ty: BuiltinParamType::Any,
65        arity: BuiltinParamArity::Required,
66        default: None,
67        description: "tokenizedDocument object or row word vector.",
68    },
69    BuiltinParamDescriptor {
70        name: "NameValue",
71        ty: BuiltinParamType::Any,
72        arity: BuiltinParamArity::Variadic,
73        default: None,
74        description: "Name-value options: DocumentsIn, ForceCellOutput.",
75    },
76];
77
78const ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
79    code: "RM.ENCODE.INVALID_INPUT",
80    identifier: Some("RunMat:encode:InvalidInput"),
81    when: "Inputs do not match a supported Text Analytics encode form.",
82    message: "encode: invalid input",
83};
84
85const ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_INVALID_INPUT];
86
87pub const ENCODE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
88    signatures: &[BuiltinSignatureDescriptor {
89        label: "counts = encode(bag, documentsOrWords, Name, Value, ...)",
90        inputs: &IN_BAG_INPUT_REST,
91        outputs: &OUT_COUNTS,
92    }],
93    output_mode: BuiltinOutputMode::Fixed,
94    completion_policy: BuiltinCompletionPolicy::Public,
95    errors: &ERRORS,
96};
97
98const ENCODE_NUMERIC_FORCE_CELL_OUTPUT_EXTENSION: BuiltinExtensionDescriptor =
99    BuiltinExtensionDescriptor {
100        id: "encode-numeric-force-cell-output",
101        mode: BuiltinExtensionMode::RunMatOnly,
102        description: "encode with a numeric ForceCellOutput value is a RunMat extension",
103        error_identifier: Some("RunMat:compatibility:EncodeNumericForceCellOutputExtension"),
104    };
105
106const ENCODE_RESIDENT_FORCE_CELL_OUTPUT_EXTENSION: BuiltinExtensionDescriptor =
107    BuiltinExtensionDescriptor {
108        id: "encode-resident-force-cell-output",
109        mode: BuiltinExtensionMode::RunMatOnly,
110        description: "encode with a resident ForceCellOutput value is a RunMat extension",
111        error_identifier: Some("RunMat:compatibility:EncodeResidentForceCellOutputExtension"),
112    };
113
114pub const ENCODE_EXTENSIONS: [BuiltinExtensionDescriptor; 2] = [
115    ENCODE_NUMERIC_FORCE_CELL_OUTPUT_EXTENSION,
116    ENCODE_RESIDENT_FORCE_CELL_OUTPUT_EXTENSION,
117];
118
119const ENCODE_REJECTED_INTEGER_DATA_INPUTS: [BuiltinIntegerInputCapability; 2] = [
120    BuiltinIntegerInputCapability {
121        name: "bag",
122        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
123        availability: BuiltinIntegerInputAvailability::Rejected,
124        scalar_double: BuiltinIntegerScalarDoubleRule::Rejected,
125        notes: "The model role requires a bagOfWords or bagOfNgrams object; integer values are not model payloads and reject before provider access.",
126    },
127    BuiltinIntegerInputCapability {
128        name: "documentsOrWords",
129        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
130        availability: BuiltinIntegerInputAvailability::Rejected,
131        scalar_double: BuiltinIntegerScalarDoubleRule::Rejected,
132        notes: "The document role is tokenizedDocument or text; integer values are not converted to words and reject before provider access.",
133    },
134];
135
136const ENCODE_INTEGER_FORCE_CELL_INPUTS: [BuiltinIntegerInputCapability; 1] =
137    [BuiltinIntegerInputCapability {
138        name: "ForceCellOutput",
139        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
140        availability: BuiltinIntegerInputAvailability::RunMatOnly,
141        scalar_double: BuiltinIntegerScalarDoubleRule::Allowed,
142        notes: "RunMat mode accepts exact scalar integer zero as false and every nonzero integer as true; MATLAB-compatible mode requires logical.",
143    }];
144
145pub const ENCODE_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 2] = [
146    BuiltinIntegerCapabilityDescriptor {
147        form: "counts = encode(bag, documentsOrWords)",
148        inputs: &ENCODE_REJECTED_INTEGER_DATA_INPUTS,
149        computation_domain: BuiltinIntegerComputationDomain::FunctionSpecific,
150        output_class: BuiltinIntegerOutputClassRule::NotApplicable,
151        overflow: BuiltinIntegerOverflowRule::NotApplicable,
152        backend: BuiltinIntegerBackendRule::HostOnly,
153        overload: BuiltinIntegerOverloadKind::FunctionSpecific,
154        notes: "encode is object/text based. Its integer-valued counts intentionally cross the documented sparse-double output boundary rather than using integer sparse storage.",
155    },
156    BuiltinIntegerCapabilityDescriptor {
157        form: "counts = encode(___, 'ForceCellOutput', integer_value)",
158        inputs: &ENCODE_INTEGER_FORCE_CELL_INPUTS,
159        computation_domain: BuiltinIntegerComputationDomain::Predicate,
160        output_class: BuiltinIntegerOutputClassRule::NotApplicable,
161        overflow: BuiltinIntegerOverflowRule::NotApplicable,
162        backend: BuiltinIntegerBackendRule::GatherFallback,
163        overload: BuiltinIntegerOverloadKind::ScalarOnly,
164        notes: "The compatibility-gated integer control only selects sparse-double versus cell-of-sparse-double representation; it never changes count storage.",
165    },
166];
167
168fn any_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
169    Type::Unknown
170}
171
172fn encode_error(message: impl Into<String>) -> crate::RuntimeError {
173    let mut builder = build_runtime_error(message).with_builtin("encode");
174    if let Some(identifier) = ERROR_INVALID_INPUT.identifier {
175        builder = builder.with_identifier(identifier);
176    }
177    builder.build()
178}
179
180#[runtime_builtin(
181    name = "encode",
182    category = "strings/text_analytics",
183    summary = "Encode documents as sparse word or n-gram count matrices.",
184    keywords = "encode,text analytics,bagOfWords,bagOfNgrams,count matrix",
185    accel = "sink",
186    type_resolver(any_type),
187    descriptor(crate::builtins::strings::text_analytics::encode::ENCODE_DESCRIPTOR),
188    extensions(crate::builtins::strings::text_analytics::encode::ENCODE_EXTENSIONS),
189    integer_capabilities(
190        crate::builtins::strings::text_analytics::encode::ENCODE_INTEGER_CAPABILITIES
191    ),
192    builtin_path = "crate::builtins::strings::text_analytics::encode"
193)]
194async fn encode_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
195    let (bag, input, options) = parse_args(args).await?;
196    let sparse = match bag {
197        Value::Object(object) if object.is_class(BAG_OF_WORDS_CLASS) => {
198            encode_words(&object, input, options.documents_in)?
199        }
200        Value::Object(object) if object.is_class(BAG_OF_NGRAMS_CLASS) => {
201            encode_ngrams(&object, input, options.documents_in)?
202        }
203        Value::Object(object) => {
204            return Err(encode_error(format!(
205                "encode: expected bagOfWords or bagOfNgrams object, got {}",
206                object.class_name
207            )))
208        }
209        other => {
210            return Err(encode_error(format!(
211                "encode: expected bagOfWords or bagOfNgrams object, got {other:?}"
212            )))
213        }
214    };
215
216    let output = Value::SparseTensor(sparse);
217    if options.force_cell_output {
218        return CellArray::new(vec![output], 1, 1)
219            .map(Value::Cell)
220            .map_err(encode_error);
221    }
222    Ok(output)
223}
224
225#[derive(Clone, Copy)]
226enum DocumentsIn {
227    Rows,
228    Columns,
229}
230
231struct EncodeOptions {
232    documents_in: DocumentsIn,
233    force_cell_output: bool,
234}
235
236impl Default for EncodeOptions {
237    fn default() -> Self {
238        Self {
239            documents_in: DocumentsIn::Rows,
240            force_cell_output: false,
241        }
242    }
243}
244
245async fn parse_args(mut args: Vec<Value>) -> BuiltinResult<(Value, Value, EncodeOptions)> {
246    if args.len() < 2 {
247        return Err(encode_error(
248            "encode: expected bag model and documents or words input",
249        ));
250    }
251    if !(args.len() - 2).is_multiple_of(2) {
252        return Err(encode_error(
253            "encode: name-value options must appear in pairs",
254        ));
255    }
256    let bag = args.remove(0);
257    let input = args.remove(0);
258    if crate::dispatcher::value_contains_gpu(&bag) {
259        return Err(encode_error(
260            "encode: bag model must be a host bagOfWords or bagOfNgrams object",
261        ));
262    }
263    if crate::dispatcher::value_contains_gpu(&input) {
264        return Err(encode_error(
265            "encode: documents or words must be host text or tokenizedDocument values",
266        ));
267    }
268    match &bag {
269        Value::Object(object)
270            if object.is_class(BAG_OF_WORDS_CLASS) || object.is_class(BAG_OF_NGRAMS_CLASS) => {}
271        Value::Object(object) => {
272            return Err(encode_error(format!(
273                "encode: expected bagOfWords or bagOfNgrams object, got {}",
274                object.class_name
275            )))
276        }
277        other => {
278            return Err(encode_error(format!(
279                "encode: expected bagOfWords or bagOfNgrams object, got {other:?}"
280            )))
281        }
282    }
283    validate_documents_outer_type(&input)?;
284    let mut options = EncodeOptions::default();
285    let mut idx = 0;
286    while idx < args.len() {
287        let name =
288            scalar_text(&args[idx], "encode").map_err(|err| encode_error(err.to_string()))?;
289        match name.to_ascii_lowercase().as_str() {
290            "documentsin" => {
291                let value = scalar_text(&args[idx + 1], "encode")
292                    .map_err(|err| encode_error(err.to_string()))?;
293                options.documents_in = match value.to_ascii_lowercase().as_str() {
294                    "rows" => DocumentsIn::Rows,
295                    "columns" => DocumentsIn::Columns,
296                    other => {
297                        return Err(encode_error(format!(
298                            "encode: DocumentsIn must be 'rows' or 'columns', got '{other}'"
299                        )))
300                    }
301                };
302            }
303            "forcecelloutput" => {
304                let raw = &args[idx + 1];
305                let resident = crate::dispatcher::value_contains_gpu(raw);
306                if resident {
307                    validate_scalar_control_shape(raw)?;
308                    crate::compatibility::ensure_builtin_extension_enabled(
309                        &ENCODE_RESIDENT_FORCE_CELL_OUTPUT_EXTENSION,
310                        "encode",
311                    )?;
312                }
313                let host = if resident {
314                    gather_if_needed_async(raw).await.map_err(|err| {
315                        encode_error(format!("encode: failed to gather ForceCellOutput: {err}"))
316                    })?
317                } else {
318                    raw.clone()
319                };
320                let parsed = parse_bool_scalar(&host)?;
321                if is_numeric_bool_value(&host) {
322                    crate::compatibility::ensure_builtin_extension_enabled(
323                        &ENCODE_NUMERIC_FORCE_CELL_OUTPUT_EXTENSION,
324                        "encode",
325                    )?;
326                }
327                options.force_cell_output = parsed;
328            }
329            other => {
330                return Err(encode_error(format!(
331                    "encode: unsupported option '{other}'"
332                )))
333            }
334        }
335        idx += 2;
336    }
337    Ok((bag, input, options))
338}
339
340fn validate_documents_outer_type(value: &Value) -> BuiltinResult<()> {
341    match value {
342        Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => Ok(()),
343        Value::String(_) | Value::StringArray(_) | Value::CharArray(_) | Value::Cell(_) => Ok(()),
344        other => Err(encode_error(format!(
345            "encode: expected tokenizedDocument or word vector, got {other:?}"
346        ))),
347    }
348}
349
350fn validate_scalar_control_shape(value: &Value) -> BuiltinResult<()> {
351    let len = match value {
352        Value::GpuTensor(handle) => handle
353            .shape
354            .iter()
355            .try_fold(1usize, |total, dimension| total.checked_mul(*dimension))
356            .unwrap_or(usize::MAX),
357        Value::Tensor(tensor) => tensor.len(),
358        Value::LogicalArray(array) => array.data.len(),
359        _ => 1,
360    };
361    if len != 1 {
362        return Err(encode_error(
363            "encode: ForceCellOutput must be a logical scalar",
364        ));
365    }
366    Ok(())
367}
368
369fn is_numeric_bool_value(value: &Value) -> bool {
370    matches!(value, Value::Num(_) | Value::Int(_) | Value::Tensor(_))
371}
372
373fn parse_bool_scalar(value: &Value) -> BuiltinResult<bool> {
374    match value {
375        Value::Bool(value) => Ok(*value),
376        Value::LogicalArray(array) if array.data.len() == 1 => Ok(array.data[0] != 0),
377        Value::Num(value) if *value == 0.0 || *value == 1.0 => Ok(*value != 0.0),
378        Value::Int(value) => Ok(!value.is_zero()),
379        Value::Tensor(tensor) if tensor.len() == 1 => {
380            if let Some(value) = tensor
381                .integer_storage()
382                .and_then(|storage| storage.value_at(0))
383            {
384                return Ok(!value.is_zero());
385            }
386            let value = crate::builtins::common::tensor::tensor_value_f64(tensor, 0);
387            if value == 0.0 || value == 1.0 {
388                Ok(value != 0.0)
389            } else {
390                Err(encode_error(
391                    "encode: ForceCellOutput numeric values must be 0 or 1",
392                ))
393            }
394        }
395        other => Err(encode_error(format!(
396            "encode: ForceCellOutput must be a logical scalar, got {other:?}"
397        ))),
398    }
399}
400
401fn encode_words(
402    object: &ObjectInstance,
403    input: Value,
404    documents_in: DocumentsIn,
405) -> BuiltinResult<SparseTensor> {
406    let vocabulary = vocabulary_from_bag(object, "encode").map_err(|err| {
407        encode_error(format!(
408            "encode: failed to read bagOfWords Vocabulary property: {err}"
409        ))
410    })?;
411    let documents = documents_from_input(input, "bagOfWords")?;
412    let positions = vocabulary
413        .iter()
414        .enumerate()
415        .map(|(idx, word)| (word.as_str(), idx))
416        .collect::<HashMap<_, _>>();
417    let counts = documents
418        .iter()
419        .map(|document| {
420            let mut row = BTreeMap::new();
421            for token in document {
422                if let Some(&col) = positions.get(token.as_str()) {
423                    *row.entry(col).or_insert(0.0) += 1.0;
424                }
425            }
426            row
427        })
428        .collect::<Vec<_>>();
429    sparse_from_document_counts(counts, vocabulary.len(), documents_in)
430}
431
432fn encode_ngrams(
433    object: &ObjectInstance,
434    input: Value,
435    documents_in: DocumentsIn,
436) -> BuiltinResult<SparseTensor> {
437    let ngrams = ngrams_from_bag(object, "encode").map_err(|err| {
438        encode_error(format!(
439            "encode: failed to read bagOfNgrams Ngrams property: {err}"
440        ))
441    })?;
442    let lengths = unique_ngram_lengths(&ngrams);
443    let documents = documents_from_input(input, "bagOfNgrams")?;
444    let positions = ngrams
445        .iter()
446        .enumerate()
447        .map(|(idx, ngram)| (ngram.as_slice(), idx))
448        .collect::<HashMap<_, _>>();
449    let counts = documents
450        .iter()
451        .map(|document| {
452            let mut row = BTreeMap::new();
453            for &length in &lengths {
454                if length > document.len() {
455                    continue;
456                }
457                for start in 0..=document.len() - length {
458                    let key = &document[start..start + length];
459                    if let Some(&col) = positions.get(key) {
460                        *row.entry(col).or_insert(0.0) += 1.0;
461                    }
462                }
463            }
464            row
465        })
466        .collect::<Vec<_>>();
467    sparse_from_document_counts(counts, ngrams.len(), documents_in)
468}
469
470fn unique_ngram_lengths(ngrams: &[Vec<String>]) -> Vec<usize> {
471    let mut seen = HashSet::new();
472    let mut lengths = Vec::new();
473    for ngram in ngrams {
474        let length = ngram.len();
475        if seen.insert(length) {
476            lengths.push(length);
477        }
478    }
479    lengths
480}
481
482fn documents_from_input(input: Value, model_name: &str) -> BuiltinResult<Vec<Vec<String>>> {
483    match input {
484        Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => {
485            documents_from_object(&object, "encode").map_err(|err| {
486                encode_error(format!(
487                    "encode: failed to read tokenizedDocument input: {err}"
488                ))
489            })
490        }
491        Value::Object(object) => Err(encode_error(format!(
492            "encode: expected tokenizedDocument or word vector for {model_name}, got {}",
493            object.class_name
494        ))),
495        other => {
496            validate_row_word_vector(&other, model_name)?;
497            Ok(vec![words_from_word_vector(&other, "encode").map_err(
498                |err| encode_error(format!("encode: failed to read word vector input: {err}")),
499            )?])
500        }
501    }
502}
503
504fn validate_row_word_vector(value: &Value, model_name: &str) -> BuiltinResult<()> {
505    match value {
506        Value::String(_) => Ok(()),
507        Value::StringArray(array) if array.rows <= 1 => Ok(()),
508        Value::CharArray(array) if array.rows <= 1 => Ok(()),
509        Value::Cell(cell) if cell.rows <= 1 => Ok(()),
510        Value::StringArray(array) => Err(encode_error(format!(
511            "encode: non-tokenized {model_name} input must be a row word vector; got string array with shape {}x{}",
512            array.rows, array.cols
513        ))),
514        Value::CharArray(array) => Err(encode_error(format!(
515            "encode: non-tokenized {model_name} input must be a row word vector; got char array with shape {}x{}",
516            array.rows, array.cols
517        ))),
518        Value::Cell(cell) => Err(encode_error(format!(
519            "encode: non-tokenized {model_name} input must be a row word vector; got cell array with shape {}x{}",
520            cell.rows, cell.cols
521        ))),
522        other => Err(encode_error(format!(
523            "encode: expected tokenizedDocument or word vector for {model_name}, got {other:?}"
524        ))),
525    }
526}
527
528fn sparse_from_document_counts(
529    counts: Vec<BTreeMap<usize, f64>>,
530    term_count: usize,
531    documents_in: DocumentsIn,
532) -> BuiltinResult<SparseTensor> {
533    match documents_in {
534        DocumentsIn::Rows => sparse_rows(counts, term_count),
535        DocumentsIn::Columns => sparse_columns(counts, term_count),
536    }
537}
538
539fn sparse_rows(
540    counts: Vec<BTreeMap<usize, f64>>,
541    term_count: usize,
542) -> BuiltinResult<SparseTensor> {
543    let rows = counts.len();
544    let cols = term_count;
545    let col_ptr_capacity = cols
546        .checked_add(1)
547        .ok_or_else(|| encode_error("encode: sparse output column count overflows"))?;
548    let mut columns = vec![Vec::<(usize, f64)>::new(); cols];
549    for (doc_idx, doc_counts) in counts.iter().enumerate() {
550        for (&term_idx, &value) in doc_counts {
551            if term_idx >= term_count {
552                return Err(encode_error(
553                    "encode: internal sparse term index exceeds model size",
554                ));
555            }
556            if value != 0.0 {
557                columns[term_idx].push((doc_idx, value));
558            }
559        }
560    }
561    let mut col_ptrs = Vec::with_capacity(col_ptr_capacity);
562    let mut row_indices = Vec::new();
563    let mut values = Vec::new();
564    col_ptrs.push(0);
565    for entries in columns {
566        for (row, value) in entries {
567            row_indices.push(row);
568            values.push(value);
569        }
570        col_ptrs.push(values.len());
571    }
572    SparseTensor::new(rows, cols, col_ptrs, row_indices, values).map_err(encode_error)
573}
574
575fn sparse_columns(
576    counts: Vec<BTreeMap<usize, f64>>,
577    term_count: usize,
578) -> BuiltinResult<SparseTensor> {
579    let rows = term_count;
580    let cols = counts.len();
581    let col_ptr_capacity = cols
582        .checked_add(1)
583        .ok_or_else(|| encode_error("encode: sparse output column count overflows"))?;
584    let mut col_ptrs = Vec::with_capacity(col_ptr_capacity);
585    let mut row_indices = Vec::new();
586    let mut values = Vec::new();
587    col_ptrs.push(0);
588    for doc_counts in &counts {
589        for (&row, &value) in doc_counts {
590            if row >= term_count {
591                return Err(encode_error(
592                    "encode: internal sparse term index exceeds model size",
593                ));
594            }
595            if value != 0.0 {
596                row_indices.push(row);
597                values.push(value);
598            }
599        }
600        col_ptrs.push(values.len());
601    }
602    SparseTensor::new(rows, cols, col_ptrs, row_indices, values).map_err(encode_error)
603}
604
605#[cfg(test)]
606mod tests {
607    use super::*;
608    use runmat_value::{IntValue, IntegerStorage, StringArray, Tensor};
609
610    fn run_encode(args: Vec<Value>) -> BuiltinResult<Value> {
611        futures::executor::block_on(encode_builtin(args))
612    }
613
614    fn sparse(value: Value) -> SparseTensor {
615        match value {
616            Value::SparseTensor(sparse) => sparse,
617            other => panic!("expected sparse tensor, got {other:?}"),
618        }
619    }
620
621    fn string_array(values: &[&str], rows: usize, cols: usize) -> Value {
622        Value::StringArray(
623            StringArray::new(
624                values.iter().map(|value| (*value).to_string()).collect(),
625                vec![rows, cols],
626            )
627            .expect("string array"),
628        )
629    }
630
631    fn tokenized(docs: &[&[&str]]) -> Value {
632        let mut data = Vec::with_capacity(docs.len());
633        for doc in docs {
634            let row = doc
635                .iter()
636                .map(|token| Value::from(*token))
637                .collect::<Vec<_>>();
638            data.push(Value::Cell(
639                CellArray::new(row, 1, doc.len()).expect("row cell"),
640            ));
641        }
642        let mut object = ObjectInstance::new(TOKENIZED_DOCUMENT_CLASS.to_string());
643        object.properties.insert(
644            "Documents".to_string(),
645            Value::Cell(CellArray::new(data, docs.len(), 1).expect("documents cell")),
646        );
647        object
648            .properties
649            .insert("NumDocuments".to_string(), Value::Num(docs.len() as f64));
650        Value::Object(object)
651    }
652
653    fn bag_of_words(vocabulary: &[&str]) -> Value {
654        let mut object = ObjectInstance::new(BAG_OF_WORDS_CLASS.to_string());
655        object.properties.insert(
656            "Vocabulary".to_string(),
657            string_array(vocabulary, 1, vocabulary.len()),
658        );
659        object.properties.insert(
660            "Counts".to_string(),
661            Value::Tensor(Tensor::zeros(vec![0, vocabulary.len()])),
662        );
663        object
664            .properties
665            .insert("NumWords".to_string(), Value::Num(vocabulary.len() as f64));
666        object
667            .properties
668            .insert("NumDocuments".to_string(), Value::Num(0.0));
669        Value::Object(object)
670    }
671
672    fn bag_of_ngrams(ngrams: &[&[&str]], lengths: &[usize]) -> Value {
673        let rows = ngrams.len();
674        let cols = ngrams.iter().map(|ngram| ngram.len()).max().unwrap_or(0);
675        let mut data = Vec::with_capacity(rows * cols);
676        for col in 0..cols {
677            for ngram in ngrams {
678                data.push(ngram.get(col).copied().unwrap_or_default().to_string());
679            }
680        }
681        let mut object = ObjectInstance::new(BAG_OF_NGRAMS_CLASS.to_string());
682        object.properties.insert(
683            "Ngrams".to_string(),
684            Value::StringArray(StringArray::new(data, vec![rows, cols]).expect("ngrams")),
685        );
686        object.properties.insert(
687            "NgramLengths".to_string(),
688            Value::Tensor(
689                Tensor::new(
690                    lengths.iter().map(|length| *length as f64).collect(),
691                    vec![1, lengths.len()],
692                )
693                .expect("lengths"),
694            ),
695        );
696        object.properties.insert(
697            "Counts".to_string(),
698            Value::Tensor(Tensor::zeros(vec![0, ngrams.len()])),
699        );
700        object
701            .properties
702            .insert("NumNgrams".to_string(), Value::Num(ngrams.len() as f64));
703        object
704            .properties
705            .insert("NumDocuments".to_string(), Value::Num(0.0));
706        Value::Object(object)
707    }
708
709    #[test]
710    fn encodes_tokenized_documents_against_bag_of_words_rows() {
711        let bag = bag_of_words(&["alpha", "beta", "gamma"]);
712        let docs = tokenized(&[&["beta", "beta", "delta"], &["alpha", "gamma"]]);
713
714        let out = sparse(run_encode(vec![bag, docs]).expect("encode"));
715        assert_eq!((out.rows, out.cols), (2, 3));
716        let dense = out.to_dense().unwrap();
717        assert_eq!(dense.shape, vec![2, 3]);
718        assert_eq!(dense.materialize_f64(), vec![0.0, 1.0, 2.0, 0.0, 0.0, 1.0]);
719    }
720
721    #[test]
722    fn encodes_word_vector_with_documents_in_columns() {
723        let bag = bag_of_words(&["alpha", "beta", "gamma"]);
724
725        let out = sparse(
726            run_encode(vec![
727                bag,
728                string_array(&["beta", "gamma", "beta"], 1, 3),
729                Value::from("DocumentsIn"),
730                Value::from("columns"),
731            ])
732            .expect("encode"),
733        );
734        assert_eq!((out.rows, out.cols), (3, 1));
735        let dense = out.to_dense().unwrap();
736        assert_eq!(dense.shape, vec![3, 1]);
737        assert_eq!(dense.materialize_f64(), vec![0.0, 2.0, 1.0]);
738    }
739
740    #[test]
741    fn encodes_multiple_documents_in_columns() {
742        let bag = bag_of_words(&["alpha", "beta", "gamma"]);
743        let docs = tokenized(&[&["beta", "beta", "delta"], &["alpha", "gamma"]]);
744
745        let out = sparse(
746            run_encode(vec![
747                bag,
748                docs,
749                Value::from("DocumentsIn"),
750                Value::from("columns"),
751            ])
752            .expect("encode"),
753        );
754        assert_eq!((out.rows, out.cols), (3, 2));
755        let dense = out.to_dense().unwrap();
756        assert_eq!(dense.shape, vec![3, 2]);
757        assert_eq!(dense.materialize_f64(), vec![0.0, 2.0, 0.0, 1.0, 0.0, 1.0]);
758    }
759
760    #[test]
761    fn returns_sparse_zeros_for_empty_bag_and_unknown_terms() {
762        let empty = sparse(run_encode(vec![bag_of_words(&[]), tokenized(&[&["alpha"]])]).unwrap());
763        assert_eq!((empty.rows, empty.cols), (1, 0));
764        assert_eq!(empty.col_ptrs, vec![0]);
765        assert!(empty.row_indices.is_empty());
766        assert!(empty.materialize_f64().is_empty());
767
768        let unknown = sparse(
769            run_encode(vec![bag_of_words(&["alpha", "beta"]), Value::from("gamma")]).unwrap(),
770        );
771        assert_eq!((unknown.rows, unknown.cols), (1, 2));
772        assert_eq!(unknown.col_ptrs, vec![0, 0, 0]);
773        assert!(unknown.row_indices.is_empty());
774        assert!(unknown.materialize_f64().is_empty());
775    }
776
777    #[test]
778    fn force_cell_output_wraps_sparse_result() {
779        let bag = bag_of_words(&["alpha", "beta"]);
780        let out = run_encode(vec![
781            bag,
782            Value::from("alpha"),
783            Value::from("ForceCellOutput"),
784            Value::Bool(true),
785        ])
786        .expect("encode");
787        let Value::Cell(cell) = out else {
788            panic!("expected cell");
789        };
790        assert_eq!((cell.rows, cell.cols), (1, 1));
791        let Value::SparseTensor(sparse) = &cell.data[0] else {
792            panic!("expected sparse cell element");
793        };
794        assert_eq!((sparse.rows, sparse.cols), (1, 2));
795        let dense = sparse.to_dense().unwrap();
796        assert_eq!(dense.shape, vec![1, 2]);
797        assert_eq!(dense.materialize_f64(), vec![1.0, 0.0]);
798    }
799
800    #[test]
801    fn encodes_bag_of_ngrams_documents() {
802        let bag = bag_of_ngrams(&[&["a"], &["b"], &["a", "b"], &["b", "a"]], &[1, 2]);
803        let docs = tokenized(&[&["a", "b", "a", "b"]]);
804
805        let out = sparse(run_encode(vec![bag, docs]).expect("encode"));
806        assert_eq!(out.rows, 1);
807        assert_eq!(out.cols, 4);
808        let dense = out.to_dense().unwrap();
809        assert_eq!(dense.shape, vec![1, 4]);
810        assert_eq!(dense.materialize_f64(), vec![2.0, 2.0, 2.0, 1.0]);
811    }
812
813    #[test]
814    fn rejects_malformed_bag_of_ngrams_metadata() {
815        let err = run_encode(vec![bag_of_ngrams(&[&[]], &[]), tokenized(&[&["a"]])])
816            .expect_err("expected empty ngram rejection");
817        assert!(err.to_string().contains("empty n-gram"));
818
819        let err = run_encode(vec![
820            bag_of_ngrams(&[&["a"], &["a"]], &[1]),
821            tokenized(&[&["a"]]),
822        ])
823        .expect_err("expected duplicate ngram rejection");
824        assert!(err.to_string().contains("duplicate n-gram"));
825    }
826
827    #[test]
828    fn rejects_bad_options_and_column_word_vectors() {
829        let bag = bag_of_words(&["alpha"]);
830        let err = run_encode(vec![
831            bag.clone(),
832            Value::from("alpha"),
833            Value::from("DocumentsIn"),
834            Value::from("pages"),
835        ])
836        .expect_err("expected bad option");
837        assert!(err.to_string().contains("DocumentsIn"));
838
839        let err = run_encode(vec![bag, string_array(&["alpha", "beta"], 2, 1)])
840            .expect_err("expected column rejection");
841        assert!(err.to_string().contains("row word vector"));
842    }
843
844    #[test]
845    fn rejects_invalid_force_cell_output_and_odd_options() {
846        let bag = bag_of_words(&["alpha"]);
847        let err = run_encode(vec![
848            bag.clone(),
849            Value::from("alpha"),
850            Value::from("ForceCellOutput"),
851            Value::from("yes"),
852        ])
853        .expect_err("expected invalid force cell output");
854        assert!(err.to_string().contains("ForceCellOutput"));
855
856        let err = run_encode(vec![bag, Value::from("alpha"), Value::from("DocumentsIn")])
857            .expect_err("expected odd options rejection");
858        assert!(err.to_string().contains("name-value options"));
859    }
860
861    #[test]
862    fn integer_force_cell_output_accepts_all_classes_exactly_in_runmat_mode() {
863        let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
864        for flag in [
865            IntValue::I8(-1),
866            IntValue::I16(1),
867            IntValue::I32(1),
868            IntValue::I64(i64::MAX),
869            IntValue::U8(1),
870            IntValue::U16(1),
871            IntValue::U32(1),
872            IntValue::U64(u64::MAX),
873        ] {
874            let out = run_encode(vec![
875                bag_of_words(&["alpha"]),
876                Value::from("alpha"),
877                Value::from("ForceCellOutput"),
878                Value::Int(flag),
879            ])
880            .unwrap();
881            assert!(matches!(out, Value::Cell(_)));
882        }
883        let zero = Tensor::new_integer(IntegerStorage::U64(vec![0]), vec![1, 1]).unwrap();
884        let out = run_encode(vec![
885            bag_of_words(&["alpha"]),
886            Value::from("alpha"),
887            Value::from("ForceCellOutput"),
888            Value::Tensor(zero),
889        ])
890        .unwrap();
891        assert!(matches!(out, Value::SparseTensor(_)));
892    }
893
894    #[test]
895    fn integer_force_cell_output_is_gated_in_matlab_mode() {
896        let _compat = crate::compatibility::push_runmat_extensions_enabled(false);
897        let err = run_encode(vec![
898            bag_of_words(&["alpha"]),
899            Value::from("alpha"),
900            Value::from("ForceCellOutput"),
901            Value::Int(IntValue::U64(u64::MAX)),
902        ])
903        .unwrap_err();
904        assert_eq!(
905            err.identifier(),
906            Some("RunMat:compatibility:EncodeNumericForceCellOutputExtension")
907        );
908    }
909
910    #[test]
911    fn resident_numeric_documents_reject_before_provider_access() {
912        let resident = Value::GpuTensor(runmat_accelerate_api::GpuTensorHandle {
913            shape: vec![1, 1],
914            device_id: u32::MAX,
915            buffer_id: u64::MAX,
916            descriptor: Default::default(),
917        });
918        let err = run_encode(vec![bag_of_words(&["alpha"]), resident]).unwrap_err();
919        assert_eq!(err.identifier(), Some("RunMat:encode:InvalidInput"));
920    }
921
922    #[test]
923    fn resident_force_cell_output_rejects_before_provider_access_in_matlab_mode() {
924        let _compat = crate::compatibility::push_runmat_extensions_enabled(false);
925        let resident = Value::GpuTensor(runmat_accelerate_api::GpuTensorHandle {
926            shape: vec![1, 1],
927            device_id: u32::MAX,
928            buffer_id: u64::MAX,
929            descriptor: Default::default(),
930        });
931        let err = run_encode(vec![
932            bag_of_words(&["alpha"]),
933            Value::from("alpha"),
934            Value::from("ForceCellOutput"),
935            resident,
936        ])
937        .unwrap_err();
938        assert_eq!(
939            err.identifier(),
940            Some("RunMat:compatibility:EncodeResidentForceCellOutputExtension")
941        );
942    }
943
944    #[test]
945    fn encode_dispatch_preserves_residency_until_builtin_preflight() {
946        assert_eq!(GPU_SPEC.residency, ResidencyPolicy::NewHandle);
947        let resident = Value::GpuTensor(runmat_accelerate_api::GpuTensorHandle {
948            shape: vec![1, 1],
949            device_id: u32::MAX,
950            buffer_id: u64::MAX - 2,
951            descriptor: Default::default(),
952        });
953        let prepared = futures::executor::block_on(runmat_accelerate::prepare_builtin_args(
954            "encode",
955            &[resident],
956        ))
957        .expect("dispatcher must retain resident argument");
958        assert!(matches!(prepared.as_slice(), [Value::GpuTensor(_)]));
959    }
960}