1use runmat_types::MemberAccess;
3
4use runmat_builtins::{
5 BuiltinExtensionDescriptor, BuiltinExtensionMode, BuiltinIntegerAuditDescriptor,
6 BuiltinIntegerAuditKind, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
7 BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
8 BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
9 BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule,
10};
11use runmat_value::IntValue;
12use std::collections::HashMap;
13
14use runmat_builtins::{
15 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
16 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
17 ResolveContext, Type,
18};
19use runmat_macros::runtime_builtin;
20use runmat_value::{CharArray, LogicalArray, ObjectInstance, StringArray, Tensor, Value};
21
22use crate::builtins::common::tensor as tensor_utils;
23use crate::builtins::strings::core::compat::scalar_text;
24use crate::builtins::strings::text_analytics::documents::{
25 documents_from_object, TOKENIZED_DOCUMENT_CLASS,
26};
27use crate::builtins::strings::text_analytics::embeddings::{
28 build_word_lookup, word_embedding_vocabulary_from_object, WORD_EMBEDDING_CLASS,
29};
30use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult};
31
32pub const WORD_ENCODING_CLASS: &str = "wordEncoding";
33
34const WORD_ENCODING_INTEGER_MAX_WORDS_EXTENSION: BuiltinExtensionDescriptor =
35 BuiltinExtensionDescriptor {
36 id: "wordencoding-integer-max-num-words",
37 mode: BuiltinExtensionMode::RunMatOnly,
38 description: "wordEncoding with a typed-integer MaxNumWords value is a RunMat extension",
39 error_identifier: Some("RunMat:compatibility:WordEncodingIntegerMaxNumWordsExtension"),
40 };
41pub const WORD_ENCODING_EXTENSIONS: [BuiltinExtensionDescriptor; 1] =
42 [WORD_ENCODING_INTEGER_MAX_WORDS_EXTENSION];
43const WORD_ENCODING_INTEGER_MAX_WORDS_INPUT: [BuiltinIntegerInputCapability; 1] =
44 [BuiltinIntegerInputCapability {
45 name: "MaxNumWords",
46 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
47 availability: BuiltinIntegerInputAvailability::RunMatOnly,
48 scalar_double: BuiltinIntegerScalarDoubleRule::Allowed,
49 notes: "The public reference specifies a positive integer value or Inf without publishing native integer storage classes. RunMat mode decodes a typed scalar exactly as a bounded vocabulary length.",
50 }];
51pub const WORD_ENCODING_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
52 [BuiltinIntegerCapabilityDescriptor {
53 form: "enc = wordEncoding(documents, 'MaxNumWords', integer_n)",
54 inputs: &WORD_ENCODING_INTEGER_MAX_WORDS_INPUT,
55 computation_domain: BuiltinIntegerComputationDomain::Structural,
56 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
57 overflow: BuiltinIntegerOverflowRule::Error,
58 backend: BuiltinIntegerBackendRule::HostOnly,
59 overload: BuiltinIntegerOverloadKind::StructuralParameter,
60 notes: "The exact positive count truncates the ranked host vocabulary and does not enter floating arithmetic. Ordinary positive integer-valued double and positive Inf retain their documented behavior.",
61 }];
62
63const IND2WORD_TYPED_INTEGER_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
64 id: "ind2word-typed-integer-indices",
65 mode: BuiltinExtensionMode::RunMatOnly,
66 description: "ind2word with a typed-integer index vector is a RunMat extension",
67 error_identifier: Some("RunMat:compatibility:Ind2wordTypedIntegerExtension"),
68};
69const IND2WORD_NONVECTOR_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
70 id: "ind2word-nonvector-indices",
71 mode: BuiltinExtensionMode::RunMatOnly,
72 description: "ind2word with matrix or multidimensional indices is a RunMat extension",
73 error_identifier: Some("RunMat:compatibility:Ind2wordNonvectorExtension"),
74};
75const IND2WORD_RESIDENT_EXTENSION: BuiltinExtensionDescriptor = BuiltinExtensionDescriptor {
76 id: "ind2word-resident-indices",
77 mode: BuiltinExtensionMode::RunMatOnly,
78 description: "ind2word with resident indices is a RunMat extension",
79 error_identifier: Some("RunMat:compatibility:Ind2wordResidentExtension"),
80};
81pub const IND2WORD_EXTENSIONS: [BuiltinExtensionDescriptor; 3] = [
82 IND2WORD_TYPED_INTEGER_EXTENSION,
83 IND2WORD_NONVECTOR_EXTENSION,
84 IND2WORD_RESIDENT_EXTENSION,
85];
86const IND2WORD_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 1] =
87 [BuiltinIntegerInputCapability {
88 name: "M",
89 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
90 availability: BuiltinIntegerInputAvailability::RunMatOnly,
91 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
92 notes: "The current public reference specifies positive integer values but does not publish a native numeric class table; typed-integer vectors are therefore gated and read from authoritative integer storage.",
93 }];
94pub const IND2WORD_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
95 [BuiltinIntegerCapabilityDescriptor {
96 form: "words = ind2word(enc, integer_M)",
97 inputs: &IND2WORD_INTEGER_INPUTS,
98 computation_domain: BuiltinIntegerComputationDomain::Structural,
99 output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
100 overflow: BuiltinIntegerOverflowRule::Error,
101 backend: BuiltinIntegerBackendRule::GatherFallback,
102 overload: BuiltinIntegerOverloadKind::Multiple,
103 notes: "The public form uses a host positive-integer vector and returns a string vector. RunMat gates typed-integer and resident forms independently, reads typed indices exactly from their native class, and always returns host strings.",
104 }];
105
106static WORD_ENCODING_CLASS_REGISTERED: crate::class_registry::ClassRegistration =
107 crate::class_registry::ClassRegistration::new(WORD_ENCODING_CLASS);
108
109const OUT_ENCODING: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
110 name: "enc",
111 ty: BuiltinParamType::Any,
112 arity: BuiltinParamArity::Required,
113 default: None,
114 description: "Word encoding compatibility object.",
115}];
116
117const OUT_INDICES: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
118 name: "M",
119 ty: BuiltinParamType::NumericArray,
120 arity: BuiltinParamArity::Required,
121 default: None,
122 description: "Word encoding indices, with NaN for words outside the vocabulary.",
123}];
124
125const OUT_WORDS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
126 name: "words",
127 ty: BuiltinParamType::Any,
128 arity: BuiltinParamArity::Required,
129 default: None,
130 description: "Words mapped from encoding indices.",
131}];
132
133const OUT_LOGICAL: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
134 name: "tf",
135 ty: BuiltinParamType::LogicalArray,
136 arity: BuiltinParamArity::Required,
137 default: None,
138 description: "Logical membership mask.",
139}];
140
141const IN_DOCUMENTS_OR_WORDS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
142 name: "documentsOrWords",
143 ty: BuiltinParamType::Any,
144 arity: BuiltinParamArity::Required,
145 default: None,
146 description: "tokenizedDocument object or word vector.",
147}];
148
149const IN_DOCUMENTS_OR_WORDS_REST: [BuiltinParamDescriptor; 2] = [
150 BuiltinParamDescriptor {
151 name: "documentsOrWords",
152 ty: BuiltinParamType::Any,
153 arity: BuiltinParamArity::Required,
154 default: None,
155 description: "tokenizedDocument object or word vector.",
156 },
157 BuiltinParamDescriptor {
158 name: "NameValue",
159 ty: BuiltinParamType::Any,
160 arity: BuiltinParamArity::Variadic,
161 default: None,
162 description: "Name-value options: Order, MaxNumWords.",
163 },
164];
165
166const IN_WORDS: [BuiltinParamDescriptor; 2] = [
167 BuiltinParamDescriptor {
168 name: "enc",
169 ty: BuiltinParamType::Any,
170 arity: BuiltinParamArity::Required,
171 default: None,
172 description: "wordEncoding object.",
173 },
174 BuiltinParamDescriptor {
175 name: "words",
176 ty: BuiltinParamType::Any,
177 arity: BuiltinParamArity::Required,
178 default: None,
179 description: "Words to map to indices.",
180 },
181];
182
183const IN_WORDS_REST: [BuiltinParamDescriptor; 3] = [
184 BuiltinParamDescriptor {
185 name: "enc",
186 ty: BuiltinParamType::Any,
187 arity: BuiltinParamArity::Required,
188 default: None,
189 description: "wordEncoding object.",
190 },
191 BuiltinParamDescriptor {
192 name: "words",
193 ty: BuiltinParamType::Any,
194 arity: BuiltinParamArity::Required,
195 default: None,
196 description: "Words to map to indices.",
197 },
198 BuiltinParamDescriptor {
199 name: "NameValue",
200 ty: BuiltinParamType::Any,
201 arity: BuiltinParamArity::Variadic,
202 default: None,
203 description: "Name-value options: IgnoreCase.",
204 },
205];
206
207const IN_INDICES: [BuiltinParamDescriptor; 2] = [
208 BuiltinParamDescriptor {
209 name: "enc",
210 ty: BuiltinParamType::Any,
211 arity: BuiltinParamArity::Required,
212 default: None,
213 description: "wordEncoding object.",
214 },
215 BuiltinParamDescriptor {
216 name: "M",
217 ty: BuiltinParamType::NumericArray,
218 arity: BuiltinParamArity::Required,
219 default: None,
220 description: "Positive integer word encoding indices.",
221 },
222];
223
224const IN_VOCABULARY_WORDS: [BuiltinParamDescriptor; 2] = [
225 BuiltinParamDescriptor {
226 name: "embOrEnc",
227 ty: BuiltinParamType::Any,
228 arity: BuiltinParamArity::Required,
229 default: None,
230 description: "wordEmbedding or wordEncoding object.",
231 },
232 BuiltinParamDescriptor {
233 name: "words",
234 ty: BuiltinParamType::Any,
235 arity: BuiltinParamArity::Required,
236 default: None,
237 description: "Words to test.",
238 },
239];
240
241const IN_VOCABULARY_WORDS_REST: [BuiltinParamDescriptor; 3] = [
242 BuiltinParamDescriptor {
243 name: "embOrEnc",
244 ty: BuiltinParamType::Any,
245 arity: BuiltinParamArity::Required,
246 default: None,
247 description: "wordEmbedding or wordEncoding object.",
248 },
249 BuiltinParamDescriptor {
250 name: "words",
251 ty: BuiltinParamType::Any,
252 arity: BuiltinParamArity::Required,
253 default: None,
254 description: "Words to test.",
255 },
256 BuiltinParamDescriptor {
257 name: "NameValue",
258 ty: BuiltinParamType::Any,
259 arity: BuiltinParamArity::Variadic,
260 default: None,
261 description: "Name-value options: IgnoreCase.",
262 },
263];
264
265const ERROR_ENCODING_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
266 code: "RM.WORDENCODING.INVALID_INPUT",
267 identifier: Some("RunMat:wordEncoding:InvalidInput"),
268 when: "Inputs do not match a supported wordEncoding form.",
269 message: "wordEncoding received invalid input",
270};
271
272const ERROR_WORD2IND_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
273 code: "RM.WORD2IND.INVALID_INPUT",
274 identifier: Some("RunMat:word2ind:InvalidInput"),
275 when: "Inputs do not match a supported word2ind form.",
276 message: "word2ind received invalid input",
277};
278
279const ERROR_IND2WORD_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
280 code: "RM.IND2WORD.INVALID_INPUT",
281 identifier: Some("RunMat:ind2word:InvalidInput"),
282 when: "Inputs do not match a supported ind2word form.",
283 message: "ind2word received invalid input",
284};
285
286const ERROR_IS_VOCABULARY_WORD_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
287 code: "RM.ISVOCABULARYWORD.INVALID_INPUT",
288 identifier: Some("RunMat:isVocabularyWord:InvalidInput"),
289 when: "Inputs do not match a supported isVocabularyWord form.",
290 message: "isVocabularyWord received invalid input",
291};
292
293const WORD_ENCODING_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_ENCODING_INVALID_INPUT];
294const WORD2IND_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_WORD2IND_INVALID_INPUT];
295const IND2WORD_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_IND2WORD_INVALID_INPUT];
296const IS_VOCABULARY_WORD_ERRORS: [BuiltinErrorDescriptor; 1] =
297 [ERROR_IS_VOCABULARY_WORD_INVALID_INPUT];
298
299fn any_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
300 Type::Unknown
301}
302
303pub const WORD_ENCODING_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
304 signatures: &[
305 BuiltinSignatureDescriptor {
306 label: "enc = wordEncoding(documents)",
307 inputs: &IN_DOCUMENTS_OR_WORDS,
308 outputs: &OUT_ENCODING,
309 },
310 BuiltinSignatureDescriptor {
311 label: "enc = wordEncoding(words)",
312 inputs: &IN_DOCUMENTS_OR_WORDS,
313 outputs: &OUT_ENCODING,
314 },
315 BuiltinSignatureDescriptor {
316 label: "enc = wordEncoding(documents, Name, Value)",
317 inputs: &IN_DOCUMENTS_OR_WORDS_REST,
318 outputs: &OUT_ENCODING,
319 },
320 ],
321 output_mode: BuiltinOutputMode::Fixed,
322 completion_policy: BuiltinCompletionPolicy::Public,
323 errors: &WORD_ENCODING_ERRORS,
324};
325
326pub const WORD2IND_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
327 signatures: &[
328 BuiltinSignatureDescriptor {
329 label: "M = word2ind(enc, words)",
330 inputs: &IN_WORDS,
331 outputs: &OUT_INDICES,
332 },
333 BuiltinSignatureDescriptor {
334 label: "M = word2ind(enc, words, 'IgnoreCase', true)",
335 inputs: &IN_WORDS_REST,
336 outputs: &OUT_INDICES,
337 },
338 ],
339 output_mode: BuiltinOutputMode::Fixed,
340 completion_policy: BuiltinCompletionPolicy::Public,
341 errors: &WORD2IND_ERRORS,
342};
343pub const WORD2IND_INTEGER_AUDIT: BuiltinIntegerAuditDescriptor = BuiltinIntegerAuditDescriptor {
344 kind: BuiltinIntegerAuditKind::NotApplicable,
345 canonical_builtin: None,
346 notes: "word2ind accepts a wordEncoding object, textual words, and a logical IgnoreCase option. It returns double indices or NaN; integer and resident numeric word/control inputs are invalid and reject before provider access.",
347};
348
349pub const IND2WORD_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
350 signatures: &[BuiltinSignatureDescriptor {
351 label: "words = ind2word(enc, M)",
352 inputs: &IN_INDICES,
353 outputs: &OUT_WORDS,
354 }],
355 output_mode: BuiltinOutputMode::Fixed,
356 completion_policy: BuiltinCompletionPolicy::Public,
357 errors: &IND2WORD_ERRORS,
358};
359
360pub const IS_VOCABULARY_WORD_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
361 signatures: &[
362 BuiltinSignatureDescriptor {
363 label: "tf = isVocabularyWord(emb, words)",
364 inputs: &IN_VOCABULARY_WORDS,
365 outputs: &OUT_LOGICAL,
366 },
367 BuiltinSignatureDescriptor {
368 label: "tf = isVocabularyWord(enc, words)",
369 inputs: &IN_VOCABULARY_WORDS,
370 outputs: &OUT_LOGICAL,
371 },
372 BuiltinSignatureDescriptor {
373 label: "tf = isVocabularyWord(___, 'IgnoreCase', true)",
374 inputs: &IN_VOCABULARY_WORDS_REST,
375 outputs: &OUT_LOGICAL,
376 },
377 ],
378 output_mode: BuiltinOutputMode::Fixed,
379 completion_policy: BuiltinCompletionPolicy::Public,
380 errors: &IS_VOCABULARY_WORD_ERRORS,
381};
382pub const IS_VOCABULARY_WORD_INTEGER_AUDIT: BuiltinIntegerAuditDescriptor =
383 BuiltinIntegerAuditDescriptor {
384 kind: BuiltinIntegerAuditKind::NotApplicable,
385 canonical_builtin: None,
386 notes: "isVocabularyWord accepts vocabulary objects, textual words, and logical options; integer and resident numeric values are invalid and reject before provider access.",
387 };
388
389#[runtime_builtin(
390 name = "wordEncoding",
391 category = "strings/text_analytics",
392 summary = "Create a word encoding object that maps words to indices and back.",
393 keywords = "wordEncoding,text analytics,words,indices,vocabulary",
394 accel = "sink",
395 type_resolver(any_type),
396 descriptor(crate::builtins::strings::text_analytics::encoding::WORD_ENCODING_DESCRIPTOR),
397 extensions(crate::builtins::strings::text_analytics::encoding::WORD_ENCODING_EXTENSIONS),
398 integer_capabilities(
399 crate::builtins::strings::text_analytics::encoding::WORD_ENCODING_INTEGER_CAPABILITIES
400 ),
401 builtin_path = "crate::builtins::strings::text_analytics::encoding"
402)]
403async fn word_encoding_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
404 for pair in args.windows(2) {
405 if scalar_text(&pair[0], "wordEncoding")
406 .is_ok_and(|name| name.eq_ignore_ascii_case("MaxNumWords"))
407 && is_typed_integer_value(&pair[1])
408 {
409 crate::compatibility::ensure_builtin_extension_enabled(
410 &WORD_ENCODING_INTEGER_MAX_WORDS_EXTENSION,
411 "wordEncoding",
412 )?;
413 }
414 }
415 let gathered = gather_args(args, "wordEncoding").await?;
416 let (source, options) = parse_word_encoding_args(gathered)?;
417 word_encoding_object(build_word_encoding(source, options)?)
418}
419
420#[runtime_builtin(
421 name = "word2ind",
422 category = "strings/text_analytics",
423 summary = "Map words to indices in a wordEncoding object.",
424 keywords = "word2ind,wordEncoding,text analytics,indices,vocabulary",
425 accel = "sink",
426 type_resolver(any_type),
427 descriptor(crate::builtins::strings::text_analytics::encoding::WORD2IND_DESCRIPTOR),
428 integer_audit(crate::builtins::strings::text_analytics::encoding::WORD2IND_INTEGER_AUDIT),
429 builtin_path = "crate::builtins::strings::text_analytics::encoding"
430)]
431async fn word2ind_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
432 if args.iter().skip(1).any(|value| {
433 crate::builtins::common::validation::value_contains_native_integer_class(value)
434 || value_contains_resident(value)
435 }) {
436 return Err(encoding_error(
437 "word2ind",
438 "word2ind: words and option names must be host text and IgnoreCase must be logical",
439 ));
440 }
441 let gathered = gather_args(args, "word2ind").await?;
442 let (object, words, options) = parse_word2ind_args(gathered)?;
443 let encoding = word_encoding_from_object(&object, "word2ind")?;
444 let lookup = build_word_lookup(&encoding.vocabulary, options.ignore_case);
445 let indices = words
446 .words
447 .into_iter()
448 .map(|word| {
449 let key = if options.ignore_case {
450 word.to_lowercase()
451 } else {
452 word
453 };
454 lookup
455 .get(&key)
456 .map(|idx| (*idx + 1) as f64)
457 .unwrap_or(f64::NAN)
458 })
459 .collect::<Vec<_>>();
460 Tensor::new(indices, words.shape)
461 .map(Value::Tensor)
462 .map_err(|err| encoding_error("word2ind", err))
463}
464
465#[runtime_builtin(
466 name = "ind2word",
467 category = "strings/text_analytics",
468 summary = "Map wordEncoding indices back to words.",
469 keywords = "ind2word,wordEncoding,text analytics,indices,vocabulary",
470 accel = "sink",
471 type_resolver(any_type),
472 descriptor(crate::builtins::strings::text_analytics::encoding::IND2WORD_DESCRIPTOR),
473 extensions(crate::builtins::strings::text_analytics::encoding::IND2WORD_EXTENSIONS),
474 integer_capabilities(
475 crate::builtins::strings::text_analytics::encoding::IND2WORD_INTEGER_CAPABILITIES
476 ),
477 builtin_path = "crate::builtins::strings::text_analytics::encoding"
478)]
479async fn ind2word_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
480 ensure_ind2word_extensions(&args)?;
481 let gathered = gather_args(args, "ind2word").await?;
482 let (object, indices) = parse_ind2word_args(gathered)?;
483 let encoding = word_encoding_from_object(&object, "ind2word")?;
484 let words = indices
485 .values
486 .into_iter()
487 .map(|idx| {
488 let word_idx = positive_index(idx, encoding.vocabulary.len(), "ind2word")?;
489 Ok(encoding.vocabulary[word_idx].clone())
490 })
491 .collect::<BuiltinResult<Vec<_>>>()?;
492 StringArray::new(words, indices.shape)
493 .map(Value::StringArray)
494 .map_err(|err| encoding_error("ind2word", err))
495}
496
497fn ensure_ind2word_extensions(args: &[Value]) -> BuiltinResult<()> {
498 if args.len() != 2 {
499 return Ok(());
500 }
501 let indices = &args[1];
502 if is_typed_integer_value(indices) {
503 crate::compatibility::ensure_builtin_extension_enabled(
504 &IND2WORD_TYPED_INTEGER_EXTENSION,
505 "ind2word",
506 )?;
507 }
508 if crate::dispatcher::value_contains_gpu(indices) {
509 crate::compatibility::ensure_builtin_extension_enabled(
510 &IND2WORD_RESIDENT_EXTENSION,
511 "ind2word",
512 )?;
513 }
514 if value_shape(indices).is_some_and(|shape| !is_vector_shape(shape)) {
515 crate::compatibility::ensure_builtin_extension_enabled(
516 &IND2WORD_NONVECTOR_EXTENSION,
517 "ind2word",
518 )?;
519 }
520 Ok(())
521}
522
523fn is_typed_integer_value(value: &Value) -> bool {
524 matches!(value, Value::Int(_))
525 || matches!(value, Value::Tensor(tensor) if tensor.integer_storage().is_some())
526 || matches!(value, Value::GpuTensor(handle) if runmat_accelerate_api::handle_integer_type(handle).is_some())
527}
528
529fn value_shape(value: &Value) -> Option<&[usize]> {
530 match value {
531 Value::Tensor(tensor) => Some(&tensor.shape),
532 Value::GpuTensor(handle) => Some(&handle.shape),
533 Value::Num(_) | Value::Int(_) => Some(&[1, 1]),
534 _ => None,
535 }
536}
537
538fn is_vector_shape(shape: &[usize]) -> bool {
539 shape.len() <= 2 && shape.iter().filter(|extent| **extent > 1).count() <= 1
540}
541
542#[runtime_builtin(
543 name = "isVocabularyWord",
544 category = "strings/text_analytics",
545 summary = "Test whether words are in a wordEmbedding or wordEncoding vocabulary.",
546 keywords = "isVocabularyWord,wordEmbedding,wordEncoding,text analytics,vocabulary",
547 accel = "sink",
548 type_resolver(any_type),
549 descriptor(crate::builtins::strings::text_analytics::encoding::IS_VOCABULARY_WORD_DESCRIPTOR),
550 integer_audit(
551 crate::builtins::strings::text_analytics::encoding::IS_VOCABULARY_WORD_INTEGER_AUDIT
552 ),
553 builtin_path = "crate::builtins::strings::text_analytics::encoding"
554)]
555async fn is_vocabulary_word_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
556 if args.iter().any(value_contains_resident) {
557 return Err(encoding_error(
558 "isVocabularyWord",
559 "isVocabularyWord: provider-resident numeric inputs are not vocabulary objects, words, or controls",
560 ));
561 }
562 let gathered = gather_args(args, "isVocabularyWord").await?;
563 let (object, words, options) = parse_is_vocabulary_word_args(gathered)?;
564 let vocabulary = if object.is_class(WORD_ENCODING_CLASS) {
565 word_encoding_from_object(&object, "isVocabularyWord")?.vocabulary
566 } else if object.is_class(WORD_EMBEDDING_CLASS) {
567 word_embedding_vocabulary_from_object(&object, "isVocabularyWord")?
568 } else {
569 return Err(encoding_error(
570 "isVocabularyWord",
571 format!(
572 "isVocabularyWord: expected wordEmbedding or wordEncoding object, got {}",
573 object.class_name
574 ),
575 ));
576 };
577 let lookup = build_word_lookup(&vocabulary, options.ignore_case);
578 let flags = words
579 .words
580 .into_iter()
581 .map(|word| {
582 let key = if options.ignore_case {
583 word.to_lowercase()
584 } else {
585 word
586 };
587 u8::from(lookup.contains_key(&key))
588 })
589 .collect::<Vec<_>>();
590 LogicalArray::new(flags, words.shape)
591 .map(Value::LogicalArray)
592 .map_err(|err| encoding_error("isVocabularyWord", err))
593}
594
595fn value_contains_resident(value: &Value) -> bool {
596 match value {
597 Value::GpuTensor(_) => true,
598 Value::Cell(value) => value.data.iter().any(value_contains_resident),
599 Value::Struct(value) => value.fields.values().any(value_contains_resident),
600 Value::Object(value) => value.properties.values().any(value_contains_resident),
601 Value::Closure(value) => value.captures.iter().any(value_contains_resident),
602 Value::OutputList(values) => values.iter().any(value_contains_resident),
603 _ => false,
604 }
605}
606
607async fn gather_args(args: Vec<Value>, fn_name: &str) -> BuiltinResult<Vec<Value>> {
608 let mut out = Vec::with_capacity(args.len());
609 for arg in args {
610 out.push(gather_if_needed_async(&arg).await.map_err(|err| {
611 encoding_error(fn_name, format!("{fn_name}: failed to gather input: {err}"))
612 })?);
613 }
614 Ok(out)
615}
616
617#[derive(Clone, Debug)]
618pub(in crate::builtins::strings::text_analytics) struct WordEncodingModel {
619 pub vocabulary: Vec<String>,
620}
621
622pub(in crate::builtins::strings::text_analytics) fn word_encoding_from_object(
623 object: &ObjectInstance,
624 fn_name: &str,
625) -> BuiltinResult<WordEncodingModel> {
626 if !object.is_class(WORD_ENCODING_CLASS) {
627 return Err(encoding_error(
628 fn_name,
629 format!(
630 "{fn_name}: expected wordEncoding object, got {}",
631 object.class_name
632 ),
633 ));
634 }
635 let vocabulary = match object.properties.get("Vocabulary") {
636 Some(Value::StringArray(array)) => array.data.clone(),
637 other => {
638 return Err(encoding_error(
639 fn_name,
640 format!(
641 "{fn_name}: wordEncoding object has invalid Vocabulary property: {other:?}"
642 ),
643 ));
644 }
645 };
646 match object.properties.get("NumWords") {
647 Some(Value::Num(value)) if *value == vocabulary.len() as f64 => {}
648 other => {
649 return Err(encoding_error(
650 fn_name,
651 format!("{fn_name}: wordEncoding object has invalid NumWords property: {other:?}"),
652 ));
653 }
654 }
655 Ok(WordEncodingModel { vocabulary })
656}
657
658fn word_encoding_object(model: WordEncodingModel) -> BuiltinResult<Value> {
659 ensure_word_encoding_class_registered();
660 let mut object = ObjectInstance::new(WORD_ENCODING_CLASS.to_string());
661 object.properties.insert(
662 "NumWords".to_string(),
663 Value::Num(model.vocabulary.len() as f64),
664 );
665 object.properties.insert(
666 "Vocabulary".to_string(),
667 Value::StringArray(
668 StringArray::new(model.vocabulary.clone(), vec![1, model.vocabulary.len()])
669 .map_err(|err| encoding_error("wordEncoding", err))?,
670 ),
671 );
672 Ok(Value::Object(object))
673}
674
675fn ensure_word_encoding_class_registered() {
676 WORD_ENCODING_CLASS_REGISTERED.ensure(|| {
677 let mut properties = HashMap::new();
678 for name in ["NumWords", "Vocabulary"] {
679 properties.insert(name.to_string(), property_def(name));
680 }
681 crate::class_registry::register_class(crate::class_registry::RuntimeClass {
682 name: WORD_ENCODING_CLASS.to_string(),
683 parent: None,
684 properties,
685 methods: HashMap::new(),
686 });
687 });
688}
689
690fn property_def(name: &str) -> crate::class_registry::RuntimeProperty {
691 crate::class_registry::RuntimeProperty {
692 name: name.to_string(),
693 is_static: false,
694 is_constant: false,
695 is_dependent: false,
696 get_access: MemberAccess::Public,
697 set_access: MemberAccess::Public,
698 default_value: None,
699 }
700}
701
702enum EncodingSource {
703 Documents(Vec<Vec<String>>),
704 Words(Vec<String>),
705}
706
707#[derive(Clone, Copy, Debug, PartialEq, Eq)]
708enum EncodingOrder {
709 FirstSeen,
710 Frequency,
711}
712
713#[derive(Clone, Copy, Debug)]
714struct WordEncodingOptions {
715 order: EncodingOrder,
716 max_num_words: Option<usize>,
717}
718
719impl Default for WordEncodingOptions {
720 fn default() -> Self {
721 Self {
722 order: EncodingOrder::FirstSeen,
723 max_num_words: None,
724 }
725 }
726}
727
728fn parse_word_encoding_args(
729 args: Vec<Value>,
730) -> BuiltinResult<(EncodingSource, WordEncodingOptions)> {
731 if args.is_empty() {
732 return Err(encoding_error(
733 "wordEncoding",
734 "wordEncoding: expected tokenizedDocument object or word vector",
735 ));
736 }
737 if !(args.len() - 1).is_multiple_of(2) {
738 return Err(encoding_error(
739 "wordEncoding",
740 "wordEncoding: name-value options must be paired",
741 ));
742 }
743 let source = match &args[0] {
744 Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => {
745 EncodingSource::Documents(documents_from_object(object, "wordEncoding")?)
746 }
747 Value::Object(object) => {
748 return Err(encoding_error(
749 "wordEncoding",
750 format!(
751 "wordEncoding: expected tokenizedDocument object or word vector, got {}",
752 object.class_name
753 ),
754 ));
755 }
756 value => EncodingSource::Words(word_input_from_value(value, "wordEncoding")?.words),
757 };
758 if matches!(source, EncodingSource::Words(_)) && args.len() > 1 {
759 return Err(encoding_error(
760 "wordEncoding",
761 "wordEncoding: Order and MaxNumWords options are only supported for tokenizedDocument input",
762 ));
763 }
764 let mut options = WordEncodingOptions::default();
765 let mut idx = 1usize;
766 while idx < args.len() {
767 let name = scalar_text(&args[idx], "wordEncoding")
768 .map_err(|err| encoding_error("wordEncoding", err.to_string()))?
769 .to_ascii_lowercase();
770 match name.as_str() {
771 "order" => {
772 let value = scalar_text(&args[idx + 1], "wordEncoding")
773 .map_err(|err| encoding_error("wordEncoding", err.to_string()))?
774 .to_ascii_lowercase();
775 options.order = match value.as_str() {
776 "first-seen" => EncodingOrder::FirstSeen,
777 "frequency" => EncodingOrder::Frequency,
778 other => {
779 return Err(encoding_error(
780 "wordEncoding",
781 format!(
782 "wordEncoding: Order must be 'first-seen' or 'frequency', got '{other}'"
783 ),
784 ));
785 }
786 };
787 }
788 "maxnumwords" => {
789 options.max_num_words = parse_max_num_words(&args[idx + 1])?;
790 }
791 other => {
792 return Err(encoding_error(
793 "wordEncoding",
794 format!("wordEncoding: unsupported option '{other}'"),
795 ));
796 }
797 }
798 idx += 2;
799 }
800 Ok((source, options))
801}
802
803fn build_word_encoding(
804 source: EncodingSource,
805 options: WordEncodingOptions,
806) -> BuiltinResult<WordEncodingModel> {
807 let words = match source {
808 EncodingSource::Documents(documents) => documents.into_iter().flatten().collect::<Vec<_>>(),
809 EncodingSource::Words(words) => words,
810 };
811 let mut counts = HashMap::<String, (usize, usize)>::new();
812 for (pos, word) in words.into_iter().enumerate() {
813 let entry = counts.entry(word).or_insert((0, pos));
814 entry.0 += 1;
815 }
816 let mut ranked = counts
817 .into_iter()
818 .map(|(word, (count, first_pos))| (word, count, first_pos))
819 .collect::<Vec<_>>();
820 match options.order {
821 EncodingOrder::FirstSeen => ranked.sort_by(|left, right| left.2.cmp(&right.2)),
822 EncodingOrder::Frequency => {
823 ranked.sort_by(|left, right| right.1.cmp(&left.1).then(left.2.cmp(&right.2)))
824 }
825 }
826 if let Some(max) = options.max_num_words {
827 ranked.truncate(max);
828 }
829 Ok(WordEncodingModel {
830 vocabulary: ranked.into_iter().map(|(word, _, _)| word).collect(),
831 })
832}
833
834#[derive(Clone, Copy, Debug, Default)]
835struct LookupOptions {
836 ignore_case: bool,
837}
838
839fn parse_word2ind_args(
840 args: Vec<Value>,
841) -> BuiltinResult<(ObjectInstance, WordInput, LookupOptions)> {
842 if args.len() < 2 {
843 return Err(encoding_error(
844 "word2ind",
845 "word2ind: expected word2ind(enc, words)",
846 ));
847 }
848 let object = object_arg(&args[0], "word2ind", "wordEncoding")?;
849 let words = word_input_from_value(&args[1], "word2ind")?;
850 let options = parse_lookup_options(&args[2..], "word2ind")?;
851 Ok((object, words, options))
852}
853
854fn parse_is_vocabulary_word_args(
855 args: Vec<Value>,
856) -> BuiltinResult<(ObjectInstance, WordInput, LookupOptions)> {
857 if args.len() < 2 {
858 return Err(encoding_error(
859 "isVocabularyWord",
860 "isVocabularyWord: expected isVocabularyWord(embOrEnc, words)",
861 ));
862 }
863 let object = object_arg(
864 &args[0],
865 "isVocabularyWord",
866 "wordEmbedding or wordEncoding",
867 )?;
868 let words = word_input_from_value(&args[1], "isVocabularyWord")?;
869 let options = parse_lookup_options(&args[2..], "isVocabularyWord")?;
870 Ok((object, words, options))
871}
872
873fn parse_ind2word_args(args: Vec<Value>) -> BuiltinResult<(ObjectInstance, NumericInput)> {
874 if args.len() != 2 {
875 return Err(encoding_error(
876 "ind2word",
877 "ind2word: expected ind2word(enc, M)",
878 ));
879 }
880 let object = object_arg(&args[0], "ind2word", "wordEncoding")?;
881 let indices = numeric_input_from_value(&args[1], "ind2word")?;
882 Ok((object, indices))
883}
884
885fn object_arg(value: &Value, fn_name: &str, expected: &str) -> BuiltinResult<ObjectInstance> {
886 match value {
887 Value::Object(object) => Ok(object.clone()),
888 other => Err(encoding_error(
889 fn_name,
890 format!("{fn_name}: expected {expected} object, got {other:?}"),
891 )),
892 }
893}
894
895fn parse_lookup_options(args: &[Value], fn_name: &str) -> BuiltinResult<LookupOptions> {
896 if !args.len().is_multiple_of(2) {
897 return Err(encoding_error(
898 fn_name,
899 format!("{fn_name}: name-value options must be paired"),
900 ));
901 }
902 let mut options = LookupOptions::default();
903 let mut idx = 0usize;
904 while idx < args.len() {
905 let name = scalar_text(&args[idx], fn_name)
906 .map_err(|err| encoding_error(fn_name, err.to_string()))?
907 .to_ascii_lowercase();
908 match name.as_str() {
909 "ignorecase" => options.ignore_case = parse_bool_scalar(&args[idx + 1], fn_name)?,
910 other => {
911 return Err(encoding_error(
912 fn_name,
913 format!("{fn_name}: unsupported option '{other}'"),
914 ));
915 }
916 }
917 idx += 2;
918 }
919 Ok(options)
920}
921
922struct WordInput {
923 words: Vec<String>,
924 shape: Vec<usize>,
925}
926
927fn word_input_from_value(value: &Value, fn_name: &str) -> BuiltinResult<WordInput> {
928 match value {
929 Value::String(text) => Ok(WordInput {
930 words: vec![text.clone()],
931 shape: vec![1, 1],
932 }),
933 Value::StringArray(array) => Ok(WordInput {
934 words: array.data.clone(),
935 shape: array.shape.clone(),
936 }),
937 Value::CharArray(array) if array.rows <= 1 => Ok(WordInput {
938 words: vec![char_row_to_string(array)],
939 shape: vec![1, 1],
940 }),
941 Value::CharArray(array) => {
942 let mut words = Vec::with_capacity(array.rows);
943 for row in 0..array.rows {
944 let mut text = String::with_capacity(array.cols);
945 for col in 0..array.cols {
946 text.push(array.data[row + col * array.rows]);
947 }
948 words.push(text.trim_end().to_string());
949 }
950 Ok(WordInput {
951 words,
952 shape: vec![array.rows, 1],
953 })
954 }
955 Value::Cell(cell) => {
956 let words = cell
957 .data
958 .iter()
959 .map(|item| {
960 scalar_text(item, fn_name)
961 .map_err(|err| encoding_error(fn_name, err.to_string()))
962 })
963 .collect::<BuiltinResult<Vec<_>>>()?;
964 Ok(WordInput {
965 words,
966 shape: cell.shape.clone(),
967 })
968 }
969 other => Err(encoding_error(
970 fn_name,
971 format!("{fn_name}: expected string, character vector, or cell array of words, got {other:?}"),
972 )),
973 }
974}
975
976struct NumericInput {
977 values: Vec<NumericIndex>,
978 shape: Vec<usize>,
979}
980
981#[derive(Clone, Debug, PartialEq)]
982enum NumericIndex {
983 Float(f64),
984 Integer(IntValue),
985}
986
987fn numeric_input_from_value(value: &Value, fn_name: &str) -> BuiltinResult<NumericInput> {
988 match value {
989 Value::Num(value) => Ok(NumericInput {
990 values: vec![NumericIndex::Float(*value)],
991 shape: vec![1, 1],
992 }),
993 Value::Int(value) => Ok(NumericInput {
994 values: vec![NumericIndex::Integer(value.clone())],
995 shape: vec![1, 1],
996 }),
997 Value::Tensor(tensor) => {
998 let values = if let Some(storage) = tensor.integer_storage() {
999 (0..storage.len())
1000 .map(|index| {
1001 storage
1002 .value_at(index)
1003 .map(NumericIndex::Integer)
1004 .expect("integer index is within storage bounds")
1005 })
1006 .collect()
1007 } else {
1008 tensor_utils::tensor_values_f64(tensor)
1009 .into_iter()
1010 .map(NumericIndex::Float)
1011 .collect()
1012 };
1013 Ok(NumericInput {
1014 values,
1015 shape: tensor.shape.clone(),
1016 })
1017 }
1018 other => Err(encoding_error(
1019 fn_name,
1020 format!("{fn_name}: expected numeric positive integer indices, got {other:?}"),
1021 )),
1022 }
1023}
1024
1025fn positive_index(value: NumericIndex, len: usize, fn_name: &str) -> BuiltinResult<usize> {
1026 let idx = match value {
1027 NumericIndex::Float(value) => {
1028 if !value.is_finite() || value < 1.0 || value.fract() != 0.0 {
1029 return Err(encoding_error(
1030 fn_name,
1031 format!("{fn_name}: indices must be positive integers, got {value}"),
1032 ));
1033 }
1034 if value > usize::MAX as f64 {
1035 return Err(encoding_error(
1036 fn_name,
1037 format!("{fn_name}: index exceeds platform limits"),
1038 ));
1039 }
1040 value as usize
1041 }
1042 NumericIndex::Integer(value) => value.try_to_usize().ok_or_else(|| {
1043 encoding_error(
1044 fn_name,
1045 format!("{fn_name}: indices must be positive integers"),
1046 )
1047 })?,
1048 };
1049 if idx == 0 {
1050 return Err(encoding_error(
1051 fn_name,
1052 format!("{fn_name}: indices must be positive integers"),
1053 ));
1054 }
1055 if idx > len {
1056 return Err(encoding_error(
1057 fn_name,
1058 format!("{fn_name}: index {idx} exceeds vocabulary size {len}"),
1059 ));
1060 }
1061 Ok(idx - 1)
1062}
1063
1064fn parse_max_num_words(value: &Value) -> BuiltinResult<Option<usize>> {
1065 if let Value::Int(value) = value {
1066 return value
1067 .try_to_usize()
1068 .filter(|value| *value >= 1)
1069 .map(Some)
1070 .ok_or_else(|| {
1071 encoding_error(
1072 "wordEncoding",
1073 "wordEncoding: MaxNumWords must be a positive integer or Inf",
1074 )
1075 });
1076 }
1077 if let Value::Tensor(tensor) = value {
1078 if tensor_utils::is_scalar_tensor(tensor) {
1079 if let Some(value) = tensor
1080 .integer_storage()
1081 .and_then(|storage| storage.value_at(0))
1082 {
1083 return value
1084 .try_to_usize()
1085 .filter(|value| *value >= 1)
1086 .map(Some)
1087 .ok_or_else(|| {
1088 encoding_error(
1089 "wordEncoding",
1090 "wordEncoding: MaxNumWords must be a positive integer or Inf",
1091 )
1092 });
1093 }
1094 }
1095 }
1096 let n = numeric_scalar(value, "wordEncoding", "MaxNumWords")?;
1097 if n.is_infinite() && n.is_sign_positive() {
1098 return Ok(None);
1099 }
1100 if !n.is_finite() || n < 1.0 || n.fract() != 0.0 {
1101 return Err(encoding_error(
1102 "wordEncoding",
1103 format!("wordEncoding: MaxNumWords must be a positive integer or Inf, got {n}"),
1104 ));
1105 }
1106 Ok(Some(n as usize))
1107}
1108
1109fn numeric_scalar(value: &Value, fn_name: &str, option: &str) -> BuiltinResult<f64> {
1110 match value {
1111 Value::Num(value) => Ok(*value),
1112 Value::Int(value) => Ok(int_value_to_f64(value)),
1113 Value::Tensor(tensor) if tensor_utils::is_scalar_tensor(tensor) => {
1114 Ok(tensor_utils::tensor_value_f64(tensor, 0))
1115 }
1116 other => Err(encoding_error(
1117 fn_name,
1118 format!("{fn_name}: {option} must be a numeric scalar, got {other:?}"),
1119 )),
1120 }
1121}
1122
1123fn parse_bool_scalar(value: &Value, fn_name: &str) -> BuiltinResult<bool> {
1124 match value {
1125 Value::Bool(value) => Ok(*value),
1126 Value::Num(value) if *value == 0.0 || *value == 1.0 => Ok(*value != 0.0),
1127 Value::Tensor(tensor) if tensor_utils::is_scalar_tensor(tensor) => {
1128 if let Some(value) = tensor
1129 .integer_storage()
1130 .and_then(|storage| storage.value_at(0))
1131 {
1132 return match value.try_to_u64() {
1133 Some(0) => Ok(false),
1134 Some(1) => Ok(true),
1135 _ => Err(encoding_error(
1136 fn_name,
1137 format!(
1138 "{fn_name}: logical scalar option must be true or false, got {value:?}"
1139 ),
1140 )),
1141 };
1142 }
1143 match tensor_utils::tensor_value_f64(tensor, 0) {
1144 0.0 => Ok(false),
1145 1.0 => Ok(true),
1146 other => Err(encoding_error(
1147 fn_name,
1148 format!("{fn_name}: logical scalar option must be true or false, got {other}"),
1149 )),
1150 }
1151 }
1152 Value::LogicalArray(array) if array.data.len() == 1 => Ok(array.data[0] != 0),
1153 other => Err(encoding_error(
1154 fn_name,
1155 format!("{fn_name}: logical scalar option must be true or false, got {other:?}"),
1156 )),
1157 }
1158}
1159
1160fn int_value_to_f64(value: &runmat_value::IntValue) -> f64 {
1161 match value {
1162 runmat_value::IntValue::I8(value) => *value as f64,
1163 runmat_value::IntValue::I16(value) => *value as f64,
1164 runmat_value::IntValue::I32(value) => *value as f64,
1165 runmat_value::IntValue::I64(value) => *value as f64,
1166 runmat_value::IntValue::U8(value) => *value as f64,
1167 runmat_value::IntValue::U16(value) => *value as f64,
1168 runmat_value::IntValue::U32(value) => *value as f64,
1169 runmat_value::IntValue::U64(value) => *value as f64,
1170 }
1171}
1172
1173fn char_row_to_string(array: &CharArray) -> String {
1174 array.data.iter().collect()
1175}
1176
1177fn encoding_error(fn_name: &str, message: impl Into<String>) -> crate::RuntimeError {
1178 let descriptor = match fn_name {
1179 "word2ind" => ERROR_WORD2IND_INVALID_INPUT,
1180 "ind2word" => ERROR_IND2WORD_INVALID_INPUT,
1181 "isVocabularyWord" => ERROR_IS_VOCABULARY_WORD_INVALID_INPUT,
1182 _ => ERROR_ENCODING_INVALID_INPUT,
1183 };
1184 let builder = build_runtime_error(message.into()).with_builtin(fn_name);
1185 match descriptor.identifier {
1186 Some(identifier) => builder.with_identifier(identifier).build(),
1187 None => builder.build(),
1188 }
1189}
1190
1191#[cfg(test)]
1192mod tests {
1193 use super::*;
1194 use runmat_value::{CellArray, IntegerStorage};
1195
1196 fn poisoned_integer_scalar(storage: IntegerStorage) -> Value {
1197 let tensor = Tensor::new_integer(storage, vec![1, 1]).expect("integer tensor");
1198 Value::Tensor(tensor)
1199 }
1200
1201 fn poisoned_integer_vector(storage: IntegerStorage, cols: usize) -> Value {
1202 let tensor = Tensor::new_integer(storage, vec![1, cols]).expect("integer tensor");
1203 Value::Tensor(tensor)
1204 }
1205
1206 fn tokenized_document_object(rows: Vec<Vec<&str>>) -> ObjectInstance {
1207 let values = rows
1208 .into_iter()
1209 .map(|row| {
1210 let len = row.len();
1211 Value::StringArray(
1212 StringArray::new(
1213 row.into_iter().map(str::to_string).collect::<Vec<_>>(),
1214 vec![1, len],
1215 )
1216 .unwrap(),
1217 )
1218 })
1219 .collect::<Vec<_>>();
1220 let rows = values.len();
1221 let documents = CellArray::new(values, rows, 1).unwrap();
1222 let mut object = ObjectInstance::new(TOKENIZED_DOCUMENT_CLASS.to_string());
1223 object
1224 .properties
1225 .insert("Documents".to_string(), Value::Cell(documents));
1226 object
1227 }
1228
1229 #[test]
1230 fn scalar_option_parsers_read_typed_integer_storage_exactly() {
1231 assert_eq!(
1232 numeric_scalar(
1233 &poisoned_integer_scalar(IntegerStorage::U16(vec![12])),
1234 "wordEncoding",
1235 "MaxNumWords"
1236 )
1237 .expect("numeric"),
1238 12.0
1239 );
1240 assert!(parse_bool_scalar(
1241 &poisoned_integer_scalar(IntegerStorage::U8(vec![1])),
1242 "wordEncoding"
1243 )
1244 .expect("bool"));
1245 assert!(!parse_bool_scalar(
1246 &poisoned_integer_scalar(IntegerStorage::I16(vec![0])),
1247 "wordEncoding"
1248 )
1249 .expect("bool"));
1250 }
1251
1252 #[test]
1253 fn numeric_input_reads_typed_integer_storage_exactly() {
1254 let input = numeric_input_from_value(
1255 &poisoned_integer_vector(IntegerStorage::I16(vec![2, 3]), 2),
1256 "ind2word",
1257 )
1258 .expect("numeric");
1259
1260 assert_eq!(
1261 input.values,
1262 vec![
1263 NumericIndex::Integer(IntValue::I16(2)),
1264 NumericIndex::Integer(IntValue::I16(3))
1265 ]
1266 );
1267 assert_eq!(input.shape, vec![1, 2]);
1268 }
1269
1270 #[tokio::test]
1271 async fn word_encoding_builds_first_seen_and_frequency_vocabularies() {
1272 let documents = Value::Object(tokenized_document_object(vec![
1273 vec!["beta", "alpha", "beta"],
1274 vec!["gamma", "alpha", "beta"],
1275 ]));
1276 let first_seen = word_encoding_builtin(vec![documents.clone()])
1277 .await
1278 .unwrap();
1279 let Value::Object(first_seen) = first_seen else {
1280 panic!("expected object");
1281 };
1282 let model = word_encoding_from_object(&first_seen, "test").unwrap();
1283 assert_eq!(model.vocabulary, vec!["beta", "alpha", "gamma"]);
1284
1285 let frequency = word_encoding_builtin(vec![
1286 documents,
1287 Value::String("Order".into()),
1288 Value::String("frequency".into()),
1289 Value::String("MaxNumWords".into()),
1290 Value::Num(2.0),
1291 ])
1292 .await
1293 .unwrap();
1294 let Value::Object(frequency) = frequency else {
1295 panic!("expected object");
1296 };
1297 let model = word_encoding_from_object(&frequency, "test").unwrap();
1298 assert_eq!(model.vocabulary, vec!["beta", "alpha"]);
1299 }
1300
1301 #[tokio::test]
1302 async fn word_encoding_accepts_word_arrays_and_validates_options() {
1303 let words = Value::StringArray(
1304 StringArray::new(vec!["red".into(), "blue".into(), "red".into()], vec![1, 3]).unwrap(),
1305 );
1306 let enc = word_encoding_builtin(vec![words]).await.unwrap();
1307 let Value::Object(enc) = enc else {
1308 panic!("expected object");
1309 };
1310 assert_eq!(enc.properties.get("NumWords"), Some(&Value::Num(2.0)));
1311
1312 let err = word_encoding_builtin(vec![
1313 Value::String("x".into()),
1314 Value::String("Order".into()),
1315 Value::String("frequency".into()),
1316 ])
1317 .await
1318 .unwrap_err();
1319 assert!(
1320 err.to_string()
1321 .contains("only supported for tokenizedDocument input"),
1322 "{err}"
1323 );
1324 }
1325
1326 #[tokio::test]
1327 async fn word_encoding_typed_maximum_is_a_gated_exact_control() {
1328 let documents = Value::Object(tokenized_document_object(vec![vec!["alpha", "beta"]]));
1329 let args = vec![
1330 documents,
1331 Value::String("MaxNumWords".into()),
1332 Value::Int(IntValue::U64(1)),
1333 ];
1334 let _strict = crate::compatibility::push_runmat_extensions_enabled(false);
1335 let error = word_encoding_builtin(args.clone())
1336 .await
1337 .expect_err("typed MaxNumWords is gated in strict mode");
1338 assert_eq!(
1339 error.identifier(),
1340 WORD_ENCODING_INTEGER_MAX_WORDS_EXTENSION.error_identifier
1341 );
1342 drop(_strict);
1343
1344 let _runmat = crate::compatibility::push_runmat_extensions_enabled(true);
1345 let value = word_encoding_builtin(args)
1346 .await
1347 .expect("RunMat mode accepts typed MaxNumWords");
1348 let Value::Object(object) = value else {
1349 panic!("expected wordEncoding object");
1350 };
1351 assert_eq!(
1352 word_encoding_from_object(&object, "test")
1353 .unwrap()
1354 .vocabulary,
1355 vec!["alpha"]
1356 );
1357 }
1358
1359 #[tokio::test]
1360 async fn word2ind_preserves_shape_and_supports_ignore_case() {
1361 let enc = word_encoding_builtin(vec![Value::StringArray(
1362 StringArray::new(vec!["Alpha".into(), "beta".into()], vec![1, 2]).unwrap(),
1363 )])
1364 .await
1365 .unwrap();
1366 let words = Value::StringArray(
1367 StringArray::new(
1368 vec![
1369 "beta".into(),
1370 "missing".into(),
1371 "alpha".into(),
1372 "Alpha".into(),
1373 ],
1374 vec![2, 2],
1375 )
1376 .unwrap(),
1377 );
1378 let out = word2ind_builtin(vec![
1379 enc,
1380 words,
1381 Value::String("IgnoreCase".into()),
1382 Value::Bool(true),
1383 ])
1384 .await
1385 .unwrap();
1386 let Value::Tensor(indices) = out else {
1387 panic!("expected tensor");
1388 };
1389 assert_eq!(indices.shape, vec![2, 2]);
1390 assert_eq!(indices.materialize_f64()[0], 2.0);
1391 assert!(indices.materialize_f64()[1].is_nan());
1392 assert_eq!(indices.materialize_f64()[2], 1.0);
1393 assert_eq!(indices.materialize_f64()[3], 1.0);
1394 }
1395
1396 #[tokio::test]
1397 async fn word2ind_rejects_integer_words_before_object_validation() {
1398 let error = word2ind_builtin(vec![
1399 Value::String("not an object".into()),
1400 Value::Int(IntValue::U8(1)),
1401 ])
1402 .await
1403 .expect_err("integer words are outside the text-only surface");
1404 assert!(error.message().contains("must be host text"));
1405 }
1406
1407 #[tokio::test]
1408 async fn ind2word_preserves_numeric_shape_and_rejects_bad_indices() {
1409 let enc = word_encoding_builtin(vec![Value::StringArray(
1410 StringArray::new(
1411 vec!["red".into(), "blue".into(), "green".into()],
1412 vec![1, 3],
1413 )
1414 .unwrap(),
1415 )])
1416 .await
1417 .unwrap();
1418 let out = ind2word_builtin(vec![
1419 enc.clone(),
1420 Value::Tensor(Tensor::new(vec![1.0, 3.0], vec![1, 2]).unwrap()),
1421 ])
1422 .await
1423 .unwrap();
1424 let Value::StringArray(words) = out else {
1425 panic!("expected string array");
1426 };
1427 assert_eq!(words.shape, vec![1, 2]);
1428 assert_eq!(words.data, vec!["red", "green"]);
1429
1430 let err = ind2word_builtin(vec![enc, Value::Num(4.0)])
1431 .await
1432 .unwrap_err();
1433 assert!(err.to_string().contains("exceeds vocabulary"), "{err}");
1434 }
1435
1436 #[test]
1437 fn ind2word_extensions_gate_before_gather_and_integer_indices_stay_exact() {
1438 let enc = futures::executor::block_on(word_encoding_builtin(vec![Value::StringArray(
1439 StringArray::new(vec!["red".into(), "blue".into()], vec![1, 2]).unwrap(),
1440 )]))
1441 .unwrap();
1442
1443 {
1444 let _compat = crate::compatibility::push_runmat_extensions_enabled(false);
1445 let integer_error = futures::executor::block_on(ind2word_builtin(vec![
1446 enc.clone(),
1447 poisoned_integer_vector(IntegerStorage::U64(vec![1]), 1),
1448 ]))
1449 .unwrap_err();
1450 assert_eq!(
1451 integer_error.identifier(),
1452 Some("RunMat:compatibility:Ind2wordTypedIntegerExtension")
1453 );
1454 let matrix_error = futures::executor::block_on(ind2word_builtin(vec![
1455 enc.clone(),
1456 Value::Tensor(Tensor::new(vec![1.0, 2.0, 1.0, 2.0], vec![2, 2]).unwrap()),
1457 ]))
1458 .unwrap_err();
1459 assert_eq!(
1460 matrix_error.identifier(),
1461 Some("RunMat:compatibility:Ind2wordNonvectorExtension")
1462 );
1463 }
1464
1465 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
1466 let wide = futures::executor::block_on(ind2word_builtin(vec![
1467 enc.clone(),
1468 poisoned_integer_vector(IntegerStorage::U64(vec![(1_u64 << 53) + 1]), 1),
1469 ]))
1470 .unwrap_err();
1471 assert!(wide.message().contains("exceeds vocabulary"));
1472
1473 crate::builtins::common::test_support::with_test_provider(|provider| {
1474 let tensor = Tensor::new(vec![1.0], vec![1, 1]).unwrap();
1475 let handle = crate::builtins::common::gpu_helpers::upload_tensor(provider, &tensor)
1476 .expect("resident indices");
1477 runmat_accelerate::fusion_residency::mark(&handle);
1478 let _strict = crate::compatibility::push_runmat_extensions_enabled(false);
1479 let error = futures::executor::block_on(ind2word_builtin(vec![
1480 enc.clone(),
1481 Value::GpuTensor(handle.clone()),
1482 ]))
1483 .unwrap_err();
1484 assert_eq!(
1485 error.identifier(),
1486 Some("RunMat:compatibility:Ind2wordResidentExtension")
1487 );
1488 assert!(runmat_accelerate::fusion_residency::is_resident(&handle));
1489 let _ = provider.free(&handle);
1490 });
1491 }
1492
1493 #[tokio::test]
1494 async fn is_vocabulary_word_supports_word_encoding() {
1495 let enc = word_encoding_builtin(vec![Value::StringArray(
1496 StringArray::new(vec!["RunMat".into(), "GPU".into()], vec![1, 2]).unwrap(),
1497 )])
1498 .await
1499 .unwrap();
1500 let words = Value::StringArray(
1501 StringArray::new(vec!["runmat".into(), "cpu".into()], vec![1, 2]).unwrap(),
1502 );
1503 let out = is_vocabulary_word_builtin(vec![
1504 enc,
1505 words,
1506 Value::String("IgnoreCase".into()),
1507 Value::Bool(true),
1508 ])
1509 .await
1510 .unwrap();
1511 let Value::LogicalArray(mask) = out else {
1512 panic!("expected logical array");
1513 };
1514 assert_eq!(mask.shape, vec![1, 2]);
1515 assert_eq!(mask.data, vec![1, 0]);
1516 }
1517}