1use std::cell::Cell;
4use std::collections::HashMap;
5
6use runmat_builtins::{
7 Access, BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
8 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
9 CharArray, ClassDef, LogicalArray, ObjectInstance, PropertyDef, ResolveContext, StringArray,
10 Tensor, Type, Value,
11};
12use runmat_macros::runtime_builtin;
13
14use crate::builtins::strings::core::compat::scalar_text;
15use crate::builtins::strings::text_analytics::documents::{
16 documents_from_object, TOKENIZED_DOCUMENT_CLASS,
17};
18use crate::builtins::strings::text_analytics::embeddings::{
19 build_word_lookup, word_embedding_vocabulary_from_object, WORD_EMBEDDING_CLASS,
20};
21use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult};
22
23pub const WORD_ENCODING_CLASS: &str = "wordEncoding";
24
25thread_local! {
26 static WORD_ENCODING_CLASS_REGISTERED: Cell<bool> = const { Cell::new(false) };
27}
28
29const OUT_ENCODING: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
30 name: "enc",
31 ty: BuiltinParamType::Any,
32 arity: BuiltinParamArity::Required,
33 default: None,
34 description: "Word encoding compatibility object.",
35}];
36
37const OUT_INDICES: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
38 name: "M",
39 ty: BuiltinParamType::NumericArray,
40 arity: BuiltinParamArity::Required,
41 default: None,
42 description: "Word encoding indices, with NaN for words outside the vocabulary.",
43}];
44
45const OUT_WORDS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
46 name: "words",
47 ty: BuiltinParamType::Any,
48 arity: BuiltinParamArity::Required,
49 default: None,
50 description: "Words mapped from encoding indices.",
51}];
52
53const OUT_LOGICAL: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
54 name: "tf",
55 ty: BuiltinParamType::LogicalArray,
56 arity: BuiltinParamArity::Required,
57 default: None,
58 description: "Logical membership mask.",
59}];
60
61const IN_DOCUMENTS_OR_WORDS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
62 name: "documentsOrWords",
63 ty: BuiltinParamType::Any,
64 arity: BuiltinParamArity::Required,
65 default: None,
66 description: "tokenizedDocument object or word vector.",
67}];
68
69const IN_DOCUMENTS_OR_WORDS_REST: [BuiltinParamDescriptor; 2] = [
70 BuiltinParamDescriptor {
71 name: "documentsOrWords",
72 ty: BuiltinParamType::Any,
73 arity: BuiltinParamArity::Required,
74 default: None,
75 description: "tokenizedDocument object or word vector.",
76 },
77 BuiltinParamDescriptor {
78 name: "NameValue",
79 ty: BuiltinParamType::Any,
80 arity: BuiltinParamArity::Variadic,
81 default: None,
82 description: "Name-value options: Order, MaxNumWords.",
83 },
84];
85
86const IN_WORDS: [BuiltinParamDescriptor; 2] = [
87 BuiltinParamDescriptor {
88 name: "enc",
89 ty: BuiltinParamType::Any,
90 arity: BuiltinParamArity::Required,
91 default: None,
92 description: "wordEncoding object.",
93 },
94 BuiltinParamDescriptor {
95 name: "words",
96 ty: BuiltinParamType::Any,
97 arity: BuiltinParamArity::Required,
98 default: None,
99 description: "Words to map to indices.",
100 },
101];
102
103const IN_WORDS_REST: [BuiltinParamDescriptor; 3] = [
104 BuiltinParamDescriptor {
105 name: "enc",
106 ty: BuiltinParamType::Any,
107 arity: BuiltinParamArity::Required,
108 default: None,
109 description: "wordEncoding object.",
110 },
111 BuiltinParamDescriptor {
112 name: "words",
113 ty: BuiltinParamType::Any,
114 arity: BuiltinParamArity::Required,
115 default: None,
116 description: "Words to map to indices.",
117 },
118 BuiltinParamDescriptor {
119 name: "NameValue",
120 ty: BuiltinParamType::Any,
121 arity: BuiltinParamArity::Variadic,
122 default: None,
123 description: "Name-value options: IgnoreCase.",
124 },
125];
126
127const IN_INDICES: [BuiltinParamDescriptor; 2] = [
128 BuiltinParamDescriptor {
129 name: "enc",
130 ty: BuiltinParamType::Any,
131 arity: BuiltinParamArity::Required,
132 default: None,
133 description: "wordEncoding object.",
134 },
135 BuiltinParamDescriptor {
136 name: "M",
137 ty: BuiltinParamType::NumericArray,
138 arity: BuiltinParamArity::Required,
139 default: None,
140 description: "Positive integer word encoding indices.",
141 },
142];
143
144const IN_VOCABULARY_WORDS: [BuiltinParamDescriptor; 2] = [
145 BuiltinParamDescriptor {
146 name: "embOrEnc",
147 ty: BuiltinParamType::Any,
148 arity: BuiltinParamArity::Required,
149 default: None,
150 description: "wordEmbedding or wordEncoding object.",
151 },
152 BuiltinParamDescriptor {
153 name: "words",
154 ty: BuiltinParamType::Any,
155 arity: BuiltinParamArity::Required,
156 default: None,
157 description: "Words to test.",
158 },
159];
160
161const IN_VOCABULARY_WORDS_REST: [BuiltinParamDescriptor; 3] = [
162 BuiltinParamDescriptor {
163 name: "embOrEnc",
164 ty: BuiltinParamType::Any,
165 arity: BuiltinParamArity::Required,
166 default: None,
167 description: "wordEmbedding or wordEncoding object.",
168 },
169 BuiltinParamDescriptor {
170 name: "words",
171 ty: BuiltinParamType::Any,
172 arity: BuiltinParamArity::Required,
173 default: None,
174 description: "Words to test.",
175 },
176 BuiltinParamDescriptor {
177 name: "NameValue",
178 ty: BuiltinParamType::Any,
179 arity: BuiltinParamArity::Variadic,
180 default: None,
181 description: "Name-value options: IgnoreCase.",
182 },
183];
184
185const ERROR_ENCODING_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
186 code: "RM.WORDENCODING.INVALID_INPUT",
187 identifier: Some("RunMat:wordEncoding:InvalidInput"),
188 when: "Inputs do not match a supported wordEncoding form.",
189 message: "wordEncoding received invalid input",
190};
191
192const ERROR_WORD2IND_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
193 code: "RM.WORD2IND.INVALID_INPUT",
194 identifier: Some("RunMat:word2ind:InvalidInput"),
195 when: "Inputs do not match a supported word2ind form.",
196 message: "word2ind received invalid input",
197};
198
199const ERROR_IND2WORD_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
200 code: "RM.IND2WORD.INVALID_INPUT",
201 identifier: Some("RunMat:ind2word:InvalidInput"),
202 when: "Inputs do not match a supported ind2word form.",
203 message: "ind2word received invalid input",
204};
205
206const ERROR_IS_VOCABULARY_WORD_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
207 code: "RM.ISVOCABULARYWORD.INVALID_INPUT",
208 identifier: Some("RunMat:isVocabularyWord:InvalidInput"),
209 when: "Inputs do not match a supported isVocabularyWord form.",
210 message: "isVocabularyWord received invalid input",
211};
212
213const WORD_ENCODING_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_ENCODING_INVALID_INPUT];
214const WORD2IND_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_WORD2IND_INVALID_INPUT];
215const IND2WORD_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_IND2WORD_INVALID_INPUT];
216const IS_VOCABULARY_WORD_ERRORS: [BuiltinErrorDescriptor; 1] =
217 [ERROR_IS_VOCABULARY_WORD_INVALID_INPUT];
218
219fn any_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
220 Type::Unknown
221}
222
223pub const WORD_ENCODING_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
224 signatures: &[
225 BuiltinSignatureDescriptor {
226 label: "enc = wordEncoding(documents)",
227 inputs: &IN_DOCUMENTS_OR_WORDS,
228 outputs: &OUT_ENCODING,
229 },
230 BuiltinSignatureDescriptor {
231 label: "enc = wordEncoding(words)",
232 inputs: &IN_DOCUMENTS_OR_WORDS,
233 outputs: &OUT_ENCODING,
234 },
235 BuiltinSignatureDescriptor {
236 label: "enc = wordEncoding(documents, Name, Value)",
237 inputs: &IN_DOCUMENTS_OR_WORDS_REST,
238 outputs: &OUT_ENCODING,
239 },
240 ],
241 output_mode: BuiltinOutputMode::Fixed,
242 completion_policy: BuiltinCompletionPolicy::Public,
243 errors: &WORD_ENCODING_ERRORS,
244};
245
246pub const WORD2IND_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
247 signatures: &[
248 BuiltinSignatureDescriptor {
249 label: "M = word2ind(enc, words)",
250 inputs: &IN_WORDS,
251 outputs: &OUT_INDICES,
252 },
253 BuiltinSignatureDescriptor {
254 label: "M = word2ind(enc, words, 'IgnoreCase', true)",
255 inputs: &IN_WORDS_REST,
256 outputs: &OUT_INDICES,
257 },
258 ],
259 output_mode: BuiltinOutputMode::Fixed,
260 completion_policy: BuiltinCompletionPolicy::Public,
261 errors: &WORD2IND_ERRORS,
262};
263
264pub const IND2WORD_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
265 signatures: &[BuiltinSignatureDescriptor {
266 label: "words = ind2word(enc, M)",
267 inputs: &IN_INDICES,
268 outputs: &OUT_WORDS,
269 }],
270 output_mode: BuiltinOutputMode::Fixed,
271 completion_policy: BuiltinCompletionPolicy::Public,
272 errors: &IND2WORD_ERRORS,
273};
274
275pub const IS_VOCABULARY_WORD_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
276 signatures: &[
277 BuiltinSignatureDescriptor {
278 label: "tf = isVocabularyWord(emb, words)",
279 inputs: &IN_VOCABULARY_WORDS,
280 outputs: &OUT_LOGICAL,
281 },
282 BuiltinSignatureDescriptor {
283 label: "tf = isVocabularyWord(enc, words)",
284 inputs: &IN_VOCABULARY_WORDS,
285 outputs: &OUT_LOGICAL,
286 },
287 BuiltinSignatureDescriptor {
288 label: "tf = isVocabularyWord(___, 'IgnoreCase', true)",
289 inputs: &IN_VOCABULARY_WORDS_REST,
290 outputs: &OUT_LOGICAL,
291 },
292 ],
293 output_mode: BuiltinOutputMode::Fixed,
294 completion_policy: BuiltinCompletionPolicy::Public,
295 errors: &IS_VOCABULARY_WORD_ERRORS,
296};
297
298#[runtime_builtin(
299 name = "wordEncoding",
300 category = "strings/text_analytics",
301 summary = "Create a word encoding object that maps words to indices and back.",
302 keywords = "wordEncoding,text analytics,words,indices,vocabulary",
303 accel = "sink",
304 type_resolver(any_type),
305 descriptor(crate::builtins::strings::text_analytics::encoding::WORD_ENCODING_DESCRIPTOR),
306 builtin_path = "crate::builtins::strings::text_analytics::encoding"
307)]
308async fn word_encoding_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
309 let gathered = gather_args(args, "wordEncoding").await?;
310 let (source, options) = parse_word_encoding_args(gathered)?;
311 word_encoding_object(build_word_encoding(source, options)?)
312}
313
314#[runtime_builtin(
315 name = "word2ind",
316 category = "strings/text_analytics",
317 summary = "Map words to indices in a wordEncoding object.",
318 keywords = "word2ind,wordEncoding,text analytics,indices,vocabulary",
319 accel = "sink",
320 type_resolver(any_type),
321 descriptor(crate::builtins::strings::text_analytics::encoding::WORD2IND_DESCRIPTOR),
322 builtin_path = "crate::builtins::strings::text_analytics::encoding"
323)]
324async fn word2ind_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
325 let gathered = gather_args(args, "word2ind").await?;
326 let (object, words, options) = parse_word2ind_args(gathered)?;
327 let encoding = word_encoding_from_object(&object, "word2ind")?;
328 let lookup = build_word_lookup(&encoding.vocabulary, options.ignore_case);
329 let indices = words
330 .words
331 .into_iter()
332 .map(|word| {
333 let key = if options.ignore_case {
334 word.to_lowercase()
335 } else {
336 word
337 };
338 lookup
339 .get(&key)
340 .map(|idx| (*idx + 1) as f64)
341 .unwrap_or(f64::NAN)
342 })
343 .collect::<Vec<_>>();
344 Tensor::new(indices, words.shape)
345 .map(Value::Tensor)
346 .map_err(|err| encoding_error("word2ind", err))
347}
348
349#[runtime_builtin(
350 name = "ind2word",
351 category = "strings/text_analytics",
352 summary = "Map wordEncoding indices back to words.",
353 keywords = "ind2word,wordEncoding,text analytics,indices,vocabulary",
354 accel = "sink",
355 type_resolver(any_type),
356 descriptor(crate::builtins::strings::text_analytics::encoding::IND2WORD_DESCRIPTOR),
357 builtin_path = "crate::builtins::strings::text_analytics::encoding"
358)]
359async fn ind2word_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
360 let gathered = gather_args(args, "ind2word").await?;
361 let (object, indices) = parse_ind2word_args(gathered)?;
362 let encoding = word_encoding_from_object(&object, "ind2word")?;
363 let words = indices
364 .values
365 .into_iter()
366 .map(|idx| {
367 let word_idx = positive_index(idx, encoding.vocabulary.len(), "ind2word")?;
368 Ok(encoding.vocabulary[word_idx].clone())
369 })
370 .collect::<BuiltinResult<Vec<_>>>()?;
371 StringArray::new(words, indices.shape)
372 .map(Value::StringArray)
373 .map_err(|err| encoding_error("ind2word", err))
374}
375
376#[runtime_builtin(
377 name = "isVocabularyWord",
378 category = "strings/text_analytics",
379 summary = "Test whether words are in a wordEmbedding or wordEncoding vocabulary.",
380 keywords = "isVocabularyWord,wordEmbedding,wordEncoding,text analytics,vocabulary",
381 accel = "sink",
382 type_resolver(any_type),
383 descriptor(crate::builtins::strings::text_analytics::encoding::IS_VOCABULARY_WORD_DESCRIPTOR),
384 builtin_path = "crate::builtins::strings::text_analytics::encoding"
385)]
386async fn is_vocabulary_word_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
387 let gathered = gather_args(args, "isVocabularyWord").await?;
388 let (object, words, options) = parse_is_vocabulary_word_args(gathered)?;
389 let vocabulary = if object.is_class(WORD_ENCODING_CLASS) {
390 word_encoding_from_object(&object, "isVocabularyWord")?.vocabulary
391 } else if object.is_class(WORD_EMBEDDING_CLASS) {
392 word_embedding_vocabulary_from_object(&object, "isVocabularyWord")?
393 } else {
394 return Err(encoding_error(
395 "isVocabularyWord",
396 format!(
397 "isVocabularyWord: expected wordEmbedding or wordEncoding object, got {}",
398 object.class_name
399 ),
400 ));
401 };
402 let lookup = build_word_lookup(&vocabulary, options.ignore_case);
403 let flags = words
404 .words
405 .into_iter()
406 .map(|word| {
407 let key = if options.ignore_case {
408 word.to_lowercase()
409 } else {
410 word
411 };
412 u8::from(lookup.contains_key(&key))
413 })
414 .collect::<Vec<_>>();
415 LogicalArray::new(flags, words.shape)
416 .map(Value::LogicalArray)
417 .map_err(|err| encoding_error("isVocabularyWord", err))
418}
419
420async fn gather_args(args: Vec<Value>, fn_name: &str) -> BuiltinResult<Vec<Value>> {
421 let mut out = Vec::with_capacity(args.len());
422 for arg in args {
423 out.push(gather_if_needed_async(&arg).await.map_err(|err| {
424 encoding_error(fn_name, format!("{fn_name}: failed to gather input: {err}"))
425 })?);
426 }
427 Ok(out)
428}
429
430#[derive(Clone, Debug)]
431pub(in crate::builtins::strings::text_analytics) struct WordEncodingModel {
432 pub vocabulary: Vec<String>,
433}
434
435pub(in crate::builtins::strings::text_analytics) fn word_encoding_from_object(
436 object: &ObjectInstance,
437 fn_name: &str,
438) -> BuiltinResult<WordEncodingModel> {
439 if !object.is_class(WORD_ENCODING_CLASS) {
440 return Err(encoding_error(
441 fn_name,
442 format!(
443 "{fn_name}: expected wordEncoding object, got {}",
444 object.class_name
445 ),
446 ));
447 }
448 let vocabulary = match object.properties.get("Vocabulary") {
449 Some(Value::StringArray(array)) => array.data.clone(),
450 other => {
451 return Err(encoding_error(
452 fn_name,
453 format!(
454 "{fn_name}: wordEncoding object has invalid Vocabulary property: {other:?}"
455 ),
456 ));
457 }
458 };
459 match object.properties.get("NumWords") {
460 Some(Value::Num(value)) if *value == vocabulary.len() as f64 => {}
461 other => {
462 return Err(encoding_error(
463 fn_name,
464 format!("{fn_name}: wordEncoding object has invalid NumWords property: {other:?}"),
465 ));
466 }
467 }
468 Ok(WordEncodingModel { vocabulary })
469}
470
471fn word_encoding_object(model: WordEncodingModel) -> BuiltinResult<Value> {
472 ensure_word_encoding_class_registered();
473 let mut object = ObjectInstance::new(WORD_ENCODING_CLASS.to_string());
474 object.properties.insert(
475 "NumWords".to_string(),
476 Value::Num(model.vocabulary.len() as f64),
477 );
478 object.properties.insert(
479 "Vocabulary".to_string(),
480 Value::StringArray(
481 StringArray::new(model.vocabulary.clone(), vec![1, model.vocabulary.len()])
482 .map_err(|err| encoding_error("wordEncoding", err))?,
483 ),
484 );
485 Ok(Value::Object(object))
486}
487
488fn ensure_word_encoding_class_registered() {
489 WORD_ENCODING_CLASS_REGISTERED.with(|registered| {
490 if registered.get() {
491 return;
492 }
493 let mut properties = HashMap::new();
494 for name in ["NumWords", "Vocabulary"] {
495 properties.insert(name.to_string(), property_def(name));
496 }
497 runmat_builtins::register_class(ClassDef {
498 name: WORD_ENCODING_CLASS.to_string(),
499 parent: None,
500 properties,
501 methods: HashMap::new(),
502 });
503 registered.set(true);
504 });
505}
506
507fn property_def(name: &str) -> PropertyDef {
508 PropertyDef {
509 name: name.to_string(),
510 is_static: false,
511 is_constant: false,
512 is_dependent: false,
513 get_access: Access::Public,
514 set_access: Access::Public,
515 default_value: None,
516 }
517}
518
519enum EncodingSource {
520 Documents(Vec<Vec<String>>),
521 Words(Vec<String>),
522}
523
524#[derive(Clone, Copy, Debug, PartialEq, Eq)]
525enum EncodingOrder {
526 FirstSeen,
527 Frequency,
528}
529
530#[derive(Clone, Copy, Debug)]
531struct WordEncodingOptions {
532 order: EncodingOrder,
533 max_num_words: Option<usize>,
534}
535
536impl Default for WordEncodingOptions {
537 fn default() -> Self {
538 Self {
539 order: EncodingOrder::FirstSeen,
540 max_num_words: None,
541 }
542 }
543}
544
545fn parse_word_encoding_args(
546 args: Vec<Value>,
547) -> BuiltinResult<(EncodingSource, WordEncodingOptions)> {
548 if args.is_empty() {
549 return Err(encoding_error(
550 "wordEncoding",
551 "wordEncoding: expected tokenizedDocument object or word vector",
552 ));
553 }
554 if !(args.len() - 1).is_multiple_of(2) {
555 return Err(encoding_error(
556 "wordEncoding",
557 "wordEncoding: name-value options must be paired",
558 ));
559 }
560 let source = match &args[0] {
561 Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => {
562 EncodingSource::Documents(documents_from_object(object, "wordEncoding")?)
563 }
564 Value::Object(object) => {
565 return Err(encoding_error(
566 "wordEncoding",
567 format!(
568 "wordEncoding: expected tokenizedDocument object or word vector, got {}",
569 object.class_name
570 ),
571 ));
572 }
573 value => EncodingSource::Words(word_input_from_value(value, "wordEncoding")?.words),
574 };
575 if matches!(source, EncodingSource::Words(_)) && args.len() > 1 {
576 return Err(encoding_error(
577 "wordEncoding",
578 "wordEncoding: Order and MaxNumWords options are only supported for tokenizedDocument input",
579 ));
580 }
581 let mut options = WordEncodingOptions::default();
582 let mut idx = 1usize;
583 while idx < args.len() {
584 let name = scalar_text(&args[idx], "wordEncoding")
585 .map_err(|err| encoding_error("wordEncoding", err.to_string()))?
586 .to_ascii_lowercase();
587 match name.as_str() {
588 "order" => {
589 let value = scalar_text(&args[idx + 1], "wordEncoding")
590 .map_err(|err| encoding_error("wordEncoding", err.to_string()))?
591 .to_ascii_lowercase();
592 options.order = match value.as_str() {
593 "first-seen" => EncodingOrder::FirstSeen,
594 "frequency" => EncodingOrder::Frequency,
595 other => {
596 return Err(encoding_error(
597 "wordEncoding",
598 format!(
599 "wordEncoding: Order must be 'first-seen' or 'frequency', got '{other}'"
600 ),
601 ));
602 }
603 };
604 }
605 "maxnumwords" => {
606 options.max_num_words = parse_max_num_words(&args[idx + 1])?;
607 }
608 other => {
609 return Err(encoding_error(
610 "wordEncoding",
611 format!("wordEncoding: unsupported option '{other}'"),
612 ));
613 }
614 }
615 idx += 2;
616 }
617 Ok((source, options))
618}
619
620fn build_word_encoding(
621 source: EncodingSource,
622 options: WordEncodingOptions,
623) -> BuiltinResult<WordEncodingModel> {
624 let words = match source {
625 EncodingSource::Documents(documents) => documents.into_iter().flatten().collect::<Vec<_>>(),
626 EncodingSource::Words(words) => words,
627 };
628 let mut counts = HashMap::<String, (usize, usize)>::new();
629 for (pos, word) in words.into_iter().enumerate() {
630 let entry = counts.entry(word).or_insert((0, pos));
631 entry.0 += 1;
632 }
633 let mut ranked = counts
634 .into_iter()
635 .map(|(word, (count, first_pos))| (word, count, first_pos))
636 .collect::<Vec<_>>();
637 match options.order {
638 EncodingOrder::FirstSeen => ranked.sort_by(|left, right| left.2.cmp(&right.2)),
639 EncodingOrder::Frequency => {
640 ranked.sort_by(|left, right| right.1.cmp(&left.1).then(left.2.cmp(&right.2)))
641 }
642 }
643 if let Some(max) = options.max_num_words {
644 ranked.truncate(max);
645 }
646 Ok(WordEncodingModel {
647 vocabulary: ranked.into_iter().map(|(word, _, _)| word).collect(),
648 })
649}
650
651#[derive(Clone, Copy, Debug, Default)]
652struct LookupOptions {
653 ignore_case: bool,
654}
655
656fn parse_word2ind_args(
657 args: Vec<Value>,
658) -> BuiltinResult<(ObjectInstance, WordInput, LookupOptions)> {
659 if args.len() < 2 {
660 return Err(encoding_error(
661 "word2ind",
662 "word2ind: expected word2ind(enc, words)",
663 ));
664 }
665 let object = object_arg(&args[0], "word2ind", "wordEncoding")?;
666 let words = word_input_from_value(&args[1], "word2ind")?;
667 let options = parse_lookup_options(&args[2..], "word2ind")?;
668 Ok((object, words, options))
669}
670
671fn parse_is_vocabulary_word_args(
672 args: Vec<Value>,
673) -> BuiltinResult<(ObjectInstance, WordInput, LookupOptions)> {
674 if args.len() < 2 {
675 return Err(encoding_error(
676 "isVocabularyWord",
677 "isVocabularyWord: expected isVocabularyWord(embOrEnc, words)",
678 ));
679 }
680 let object = object_arg(
681 &args[0],
682 "isVocabularyWord",
683 "wordEmbedding or wordEncoding",
684 )?;
685 let words = word_input_from_value(&args[1], "isVocabularyWord")?;
686 let options = parse_lookup_options(&args[2..], "isVocabularyWord")?;
687 Ok((object, words, options))
688}
689
690fn parse_ind2word_args(args: Vec<Value>) -> BuiltinResult<(ObjectInstance, NumericInput)> {
691 if args.len() != 2 {
692 return Err(encoding_error(
693 "ind2word",
694 "ind2word: expected ind2word(enc, M)",
695 ));
696 }
697 let object = object_arg(&args[0], "ind2word", "wordEncoding")?;
698 let indices = numeric_input_from_value(&args[1], "ind2word")?;
699 Ok((object, indices))
700}
701
702fn object_arg(value: &Value, fn_name: &str, expected: &str) -> BuiltinResult<ObjectInstance> {
703 match value {
704 Value::Object(object) => Ok(object.clone()),
705 other => Err(encoding_error(
706 fn_name,
707 format!("{fn_name}: expected {expected} object, got {other:?}"),
708 )),
709 }
710}
711
712fn parse_lookup_options(args: &[Value], fn_name: &str) -> BuiltinResult<LookupOptions> {
713 if !args.len().is_multiple_of(2) {
714 return Err(encoding_error(
715 fn_name,
716 format!("{fn_name}: name-value options must be paired"),
717 ));
718 }
719 let mut options = LookupOptions::default();
720 let mut idx = 0usize;
721 while idx < args.len() {
722 let name = scalar_text(&args[idx], fn_name)
723 .map_err(|err| encoding_error(fn_name, err.to_string()))?
724 .to_ascii_lowercase();
725 match name.as_str() {
726 "ignorecase" => options.ignore_case = parse_bool_scalar(&args[idx + 1], fn_name)?,
727 other => {
728 return Err(encoding_error(
729 fn_name,
730 format!("{fn_name}: unsupported option '{other}'"),
731 ));
732 }
733 }
734 idx += 2;
735 }
736 Ok(options)
737}
738
739struct WordInput {
740 words: Vec<String>,
741 shape: Vec<usize>,
742}
743
744fn word_input_from_value(value: &Value, fn_name: &str) -> BuiltinResult<WordInput> {
745 match value {
746 Value::String(text) => Ok(WordInput {
747 words: vec![text.clone()],
748 shape: vec![1, 1],
749 }),
750 Value::StringArray(array) => Ok(WordInput {
751 words: array.data.clone(),
752 shape: array.shape.clone(),
753 }),
754 Value::CharArray(array) if array.rows <= 1 => Ok(WordInput {
755 words: vec![char_row_to_string(array)],
756 shape: vec![1, 1],
757 }),
758 Value::CharArray(array) => {
759 let mut words = Vec::with_capacity(array.rows);
760 for row in 0..array.rows {
761 let mut text = String::with_capacity(array.cols);
762 for col in 0..array.cols {
763 text.push(array.data[row + col * array.rows]);
764 }
765 words.push(text.trim_end().to_string());
766 }
767 Ok(WordInput {
768 words,
769 shape: vec![array.rows, 1],
770 })
771 }
772 Value::Cell(cell) => {
773 let words = cell
774 .data
775 .iter()
776 .map(|item| {
777 scalar_text(item, fn_name)
778 .map_err(|err| encoding_error(fn_name, err.to_string()))
779 })
780 .collect::<BuiltinResult<Vec<_>>>()?;
781 Ok(WordInput {
782 words,
783 shape: cell.shape.clone(),
784 })
785 }
786 other => Err(encoding_error(
787 fn_name,
788 format!("{fn_name}: expected string, character vector, or cell array of words, got {other:?}"),
789 )),
790 }
791}
792
793struct NumericInput {
794 values: Vec<f64>,
795 shape: Vec<usize>,
796}
797
798fn numeric_input_from_value(value: &Value, fn_name: &str) -> BuiltinResult<NumericInput> {
799 match value {
800 Value::Num(value) => Ok(NumericInput {
801 values: vec![*value],
802 shape: vec![1, 1],
803 }),
804 Value::Int(value) => Ok(NumericInput {
805 values: vec![int_value_to_f64(value)],
806 shape: vec![1, 1],
807 }),
808 Value::Tensor(tensor) => Ok(NumericInput {
809 values: tensor.data.clone(),
810 shape: tensor.shape.clone(),
811 }),
812 other => Err(encoding_error(
813 fn_name,
814 format!("{fn_name}: expected numeric positive integer indices, got {other:?}"),
815 )),
816 }
817}
818
819fn positive_index(value: f64, len: usize, fn_name: &str) -> BuiltinResult<usize> {
820 if !value.is_finite() || value < 1.0 || value.fract() != 0.0 {
821 return Err(encoding_error(
822 fn_name,
823 format!("{fn_name}: indices must be positive integers, got {value}"),
824 ));
825 }
826 let idx = value as usize;
827 if idx > len {
828 return Err(encoding_error(
829 fn_name,
830 format!("{fn_name}: index {idx} exceeds vocabulary size {len}"),
831 ));
832 }
833 Ok(idx - 1)
834}
835
836fn parse_max_num_words(value: &Value) -> BuiltinResult<Option<usize>> {
837 let n = numeric_scalar(value, "wordEncoding", "MaxNumWords")?;
838 if n.is_infinite() && n.is_sign_positive() {
839 return Ok(None);
840 }
841 if !n.is_finite() || n < 1.0 || n.fract() != 0.0 {
842 return Err(encoding_error(
843 "wordEncoding",
844 format!("wordEncoding: MaxNumWords must be a positive integer or Inf, got {n}"),
845 ));
846 }
847 Ok(Some(n as usize))
848}
849
850fn numeric_scalar(value: &Value, fn_name: &str, option: &str) -> BuiltinResult<f64> {
851 match value {
852 Value::Num(value) => Ok(*value),
853 Value::Int(value) => Ok(int_value_to_f64(value)),
854 Value::Tensor(tensor) if tensor.data.len() == 1 => Ok(tensor.data[0]),
855 other => Err(encoding_error(
856 fn_name,
857 format!("{fn_name}: {option} must be a numeric scalar, got {other:?}"),
858 )),
859 }
860}
861
862fn parse_bool_scalar(value: &Value, fn_name: &str) -> BuiltinResult<bool> {
863 match value {
864 Value::Bool(value) => Ok(*value),
865 Value::Num(value) if *value == 0.0 || *value == 1.0 => Ok(*value != 0.0),
866 Value::Tensor(tensor) if tensor.data.len() == 1 => match tensor.data[0] {
867 0.0 => Ok(false),
868 1.0 => Ok(true),
869 other => Err(encoding_error(
870 fn_name,
871 format!("{fn_name}: logical scalar option must be true or false, got {other}"),
872 )),
873 },
874 Value::LogicalArray(array) if array.data.len() == 1 => Ok(array.data[0] != 0),
875 other => Err(encoding_error(
876 fn_name,
877 format!("{fn_name}: logical scalar option must be true or false, got {other:?}"),
878 )),
879 }
880}
881
882fn int_value_to_f64(value: &runmat_builtins::IntValue) -> f64 {
883 match value {
884 runmat_builtins::IntValue::I8(value) => *value as f64,
885 runmat_builtins::IntValue::I16(value) => *value as f64,
886 runmat_builtins::IntValue::I32(value) => *value as f64,
887 runmat_builtins::IntValue::I64(value) => *value as f64,
888 runmat_builtins::IntValue::U8(value) => *value as f64,
889 runmat_builtins::IntValue::U16(value) => *value as f64,
890 runmat_builtins::IntValue::U32(value) => *value as f64,
891 runmat_builtins::IntValue::U64(value) => *value as f64,
892 }
893}
894
895fn char_row_to_string(array: &CharArray) -> String {
896 array.data.iter().collect()
897}
898
899fn encoding_error(fn_name: &str, message: impl Into<String>) -> crate::RuntimeError {
900 let descriptor = match fn_name {
901 "word2ind" => ERROR_WORD2IND_INVALID_INPUT,
902 "ind2word" => ERROR_IND2WORD_INVALID_INPUT,
903 "isVocabularyWord" => ERROR_IS_VOCABULARY_WORD_INVALID_INPUT,
904 _ => ERROR_ENCODING_INVALID_INPUT,
905 };
906 let builder = build_runtime_error(message.into()).with_builtin(fn_name);
907 match descriptor.identifier {
908 Some(identifier) => builder.with_identifier(identifier).build(),
909 None => builder.build(),
910 }
911}
912
913#[cfg(test)]
914mod tests {
915 use super::*;
916 use runmat_builtins::CellArray;
917
918 fn tokenized_document_object(rows: Vec<Vec<&str>>) -> ObjectInstance {
919 let values = rows
920 .into_iter()
921 .map(|row| {
922 let len = row.len();
923 Value::StringArray(
924 StringArray::new(
925 row.into_iter().map(str::to_string).collect::<Vec<_>>(),
926 vec![1, len],
927 )
928 .unwrap(),
929 )
930 })
931 .collect::<Vec<_>>();
932 let rows = values.len();
933 let documents = CellArray::new(values, rows, 1).unwrap();
934 let mut object = ObjectInstance::new(TOKENIZED_DOCUMENT_CLASS.to_string());
935 object
936 .properties
937 .insert("Documents".to_string(), Value::Cell(documents));
938 object
939 }
940
941 #[tokio::test]
942 async fn word_encoding_builds_first_seen_and_frequency_vocabularies() {
943 let documents = Value::Object(tokenized_document_object(vec![
944 vec!["beta", "alpha", "beta"],
945 vec!["gamma", "alpha", "beta"],
946 ]));
947 let first_seen = word_encoding_builtin(vec![documents.clone()])
948 .await
949 .unwrap();
950 let Value::Object(first_seen) = first_seen else {
951 panic!("expected object");
952 };
953 let model = word_encoding_from_object(&first_seen, "test").unwrap();
954 assert_eq!(model.vocabulary, vec!["beta", "alpha", "gamma"]);
955
956 let frequency = word_encoding_builtin(vec![
957 documents,
958 Value::String("Order".into()),
959 Value::String("frequency".into()),
960 Value::String("MaxNumWords".into()),
961 Value::Num(2.0),
962 ])
963 .await
964 .unwrap();
965 let Value::Object(frequency) = frequency else {
966 panic!("expected object");
967 };
968 let model = word_encoding_from_object(&frequency, "test").unwrap();
969 assert_eq!(model.vocabulary, vec!["beta", "alpha"]);
970 }
971
972 #[tokio::test]
973 async fn word_encoding_accepts_word_arrays_and_validates_options() {
974 let words = Value::StringArray(
975 StringArray::new(vec!["red".into(), "blue".into(), "red".into()], vec![1, 3]).unwrap(),
976 );
977 let enc = word_encoding_builtin(vec![words]).await.unwrap();
978 let Value::Object(enc) = enc else {
979 panic!("expected object");
980 };
981 assert_eq!(enc.properties.get("NumWords"), Some(&Value::Num(2.0)));
982
983 let err = word_encoding_builtin(vec![
984 Value::String("x".into()),
985 Value::String("Order".into()),
986 Value::String("frequency".into()),
987 ])
988 .await
989 .unwrap_err();
990 assert!(
991 err.to_string()
992 .contains("only supported for tokenizedDocument input"),
993 "{err}"
994 );
995 }
996
997 #[tokio::test]
998 async fn word2ind_preserves_shape_and_supports_ignore_case() {
999 let enc = word_encoding_builtin(vec![Value::StringArray(
1000 StringArray::new(vec!["Alpha".into(), "beta".into()], vec![1, 2]).unwrap(),
1001 )])
1002 .await
1003 .unwrap();
1004 let words = Value::StringArray(
1005 StringArray::new(
1006 vec![
1007 "beta".into(),
1008 "missing".into(),
1009 "alpha".into(),
1010 "Alpha".into(),
1011 ],
1012 vec![2, 2],
1013 )
1014 .unwrap(),
1015 );
1016 let out = word2ind_builtin(vec![
1017 enc,
1018 words,
1019 Value::String("IgnoreCase".into()),
1020 Value::Bool(true),
1021 ])
1022 .await
1023 .unwrap();
1024 let Value::Tensor(indices) = out else {
1025 panic!("expected tensor");
1026 };
1027 assert_eq!(indices.shape, vec![2, 2]);
1028 assert_eq!(indices.data[0], 2.0);
1029 assert!(indices.data[1].is_nan());
1030 assert_eq!(indices.data[2], 1.0);
1031 assert_eq!(indices.data[3], 1.0);
1032 }
1033
1034 #[tokio::test]
1035 async fn ind2word_preserves_numeric_shape_and_rejects_bad_indices() {
1036 let enc = word_encoding_builtin(vec![Value::StringArray(
1037 StringArray::new(
1038 vec!["red".into(), "blue".into(), "green".into()],
1039 vec![1, 3],
1040 )
1041 .unwrap(),
1042 )])
1043 .await
1044 .unwrap();
1045 let out = ind2word_builtin(vec![
1046 enc.clone(),
1047 Value::Tensor(Tensor::new(vec![1.0, 3.0], vec![1, 2]).unwrap()),
1048 ])
1049 .await
1050 .unwrap();
1051 let Value::StringArray(words) = out else {
1052 panic!("expected string array");
1053 };
1054 assert_eq!(words.shape, vec![1, 2]);
1055 assert_eq!(words.data, vec!["red", "green"]);
1056
1057 let err = ind2word_builtin(vec![enc, Value::Num(4.0)])
1058 .await
1059 .unwrap_err();
1060 assert!(err.to_string().contains("exceeds vocabulary"), "{err}");
1061 }
1062
1063 #[tokio::test]
1064 async fn is_vocabulary_word_supports_word_encoding() {
1065 let enc = word_encoding_builtin(vec![Value::StringArray(
1066 StringArray::new(vec!["RunMat".into(), "GPU".into()], vec![1, 2]).unwrap(),
1067 )])
1068 .await
1069 .unwrap();
1070 let words = Value::StringArray(
1071 StringArray::new(vec!["runmat".into(), "cpu".into()], vec![1, 2]).unwrap(),
1072 );
1073 let out = is_vocabulary_word_builtin(vec![
1074 enc,
1075 words,
1076 Value::String("IgnoreCase".into()),
1077 Value::Bool(true),
1078 ])
1079 .await
1080 .unwrap();
1081 let Value::LogicalArray(mask) = out else {
1082 panic!("expected logical array");
1083 };
1084 assert_eq!(mask.shape, vec![1, 2]);
1085 assert_eq!(mask.data, vec![1, 0]);
1086 }
1087}