1use 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}