Skip to main content

runmat_runtime/builtins/strings/text_analytics/
encoding.rs

1//! Word encoding compatibility objects and word/index lookup helpers.
2
3use 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}