Skip to main content

runmat_runtime/builtins/strings/text_analytics/
embeddings.rs

1//! Word embedding compatibility objects and lookup helpers.
2
3use std::cell::Cell;
4use std::cmp::Ordering;
5use std::collections::HashMap;
6use std::io::{Cursor, Read, Write};
7use std::path::Path;
8
9use runmat_builtins::{
10    Access, BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
11    BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
12    CellArray, CharArray, ClassDef, ObjectInstance, PropertyDef, ResolveContext, StringArray,
13    Tensor, Type, Value,
14};
15use runmat_filesystem::File;
16use runmat_macros::runtime_builtin;
17
18use crate::builtins::strings::core::compat::scalar_text;
19use crate::builtins::strings::text_analytics::documents::{
20    document_shape_from_object, documents_from_object, TOKENIZED_DOCUMENT_CLASS,
21};
22use crate::builtins::strings::text_analytics::encoding::{
23    word_encoding_from_object, WordEncodingModel, WORD_ENCODING_CLASS,
24};
25use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult};
26
27pub const WORD_EMBEDDING_CLASS: &str = "wordEmbedding";
28const VECTOR_PROPERTY: &str = "__Vectors";
29const MAX_EMBEDDING_FILE_BYTES: u64 = 512 * 1024 * 1024;
30const MAX_ZIP_ENTRIES: usize = 256;
31const MAX_TRAINED_DENSE_VALUES: usize = 20_000_000;
32const MAX_DOC2SEQUENCE_DENSE_VALUES: usize = 50_000_000;
33
34thread_local! {
35    static WORD_EMBEDDING_CLASS_REGISTERED: Cell<bool> = const { Cell::new(false) };
36}
37
38const OUT_EMBEDDING: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
39    name: "emb",
40    ty: BuiltinParamType::Any,
41    arity: BuiltinParamArity::Required,
42    default: None,
43    description: "Word embedding compatibility object.",
44}];
45
46const OUT_MATRIX: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
47    name: "M",
48    ty: BuiltinParamType::NumericArray,
49    arity: BuiltinParamArity::Required,
50    default: None,
51    description: "Embedding vectors, one word per row.",
52}];
53
54const OUT_WORDS: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
55    name: "words",
56    ty: BuiltinParamType::Any,
57    arity: BuiltinParamArity::Required,
58    default: None,
59    description: "Closest vocabulary words.",
60}];
61
62const OUT_WORDS_DIST: [BuiltinParamDescriptor; 2] = [
63    BuiltinParamDescriptor {
64        name: "words",
65        ty: BuiltinParamType::Any,
66        arity: BuiltinParamArity::Required,
67        default: None,
68        description: "Closest vocabulary words.",
69    },
70    BuiltinParamDescriptor {
71        name: "dist",
72        ty: BuiltinParamType::NumericArray,
73        arity: BuiltinParamArity::Required,
74        default: None,
75        description: "Distances to input vectors.",
76    },
77];
78
79const OUT_SEQUENCES: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
80    name: "sequences",
81    ty: BuiltinParamType::Any,
82    arity: BuiltinParamArity::Required,
83    default: None,
84    description: "Cell array of document embedding-vector or word-index sequences.",
85}];
86
87const OUT_NONE: [BuiltinParamDescriptor; 0] = [];
88const NO_INPUTS: [BuiltinParamDescriptor; 0] = [];
89
90const IN_FILENAME: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
91    name: "filename",
92    ty: BuiltinParamType::Any,
93    arity: BuiltinParamArity::Required,
94    default: None,
95    description: "UTF-8 word2vec/GloVe text file or zip file containing one.",
96}];
97
98const IN_EMBEDDING_FILENAME: [BuiltinParamDescriptor; 2] = [
99    BuiltinParamDescriptor {
100        name: "emb",
101        ty: BuiltinParamType::Any,
102        arity: BuiltinParamArity::Required,
103        default: None,
104        description: "wordEmbedding object.",
105    },
106    BuiltinParamDescriptor {
107        name: "filename",
108        ty: BuiltinParamType::Any,
109        arity: BuiltinParamArity::Required,
110        default: None,
111        description: "Target UTF-8 word2vec text file.",
112    },
113];
114
115const IN_TRAIN_SOURCE: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
116    name: "source",
117    ty: BuiltinParamType::Any,
118    arity: BuiltinParamArity::Required,
119    default: None,
120    description: "UTF-8 text filename or tokenizedDocument object.",
121}];
122
123const IN_TRAIN_SOURCE_REST: [BuiltinParamDescriptor; 2] = [
124    BuiltinParamDescriptor {
125        name: "source",
126        ty: BuiltinParamType::Any,
127        arity: BuiltinParamArity::Required,
128        default: None,
129        description: "UTF-8 text filename or tokenizedDocument object.",
130    },
131    BuiltinParamDescriptor {
132        name: "NameValue",
133        ty: BuiltinParamType::Any,
134        arity: BuiltinParamArity::Variadic,
135        default: None,
136        description: "Name-value options controlling local deterministic embedding training.",
137    },
138];
139
140const IN_WORDS: [BuiltinParamDescriptor; 2] = [
141    BuiltinParamDescriptor {
142        name: "emb",
143        ty: BuiltinParamType::Any,
144        arity: BuiltinParamArity::Required,
145        default: None,
146        description: "wordEmbedding object.",
147    },
148    BuiltinParamDescriptor {
149        name: "words",
150        ty: BuiltinParamType::Any,
151        arity: BuiltinParamArity::Required,
152        default: None,
153        description: "Words to map to vectors.",
154    },
155];
156
157const IN_WORDS_REST: [BuiltinParamDescriptor; 3] = [
158    BuiltinParamDescriptor {
159        name: "emb",
160        ty: BuiltinParamType::Any,
161        arity: BuiltinParamArity::Required,
162        default: None,
163        description: "wordEmbedding object.",
164    },
165    BuiltinParamDescriptor {
166        name: "words",
167        ty: BuiltinParamType::Any,
168        arity: BuiltinParamArity::Required,
169        default: None,
170        description: "Words to map to vectors.",
171    },
172    BuiltinParamDescriptor {
173        name: "NameValue",
174        ty: BuiltinParamType::Any,
175        arity: BuiltinParamArity::Variadic,
176        default: None,
177        description: "Name-value options: IgnoreCase.",
178    },
179];
180
181const IN_VECTORS: [BuiltinParamDescriptor; 2] = [
182    BuiltinParamDescriptor {
183        name: "emb",
184        ty: BuiltinParamType::Any,
185        arity: BuiltinParamArity::Required,
186        default: None,
187        description: "wordEmbedding object.",
188    },
189    BuiltinParamDescriptor {
190        name: "M",
191        ty: BuiltinParamType::NumericArray,
192        arity: BuiltinParamArity::Required,
193        default: None,
194        description: "Embedding vectors, one vector per row.",
195    },
196];
197
198const IN_VECTORS_REST: [BuiltinParamDescriptor; 4] = [
199    BuiltinParamDescriptor {
200        name: "emb",
201        ty: BuiltinParamType::Any,
202        arity: BuiltinParamArity::Required,
203        default: None,
204        description: "wordEmbedding object.",
205    },
206    BuiltinParamDescriptor {
207        name: "M",
208        ty: BuiltinParamType::NumericArray,
209        arity: BuiltinParamArity::Required,
210        default: None,
211        description: "Embedding vectors, one vector per row.",
212    },
213    BuiltinParamDescriptor {
214        name: "k",
215        ty: BuiltinParamType::NumericScalar,
216        arity: BuiltinParamArity::Optional,
217        default: Some("1"),
218        description: "Number of nearest words.",
219    },
220    BuiltinParamDescriptor {
221        name: "NameValue",
222        ty: BuiltinParamType::Any,
223        arity: BuiltinParamArity::Variadic,
224        default: None,
225        description: "Name-value options: Distance ('cosine' or 'euclidean').",
226    },
227];
228
229const IN_MAP_DOCUMENTS: [BuiltinParamDescriptor; 2] = [
230    BuiltinParamDescriptor {
231        name: "embOrEnc",
232        ty: BuiltinParamType::Any,
233        arity: BuiltinParamArity::Required,
234        default: None,
235        description: "wordEmbedding or wordEncoding object.",
236    },
237    BuiltinParamDescriptor {
238        name: "documents",
239        ty: BuiltinParamType::Any,
240        arity: BuiltinParamArity::Required,
241        default: None,
242        description: "tokenizedDocument object.",
243    },
244];
245
246const IN_MAP_DOCUMENTS_REST: [BuiltinParamDescriptor; 3] = [
247    BuiltinParamDescriptor {
248        name: "embOrEnc",
249        ty: BuiltinParamType::Any,
250        arity: BuiltinParamArity::Required,
251        default: None,
252        description: "wordEmbedding or wordEncoding object.",
253    },
254    BuiltinParamDescriptor {
255        name: "documents",
256        ty: BuiltinParamType::Any,
257        arity: BuiltinParamArity::Required,
258        default: None,
259        description: "tokenizedDocument object.",
260    },
261    BuiltinParamDescriptor {
262        name: "NameValue",
263        ty: BuiltinParamType::Any,
264        arity: BuiltinParamArity::Variadic,
265        default: None,
266        description: "Name-value options: UnknownWord, PaddingDirection, PaddingValue, Length.",
267    },
268];
269
270const ERROR_READ_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
271    code: "RM.READWORDEMBEDDING.INVALID_INPUT",
272    identifier: Some("RunMat:readWordEmbedding:InvalidInput"),
273    when: "Inputs do not match a supported readWordEmbedding form.",
274    message: "readWordEmbedding received invalid input",
275};
276
277const ERROR_WRITE_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
278    code: "RM.WRITEWORDEMBEDDING.INVALID_INPUT",
279    identifier: Some("RunMat:writeWordEmbedding:InvalidInput"),
280    when: "Inputs do not match the supported writeWordEmbedding form.",
281    message: "writeWordEmbedding received invalid input",
282};
283
284const ERROR_WORD2VEC_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
285    code: "RM.WORD2VEC.INVALID_INPUT",
286    identifier: Some("RunMat:word2vec:InvalidInput"),
287    when: "Inputs do not match a supported word2vec form.",
288    message: "word2vec received invalid input",
289};
290
291const ERROR_VEC2WORD_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
292    code: "RM.VEC2WORD.INVALID_INPUT",
293    identifier: Some("RunMat:vec2word:InvalidInput"),
294    when: "Inputs do not match a supported vec2word form.",
295    message: "vec2word received invalid input",
296};
297
298const ERROR_TRAIN_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
299    code: "RM.TRAINWORDEMBEDDING.INVALID_INPUT",
300    identifier: Some("RunMat:trainWordEmbedding:InvalidInput"),
301    when: "Inputs do not match a supported trainWordEmbedding form.",
302    message: "trainWordEmbedding received invalid input",
303};
304
305const ERROR_FASTTEXT_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
306    code: "RM.FASTTEXTWORDEMBEDDING.INVALID_INPUT",
307    identifier: Some("RunMat:fastTextWordEmbedding:InvalidInput"),
308    when: "Inputs do not match the supported fastTextWordEmbedding form.",
309    message: "fastTextWordEmbedding received invalid input",
310};
311
312const ERROR_DOC2SEQUENCE_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
313    code: "RM.DOC2SEQUENCE.INVALID_INPUT",
314    identifier: Some("RunMat:doc2sequence:InvalidInput"),
315    when: "Inputs do not match a supported doc2sequence form.",
316    message: "doc2sequence received invalid input",
317};
318
319const ERROR_READ_IO: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
320    code: "RM.READWORDEMBEDDING.IO",
321    identifier: Some("RunMat:readWordEmbedding:IOError"),
322    when: "The requested word embedding file cannot be read.",
323    message: "Unable to read word embedding file",
324};
325
326const ERROR_WRITE_IO: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
327    code: "RM.WRITEWORDEMBEDDING.IO",
328    identifier: Some("RunMat:writeWordEmbedding:IOError"),
329    when: "The requested word embedding file cannot be written.",
330    message: "Unable to write word embedding file",
331};
332
333const ERROR_WORD_EMBEDDING_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
334    code: "RM.WORDEMBEDDING.INVALID_INPUT",
335    identifier: Some("RunMat:wordEmbedding:InvalidInput"),
336    when: "Internal wordEmbedding object construction receives invalid data.",
337    message: "wordEmbedding received invalid input",
338};
339
340const FASTTEXT_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_FASTTEXT_INVALID_INPUT];
341const READ_ERRORS: [BuiltinErrorDescriptor; 2] = [ERROR_READ_INVALID_INPUT, ERROR_READ_IO];
342const WRITE_ERRORS: [BuiltinErrorDescriptor; 2] = [ERROR_WRITE_INVALID_INPUT, ERROR_WRITE_IO];
343const WORD2VEC_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_WORD2VEC_INVALID_INPUT];
344const VEC2WORD_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_VEC2WORD_INVALID_INPUT];
345const TRAIN_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_TRAIN_INVALID_INPUT];
346const DOC2SEQUENCE_ERRORS: [BuiltinErrorDescriptor; 1] = [ERROR_DOC2SEQUENCE_INVALID_INPUT];
347
348fn any_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
349    Type::Unknown
350}
351
352pub const READ_WORD_EMBEDDING_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
353    signatures: &[BuiltinSignatureDescriptor {
354        label: "emb = readWordEmbedding(filename)",
355        inputs: &IN_FILENAME,
356        outputs: &OUT_EMBEDDING,
357    }],
358    output_mode: BuiltinOutputMode::Fixed,
359    completion_policy: BuiltinCompletionPolicy::Public,
360    errors: &READ_ERRORS,
361};
362
363pub const WRITE_WORD_EMBEDDING_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
364    signatures: &[BuiltinSignatureDescriptor {
365        label: "writeWordEmbedding(emb, filename)",
366        inputs: &IN_EMBEDDING_FILENAME,
367        outputs: &OUT_NONE,
368    }],
369    output_mode: BuiltinOutputMode::Fixed,
370    completion_policy: BuiltinCompletionPolicy::Public,
371    errors: &WRITE_ERRORS,
372};
373
374pub const FASTTEXT_WORD_EMBEDDING_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
375    signatures: &[BuiltinSignatureDescriptor {
376        label: "emb = fastTextWordEmbedding",
377        inputs: &NO_INPUTS,
378        outputs: &OUT_EMBEDDING,
379    }],
380    output_mode: BuiltinOutputMode::Fixed,
381    completion_policy: BuiltinCompletionPolicy::Public,
382    errors: &FASTTEXT_ERRORS,
383};
384
385pub const WORD2VEC_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
386    signatures: &[
387        BuiltinSignatureDescriptor {
388            label: "M = word2vec(emb, words)",
389            inputs: &IN_WORDS,
390            outputs: &OUT_MATRIX,
391        },
392        BuiltinSignatureDescriptor {
393            label: "M = word2vec(emb, words, 'IgnoreCase', true)",
394            inputs: &IN_WORDS_REST,
395            outputs: &OUT_MATRIX,
396        },
397    ],
398    output_mode: BuiltinOutputMode::Fixed,
399    completion_policy: BuiltinCompletionPolicy::Public,
400    errors: &WORD2VEC_ERRORS,
401};
402
403pub const VEC2WORD_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
404    signatures: &[
405        BuiltinSignatureDescriptor {
406            label: "words = vec2word(emb, M)",
407            inputs: &IN_VECTORS,
408            outputs: &OUT_WORDS,
409        },
410        BuiltinSignatureDescriptor {
411            label: "[words, dist] = vec2word(emb, M, k, 'Distance', distance)",
412            inputs: &IN_VECTORS_REST,
413            outputs: &OUT_WORDS_DIST,
414        },
415    ],
416    output_mode: BuiltinOutputMode::ByRequestedOutputCount,
417    completion_policy: BuiltinCompletionPolicy::Public,
418    errors: &VEC2WORD_ERRORS,
419};
420
421pub const DOC2SEQUENCE_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
422    signatures: &[
423        BuiltinSignatureDescriptor {
424            label: "sequences = doc2sequence(emb, documents)",
425            inputs: &IN_MAP_DOCUMENTS,
426            outputs: &OUT_SEQUENCES,
427        },
428        BuiltinSignatureDescriptor {
429            label: "sequences = doc2sequence(enc, documents)",
430            inputs: &IN_MAP_DOCUMENTS,
431            outputs: &OUT_SEQUENCES,
432        },
433        BuiltinSignatureDescriptor {
434            label: "sequences = doc2sequence(emb, documents, Name, Value)",
435            inputs: &IN_MAP_DOCUMENTS_REST,
436            outputs: &OUT_SEQUENCES,
437        },
438        BuiltinSignatureDescriptor {
439            label: "sequences = doc2sequence(enc, documents, Name, Value)",
440            inputs: &IN_MAP_DOCUMENTS_REST,
441            outputs: &OUT_SEQUENCES,
442        },
443    ],
444    output_mode: BuiltinOutputMode::Fixed,
445    completion_policy: BuiltinCompletionPolicy::Public,
446    errors: &DOC2SEQUENCE_ERRORS,
447};
448
449pub const TRAIN_WORD_EMBEDDING_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
450    signatures: &[
451        BuiltinSignatureDescriptor {
452            label: "emb = trainWordEmbedding(filename)",
453            inputs: &IN_TRAIN_SOURCE,
454            outputs: &OUT_EMBEDDING,
455        },
456        BuiltinSignatureDescriptor {
457            label: "emb = trainWordEmbedding(documents)",
458            inputs: &IN_TRAIN_SOURCE,
459            outputs: &OUT_EMBEDDING,
460        },
461        BuiltinSignatureDescriptor {
462            label: "emb = trainWordEmbedding(___, Name, Value)",
463            inputs: &IN_TRAIN_SOURCE_REST,
464            outputs: &OUT_EMBEDDING,
465        },
466    ],
467    output_mode: BuiltinOutputMode::Fixed,
468    completion_policy: BuiltinCompletionPolicy::Public,
469    errors: &TRAIN_ERRORS,
470};
471
472#[runtime_builtin(
473    name = "fastTextWordEmbedding",
474    category = "strings/text_analytics",
475    summary = "Return a bundled fastText-style word embedding compatibility model.",
476    keywords = "fastTextWordEmbedding,wordEmbedding,text analytics,fastText,pretrained",
477    accel = "sink",
478    type_resolver(any_type),
479    descriptor(
480        crate::builtins::strings::text_analytics::embeddings::FASTTEXT_WORD_EMBEDDING_DESCRIPTOR
481    ),
482    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
483)]
484async fn fast_text_word_embedding_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
485    if !args.is_empty() {
486        return Err(embedding_error(
487            "fastTextWordEmbedding",
488            "fastTextWordEmbedding: expected no input arguments",
489        ));
490    }
491    embedding_object(compact_fast_text_embedding())
492}
493
494#[runtime_builtin(
495    name = "readWordEmbedding",
496    category = "strings/text_analytics",
497    summary = "Read word embedding models from UTF-8 text or zip files.",
498    keywords = "readWordEmbedding,wordEmbedding,text analytics,word2vec,GloVe",
499    accel = "sink",
500    type_resolver(any_type),
501    descriptor(
502        crate::builtins::strings::text_analytics::embeddings::READ_WORD_EMBEDDING_DESCRIPTOR
503    ),
504    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
505)]
506async fn read_word_embedding_builtin(filename: Value) -> BuiltinResult<Value> {
507    let filename = gather_if_needed_async(&filename)
508        .await
509        .map_err(|err| embedding_error("readWordEmbedding", err.to_string()))?;
510    let filename = embedding_filename_text(&filename, "readWordEmbedding")
511        .map_err(|err| embedding_error("readWordEmbedding", err.to_string()))?;
512    let path = Path::new(&filename);
513    let bytes = read_limited_file_bytes(path, "readWordEmbedding").await?;
514    let text = if is_zip_path(path) || looks_like_zip(&bytes) {
515        read_embedding_text_from_zip(&bytes)?
516    } else {
517        String::from_utf8(bytes).map_err(|err| {
518            embedding_error(
519                "readWordEmbedding",
520                format!("readWordEmbedding: embedding file must be UTF-8 text: {err}"),
521            )
522        })?
523    };
524    embedding_object(parse_embedding_text(&text, "readWordEmbedding")?)
525}
526
527#[runtime_builtin(
528    name = "writeWordEmbedding",
529    category = "strings/text_analytics",
530    summary = "Write a word embedding model as UTF-8 word2vec text.",
531    keywords = "writeWordEmbedding,wordEmbedding,text analytics,word2vec,write",
532    accel = "sink",
533    type_resolver(any_type),
534    descriptor(
535        crate::builtins::strings::text_analytics::embeddings::WRITE_WORD_EMBEDDING_DESCRIPTOR
536    ),
537    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
538)]
539async fn write_word_embedding_builtin(emb: Value, filename: Value) -> BuiltinResult<Value> {
540    let emb = gather_if_needed_async(&emb)
541        .await
542        .map_err(|err| embedding_error("writeWordEmbedding", err.to_string()))?;
543    let filename = gather_if_needed_async(&filename)
544        .await
545        .map_err(|err| embedding_error("writeWordEmbedding", err.to_string()))?;
546    let filename = embedding_filename_text(&filename, "writeWordEmbedding")
547        .map_err(|err| embedding_error("writeWordEmbedding", err.to_string()))?;
548    let object = match emb {
549        Value::Object(object) => object,
550        other => {
551            return Err(embedding_error(
552                "writeWordEmbedding",
553                format!("writeWordEmbedding: expected wordEmbedding object, got {other:?}"),
554            ));
555        }
556    };
557    let embedding = embedding_from_object(&object, "writeWordEmbedding")?;
558    let mut file = File::create_async(Path::new(&filename))
559        .await
560        .map_err(|err| {
561            embedding_error_with_source(
562                "writeWordEmbedding",
563                format!("writeWordEmbedding: unable to create '{filename}': {err}"),
564                err,
565            )
566        })?;
567    write_embedding_text(&embedding, &mut file, &filename)?;
568    file.flush().map_err(|err| {
569        embedding_error_with_source(
570            "writeWordEmbedding",
571            format!("writeWordEmbedding: unable to flush '{filename}': {err}"),
572            err,
573        )
574    })?;
575    Ok(Value::Num(0.0))
576}
577
578#[runtime_builtin(
579    name = "trainWordEmbedding",
580    category = "strings/text_analytics",
581    summary = "Train a local word embedding compatibility model.",
582    keywords = "trainWordEmbedding,wordEmbedding,text analytics,training",
583    accel = "sink",
584    type_resolver(any_type),
585    descriptor(
586        crate::builtins::strings::text_analytics::embeddings::TRAIN_WORD_EMBEDDING_DESCRIPTOR
587    ),
588    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
589)]
590async fn train_word_embedding_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
591    let gathered = gather_args(args, "trainWordEmbedding").await?;
592    let (source, options) = parse_train_word_embedding_args(gathered)?;
593    let documents = match source {
594        TrainSource::Documents(documents) => documents,
595        TrainSource::Filename(filename) => {
596            let bytes = read_limited_file_bytes(Path::new(&filename), "trainWordEmbedding").await?;
597            let text = String::from_utf8(bytes).map_err(|err| {
598                embedding_error(
599                    "trainWordEmbedding",
600                    format!("trainWordEmbedding: training file must be UTF-8 text: {err}"),
601                )
602            })?;
603            documents_from_training_text(&text)
604        }
605    };
606    embedding_object(train_embedding_model(documents, options)?)
607}
608
609#[runtime_builtin(
610    name = "doc2sequence",
611    category = "strings/text_analytics",
612    summary = "Convert tokenized documents to word-vector or word-index sequences.",
613    keywords = "doc2sequence,wordEmbedding,wordEncoding,tokenizedDocument,text analytics,sequences",
614    accel = "sink",
615    type_resolver(any_type),
616    descriptor(crate::builtins::strings::text_analytics::embeddings::DOC2SEQUENCE_DESCRIPTOR),
617    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
618)]
619async fn doc2sequence_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
620    let gathered = gather_args(args, "doc2sequence").await?;
621    let (sequence_object, document_object, options) = parse_doc2sequence_args(gathered)?;
622    let document_shape = document_shape_from_object(&document_object, "doc2sequence")?;
623    let documents = documents_from_object(&document_object, "doc2sequence")?;
624    if sequence_object.is_class(WORD_EMBEDDING_CLASS) {
625        let embedding = embedding_from_object(&sequence_object, "doc2sequence")?;
626        doc2sequence_value(&embedding, &documents, &document_shape, options)
627    } else if sequence_object.is_class(WORD_ENCODING_CLASS) {
628        let encoding = word_encoding_from_object(&sequence_object, "doc2sequence")?;
629        doc2sequence_indices_value(&encoding, &documents, &document_shape, options)
630    } else {
631        Err(embedding_error(
632            "doc2sequence",
633            format!(
634                "doc2sequence: expected wordEmbedding or wordEncoding object, got {}",
635                sequence_object.class_name
636            ),
637        ))
638    }
639}
640
641#[runtime_builtin(
642    name = "word2vec",
643    category = "strings/text_analytics",
644    summary = "Map words to rows of a word embedding matrix.",
645    keywords = "word2vec,wordEmbedding,text analytics,vectors",
646    accel = "sink",
647    type_resolver(any_type),
648    descriptor(crate::builtins::strings::text_analytics::embeddings::WORD2VEC_DESCRIPTOR),
649    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
650)]
651async fn word2vec_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
652    let gathered = gather_args(args, "word2vec").await?;
653    let (object, words, options) = parse_word2vec_args(gathered)?;
654    let embedding = embedding_from_object(&object, "word2vec")?;
655    let lookup = build_word_lookup(&embedding.vocabulary, options.ignore_case);
656    let word_count = words.len();
657    let mut rows = Vec::with_capacity(word_count);
658    for word in words {
659        let key = if options.ignore_case {
660            word.to_lowercase()
661        } else {
662            word
663        };
664        if let Some(&row) = lookup.get(&key) {
665            let start = row * embedding.dimension;
666            rows.push(embedding.vectors[start..start + embedding.dimension].to_vec());
667        } else {
668            rows.push(vec![f64::NAN; embedding.dimension]);
669        }
670    }
671    let mut out = Vec::with_capacity(word_count * embedding.dimension);
672    for col in 0..embedding.dimension {
673        for row in &rows {
674            out.push(row[col]);
675        }
676    }
677    Tensor::new(out, vec![word_count, embedding.dimension])
678        .map(Value::Tensor)
679        .map_err(|err| embedding_error("word2vec", err))
680}
681
682#[runtime_builtin(
683    name = "vec2word",
684    category = "strings/text_analytics",
685    summary = "Map embedding vectors to nearest vocabulary words.",
686    keywords = "vec2word,wordEmbedding,text analytics,nearest",
687    accel = "sink",
688    type_resolver(any_type),
689    descriptor(crate::builtins::strings::text_analytics::embeddings::VEC2WORD_DESCRIPTOR),
690    builtin_path = "crate::builtins::strings::text_analytics::embeddings"
691)]
692async fn vec2word_builtin(args: Vec<Value>) -> BuiltinResult<Value> {
693    let gathered = gather_args(args, "vec2word").await?;
694    let (object, matrix, options) = parse_vec2word_args(gathered)?;
695    let embedding = embedding_from_object(&object, "vec2word")?;
696    if matrix.cols != embedding.dimension {
697        return Err(embedding_error(
698            "vec2word",
699            format!(
700                "vec2word: input matrix must have {} columns, got {}",
701                embedding.dimension, matrix.cols
702            ),
703        ));
704    }
705    if options.k == 0 || options.k > embedding.vocabulary.len() {
706        return Err(embedding_error(
707            "vec2word",
708            format!(
709                "vec2word: k must be between 1 and vocabulary size ({})",
710                embedding.vocabulary.len()
711            ),
712        ));
713    }
714
715    let mut row_words = Vec::with_capacity(matrix.rows);
716    let mut row_distances = Vec::with_capacity(matrix.rows);
717    for row in 0..matrix.rows {
718        let query = row_slice(&matrix, row);
719        let mut scored = embedding
720            .vectors
721            .chunks(embedding.dimension)
722            .enumerate()
723            .map(|(idx, candidate)| {
724                (
725                    idx,
726                    match options.distance {
727                        DistanceMetric::Cosine => cosine_distance(&query, candidate),
728                        DistanceMetric::Euclidean => euclidean_distance(&query, candidate),
729                    },
730                )
731            })
732            .collect::<Vec<_>>();
733        scored.sort_by(|left, right| compare_scores(left.1, right.1).then(left.0.cmp(&right.0)));
734        let mut words = Vec::with_capacity(options.k);
735        let mut distances = Vec::with_capacity(options.k);
736        for (idx, distance) in scored.into_iter().take(options.k) {
737            words.push(embedding.vocabulary[idx].clone());
738            distances.push(distance);
739        }
740        row_words.push(words);
741        row_distances.push(distances);
742    }
743
744    let word_shape = if options.k == 1 {
745        vec![matrix.rows, 1]
746    } else {
747        vec![matrix.rows, options.k]
748    };
749    let mut words = Vec::with_capacity(matrix.rows * options.k);
750    let mut distances = Vec::with_capacity(matrix.rows * options.k);
751    for col in 0..options.k {
752        for row in 0..matrix.rows {
753            words.push(row_words[row][col].clone());
754            distances.push(row_distances[row][col]);
755        }
756    }
757    let words = Value::StringArray(
758        StringArray::new(words, word_shape).map_err(|err| embedding_error("vec2word", err))?,
759    );
760    let dist = Value::Tensor(
761        Tensor::new(distances, vec![matrix.rows, options.k])
762            .map_err(|err| embedding_error("vec2word", err))?,
763    );
764    Ok(Value::OutputList(vec![words, dist]))
765}
766
767async fn gather_args(args: Vec<Value>, fn_name: &str) -> BuiltinResult<Vec<Value>> {
768    let mut out = Vec::with_capacity(args.len());
769    for arg in args {
770        out.push(gather_if_needed_async(&arg).await.map_err(|err| {
771            embedding_error(fn_name, format!("{fn_name}: failed to gather input: {err}"))
772        })?);
773    }
774    Ok(out)
775}
776
777async fn read_limited_file_bytes(path: &Path, fn_name: &str) -> BuiltinResult<Vec<u8>> {
778    let file = File::open_async(path).await.map_err(|err| {
779        embedding_error_with_source(
780            fn_name,
781            format!("{fn_name}: unable to open '{}': {err}", path.display()),
782            err,
783        )
784    })?;
785    let mut limited = file.take(MAX_EMBEDDING_FILE_BYTES + 1);
786    let mut bytes = Vec::new();
787    limited.read_to_end(&mut bytes).map_err(|err| {
788        embedding_error_with_source(
789            fn_name,
790            format!("{fn_name}: unable to read '{}': {err}", path.display()),
791            err,
792        )
793    })?;
794    if bytes.len() as u64 > MAX_EMBEDDING_FILE_BYTES {
795        return Err(embedding_error(
796            fn_name,
797            format!(
798                "{fn_name}: embedding file exceeds maximum supported size of {MAX_EMBEDDING_FILE_BYTES} bytes"
799            ),
800        ));
801    }
802    Ok(bytes)
803}
804
805fn read_embedding_text_from_zip(bytes: &[u8]) -> BuiltinResult<String> {
806    let mut archive = zip::ZipArchive::new(Cursor::new(bytes)).map_err(|err| {
807        embedding_error(
808            "readWordEmbedding",
809            format!("readWordEmbedding: unable to read zip archive: {err}"),
810        )
811    })?;
812    if archive.len() > MAX_ZIP_ENTRIES {
813        return Err(embedding_error(
814            "readWordEmbedding",
815            format!("readWordEmbedding: zip archive contains more than {MAX_ZIP_ENTRIES} entries"),
816        ));
817    }
818
819    let mut selected = None;
820    for idx in 0..archive.len() {
821        let mut entry = archive.by_index(idx).map_err(|err| {
822            embedding_error(
823                "readWordEmbedding",
824                format!("readWordEmbedding: unable to read zip entry: {err}"),
825            )
826        })?;
827        if entry.is_dir() {
828            continue;
829        }
830        if entry.size() > MAX_EMBEDDING_FILE_BYTES {
831            return Err(embedding_error(
832                "readWordEmbedding",
833                format!(
834                    "readWordEmbedding: zip entry '{}' exceeds maximum supported size of {MAX_EMBEDDING_FILE_BYTES} bytes",
835                    entry.name()
836                ),
837            ));
838        }
839        let name = entry.name().to_ascii_lowercase();
840        let text_like = matches!(
841            Path::new(&name)
842                .extension()
843                .and_then(|ext| ext.to_str())
844                .map(|ext| ext.to_ascii_lowercase())
845                .as_deref(),
846            Some("txt") | Some("vec") | Some("glove") | Some("emb")
847        );
848        if !text_like && selected.is_some() {
849            continue;
850        }
851        let mut text = String::new();
852        entry.read_to_string(&mut text).map_err(|err| {
853            embedding_error(
854                "readWordEmbedding",
855                format!("readWordEmbedding: zip entry must contain UTF-8 text: {err}"),
856            )
857        })?;
858        selected = Some(text);
859        if text_like {
860            break;
861        }
862    }
863    selected.ok_or_else(|| {
864        embedding_error(
865            "readWordEmbedding",
866            "readWordEmbedding: zip archive does not contain an embedding text file",
867        )
868    })
869}
870
871fn parse_embedding_text(text: &str, fn_name: &str) -> BuiltinResult<EmbeddingModel> {
872    let mut lines = text.lines().enumerate().filter_map(|(idx, raw)| {
873        let trimmed = raw.trim();
874        (!trimmed.is_empty()).then_some((idx + 1, trimmed))
875    });
876    let Some((first_line_no, first_line)) = lines.next() else {
877        return Err(embedding_error(
878            fn_name,
879            format!("{fn_name}: embedding file is empty"),
880        ));
881    };
882
883    let first_parts = first_line.split_whitespace().collect::<Vec<_>>();
884    let (dimension_hint, expected_rows, pending_first) = if first_parts.len() == 2 {
885        match (
886            first_parts[0].parse::<usize>(),
887            first_parts[1].parse::<usize>(),
888        ) {
889            (Ok(rows), Ok(dim)) if rows > 0 && dim > 0 => (Some(dim), Some(rows), None),
890            _ => (None, None, Some((first_line_no, first_line))),
891        }
892    } else {
893        (None, None, Some((first_line_no, first_line)))
894    };
895
896    let mut vocabulary = Vec::new();
897    let mut vectors = Vec::new();
898    let mut positions = HashMap::new();
899    let mut dimension = dimension_hint;
900    let mut parsed_rows = 0usize;
901    let rows = pending_first.into_iter().chain(lines);
902    for (line_no, line) in rows {
903        let (word, vector) = parse_embedding_line(line, dimension, fn_name, line_no)?;
904        parsed_rows += 1;
905        let dim = vector.len();
906        if dim == 0 {
907            return Err(embedding_error(
908                fn_name,
909                format!("{fn_name}: line {line_no} has no vector values"),
910            ));
911        }
912        match dimension {
913            Some(expected) if expected != dim => {
914                return Err(embedding_error(
915                    fn_name,
916                    format!(
917                        "{fn_name}: line {line_no} has {dim} dimensions but expected {expected}"
918                    ),
919                ));
920            }
921            Some(_) => {}
922            None => dimension = Some(dim),
923        }
924        if let Some(old_pos) = positions.remove(&word) {
925            vocabulary.remove(old_pos);
926            let start = old_pos * dim;
927            vectors.drain(start..start + dim);
928            for pos in positions.values_mut() {
929                if *pos > old_pos {
930                    *pos -= 1;
931                }
932            }
933        }
934        positions.insert(word.clone(), vocabulary.len());
935        vocabulary.push(word);
936        vectors.extend(vector);
937    }
938
939    let dimension = dimension.ok_or_else(|| {
940        embedding_error(
941            fn_name,
942            format!("{fn_name}: embedding file contains no vectors"),
943        )
944    })?;
945    if vocabulary.is_empty() {
946        return Err(embedding_error(
947            fn_name,
948            format!("{fn_name}: embedding file contains no words"),
949        ));
950    }
951    if let Some(expected_rows) = expected_rows {
952        if expected_rows != parsed_rows {
953            return Err(embedding_error(
954                fn_name,
955                format!(
956                    "{fn_name}: header declares {expected_rows} words but parsed {parsed_rows} rows"
957                ),
958            ));
959        }
960    }
961    Ok(EmbeddingModel {
962        vocabulary,
963        vectors,
964        dimension,
965    })
966}
967
968fn parse_embedding_line(
969    line: &str,
970    dimension_hint: Option<usize>,
971    fn_name: &str,
972    line_no: usize,
973) -> BuiltinResult<(String, Vec<f64>)> {
974    let parts = line.split_whitespace().collect::<Vec<_>>();
975    if parts.len() < 2 {
976        return Err(embedding_error(
977            fn_name,
978            format!("{fn_name}: line {line_no} must contain a word and vector values"),
979        ));
980    }
981    let dimension = dimension_hint.unwrap_or(parts.len() - 1);
982    if parts.len() != dimension + 1 {
983        return Err(embedding_error(
984            fn_name,
985            format!(
986                "{fn_name}: line {line_no} has {} vector values but expected {dimension}",
987                parts.len().saturating_sub(1)
988            ),
989        ));
990    }
991    let word = parts[0].to_string();
992    if word.is_empty() {
993        return Err(embedding_error(
994            fn_name,
995            format!("{fn_name}: line {line_no} has an empty word"),
996        ));
997    }
998    let vector = parts[1..]
999        .iter()
1000        .map(|part| {
1001            let value = part.parse::<f64>().map_err(|err| {
1002                embedding_error(
1003                    fn_name,
1004                    format!("{fn_name}: invalid numeric value on line {line_no}: {err}"),
1005                )
1006            })?;
1007            if !value.is_finite() {
1008                return Err(embedding_error(
1009                    fn_name,
1010                    format!("{fn_name}: non-finite vector value on line {line_no}"),
1011                ));
1012            }
1013            Ok(value)
1014        })
1015        .collect::<BuiltinResult<Vec<_>>>()?;
1016    Ok((word, vector))
1017}
1018
1019fn write_embedding_text(
1020    model: &EmbeddingModel,
1021    writer: &mut impl Write,
1022    filename: &str,
1023) -> BuiltinResult<()> {
1024    if model.dimension == 0 || model.vocabulary.is_empty() {
1025        return Err(embedding_error(
1026            "writeWordEmbedding",
1027            "writeWordEmbedding: wordEmbedding object must contain at least one word and one dimension",
1028        ));
1029    }
1030    if model.vectors.len() != model.vocabulary.len() * model.dimension {
1031        return Err(embedding_error(
1032            "writeWordEmbedding",
1033            "writeWordEmbedding: wordEmbedding object has inconsistent vector storage",
1034        ));
1035    }
1036    writeln!(writer, "{} {}", model.vocabulary.len(), model.dimension).map_err(|err| {
1037        embedding_error_with_source(
1038            "writeWordEmbedding",
1039            format!("writeWordEmbedding: unable to write '{filename}': {err}"),
1040            err,
1041        )
1042    })?;
1043    for (row, word) in model.vocabulary.iter().enumerate() {
1044        if word.split_whitespace().count() != 1 {
1045            return Err(embedding_error(
1046                "writeWordEmbedding",
1047                format!(
1048                    "writeWordEmbedding: vocabulary word at index {} cannot contain whitespace",
1049                    row + 1
1050                ),
1051            ));
1052        }
1053        write!(writer, "{word}").map_err(|err| {
1054            embedding_error_with_source(
1055                "writeWordEmbedding",
1056                format!("writeWordEmbedding: unable to write '{filename}': {err}"),
1057                err,
1058            )
1059        })?;
1060        let start = row * model.dimension;
1061        for value in &model.vectors[start..start + model.dimension] {
1062            if !value.is_finite() {
1063                return Err(embedding_error(
1064                    "writeWordEmbedding",
1065                    format!(
1066                        "writeWordEmbedding: vector value for word '{}' must be finite",
1067                        word
1068                    ),
1069                ));
1070            }
1071            write!(writer, " {value}").map_err(|err| {
1072                embedding_error_with_source(
1073                    "writeWordEmbedding",
1074                    format!("writeWordEmbedding: unable to write '{filename}': {err}"),
1075                    err,
1076                )
1077            })?;
1078        }
1079        writeln!(writer).map_err(|err| {
1080            embedding_error_with_source(
1081                "writeWordEmbedding",
1082                format!("writeWordEmbedding: unable to write '{filename}': {err}"),
1083                err,
1084            )
1085        })?;
1086    }
1087    Ok(())
1088}
1089
1090fn embedding_filename_text(value: &Value, fn_name: &str) -> BuiltinResult<String> {
1091    match value {
1092        Value::Cell(cell) if cell.data.len() == 1 => match &cell.data[0] {
1093            Value::CharArray(array) if array.rows == 0 => Ok(String::new()),
1094            Value::CharArray(array) if array.rows == 1 => Ok(char_row_to_string(array)),
1095            other => Err(embedding_error(
1096                fn_name,
1097                format!("{fn_name}: 1-by-1 filename cell must contain a character vector, got {other:?}"),
1098            )),
1099        },
1100        Value::Cell(cell) => Err(embedding_error(
1101            fn_name,
1102            format!(
1103                "{fn_name}: filename cell array must be 1-by-1, got {} elements",
1104                cell.data.len()
1105            ),
1106        )),
1107        _ => scalar_text(value, fn_name),
1108    }
1109}
1110
1111#[derive(Clone, Debug)]
1112struct EmbeddingModel {
1113    vocabulary: Vec<String>,
1114    vectors: Vec<f64>,
1115    dimension: usize,
1116}
1117
1118fn embedding_object(model: EmbeddingModel) -> BuiltinResult<Value> {
1119    ensure_word_embedding_class_registered();
1120    let mut object = ObjectInstance::new(WORD_EMBEDDING_CLASS.to_string());
1121    object
1122        .properties
1123        .insert("Dimension".to_string(), Value::Num(model.dimension as f64));
1124    object.properties.insert(
1125        "Vocabulary".to_string(),
1126        Value::StringArray(
1127            StringArray::new(model.vocabulary.clone(), vec![1, model.vocabulary.len()])
1128                .map_err(|err| embedding_error("wordEmbedding", err))?,
1129        ),
1130    );
1131    object.properties.insert(
1132        VECTOR_PROPERTY.to_string(),
1133        Value::Tensor(
1134            Tensor::new(model.vectors, vec![model.vocabulary.len(), model.dimension])
1135                .map_err(|err| embedding_error("wordEmbedding", err))?,
1136        ),
1137    );
1138    Ok(Value::Object(object))
1139}
1140
1141fn embedding_from_object(object: &ObjectInstance, fn_name: &str) -> BuiltinResult<EmbeddingModel> {
1142    if !object.is_class(WORD_EMBEDDING_CLASS) {
1143        return Err(embedding_error(
1144            fn_name,
1145            format!(
1146                "{fn_name}: expected wordEmbedding object, got {}",
1147                object.class_name
1148            ),
1149        ));
1150    }
1151    let vocabulary = match object.properties.get("Vocabulary") {
1152        Some(Value::StringArray(array)) => array.data.clone(),
1153        other => {
1154            return Err(embedding_error(
1155                fn_name,
1156                format!(
1157                    "{fn_name}: wordEmbedding object has invalid Vocabulary property: {other:?}"
1158                ),
1159            ));
1160        }
1161    };
1162    let dimension = match object.properties.get("Dimension") {
1163        Some(Value::Num(value)) if value.is_finite() && *value >= 1.0 => *value as usize,
1164        other => {
1165            return Err(embedding_error(
1166                fn_name,
1167                format!(
1168                    "{fn_name}: wordEmbedding object has invalid Dimension property: {other:?}"
1169                ),
1170            ));
1171        }
1172    };
1173    let vectors = match object.properties.get(VECTOR_PROPERTY) {
1174        Some(Value::Tensor(tensor))
1175            if tensor.rows == vocabulary.len() && tensor.cols == dimension =>
1176        {
1177            tensor.data.clone()
1178        }
1179        other => {
1180            return Err(embedding_error(
1181                fn_name,
1182                format!("{fn_name}: wordEmbedding object has invalid vector storage: {other:?}"),
1183            ));
1184        }
1185    };
1186    Ok(EmbeddingModel {
1187        vocabulary,
1188        vectors,
1189        dimension,
1190    })
1191}
1192
1193pub(in crate::builtins::strings::text_analytics) fn word_embedding_vocabulary_from_object(
1194    object: &ObjectInstance,
1195    fn_name: &str,
1196) -> BuiltinResult<Vec<String>> {
1197    embedding_from_object(object, fn_name).map(|model| model.vocabulary)
1198}
1199
1200fn compact_fast_text_embedding() -> EmbeddingModel {
1201    let vocabulary = [
1202        "France",
1203        "Italy",
1204        "Rome",
1205        "Paris",
1206        "king",
1207        "queen",
1208        "man",
1209        "woman",
1210        "good",
1211        "bad",
1212        "excellent",
1213        "terrible",
1214        "data",
1215        "model",
1216        "analysis",
1217        "report",
1218        "signal",
1219        "image",
1220        "learning",
1221        "network",
1222        "algorithm",
1223        "matrix",
1224        "vector",
1225        "science",
1226        "engineering",
1227        "physics",
1228        "compute",
1229        "runtime",
1230        "test",
1231        "train",
1232        "document",
1233        "sequence",
1234    ]
1235    .into_iter()
1236    .map(str::to_string)
1237    .collect::<Vec<_>>();
1238    let dimension = 300usize;
1239    let mut vectors = Vec::with_capacity(vocabulary.len() * dimension);
1240    for word in &vocabulary {
1241        vectors.extend(compact_fast_text_vector(word, dimension));
1242    }
1243    EmbeddingModel {
1244        vocabulary,
1245        vectors,
1246        dimension,
1247    }
1248}
1249
1250fn compact_fast_text_vector(word: &str, dimension: usize) -> Vec<f64> {
1251    let mut vector = vec![0.0; dimension];
1252    let has_curated_vector = match word {
1253        "France" => {
1254            vector[0] = 1.0;
1255            vector[2] = 1.0;
1256            true
1257        }
1258        "Italy" => {
1259            vector[2] = 1.0;
1260            true
1261        }
1262        "Rome" => {
1263            vector[1] = 1.0;
1264            vector[2] = 1.0;
1265            true
1266        }
1267        "Paris" => {
1268            vector[0] = 1.0;
1269            vector[1] = 1.0;
1270            vector[2] = 1.0;
1271            true
1272        }
1273        "king" => {
1274            vector[3] = 1.0;
1275            vector[5] = 1.0;
1276            true
1277        }
1278        "queen" => {
1279            vector[4] = 1.0;
1280            vector[5] = 1.0;
1281            true
1282        }
1283        "man" => {
1284            vector[3] = 1.0;
1285            true
1286        }
1287        "woman" => {
1288            vector[4] = 1.0;
1289            true
1290        }
1291        "good" => {
1292            vector[6] = 1.0;
1293            true
1294        }
1295        "bad" => {
1296            vector[6] = -1.0;
1297            true
1298        }
1299        "excellent" => {
1300            vector[6] = 1.4;
1301            vector[7] = 0.4;
1302            true
1303        }
1304        "terrible" => {
1305            vector[6] = -1.4;
1306            vector[7] = -0.4;
1307            true
1308        }
1309        _ => false,
1310    };
1311    if !has_curated_vector {
1312        let seed = stable_word_hash(word);
1313        for (idx, slot) in vector.iter_mut().enumerate() {
1314            let bit = ((seed.rotate_left((idx % 31) as u32) ^ idx as u64) & 0x0f) as f64;
1315            *slot = (bit - 7.5) / 64.0;
1316        }
1317    }
1318    vector
1319}
1320
1321fn stable_word_hash(word: &str) -> u64 {
1322    let mut hash = 0xcbf29ce484222325u64;
1323    for byte in word.bytes() {
1324        hash ^= byte as u64;
1325        hash = hash.wrapping_mul(0x100000001b3);
1326    }
1327    hash
1328}
1329
1330fn ensure_word_embedding_class_registered() {
1331    WORD_EMBEDDING_CLASS_REGISTERED.with(|registered| {
1332        if registered.get() {
1333            return;
1334        }
1335        let mut properties = HashMap::new();
1336        for name in ["Dimension", "Vocabulary", VECTOR_PROPERTY] {
1337            properties.insert(name.to_string(), property_def(name));
1338        }
1339        runmat_builtins::register_class(ClassDef {
1340            name: WORD_EMBEDDING_CLASS.to_string(),
1341            parent: None,
1342            properties,
1343            methods: HashMap::new(),
1344        });
1345        registered.set(true);
1346    });
1347}
1348
1349fn property_def(name: &str) -> PropertyDef {
1350    PropertyDef {
1351        name: name.to_string(),
1352        is_static: false,
1353        is_constant: false,
1354        is_dependent: false,
1355        get_access: Access::Public,
1356        set_access: Access::Public,
1357        default_value: None,
1358    }
1359}
1360
1361enum TrainSource {
1362    Filename(String),
1363    Documents(Vec<Vec<String>>),
1364}
1365
1366#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1367enum TrainModelKind {
1368    SkipGram,
1369    Cbow,
1370}
1371
1372#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1373enum TrainLossFunction {
1374    NegativeSampling,
1375    HierarchicalSoftmax,
1376    Softmax,
1377}
1378
1379#[derive(Clone, Copy, Debug)]
1380struct TrainWordEmbeddingOptions {
1381    dimension: usize,
1382    window: usize,
1383    model: TrainModelKind,
1384    discard_factor: f64,
1385    loss_function: TrainLossFunction,
1386    num_negative_samples: usize,
1387    num_negative_samples_was_set: bool,
1388    num_epochs: usize,
1389    min_count: usize,
1390    ngram_range: (usize, usize),
1391    initial_learn_rate: f64,
1392    update_rate: usize,
1393    verbose: bool,
1394}
1395
1396impl Default for TrainWordEmbeddingOptions {
1397    fn default() -> Self {
1398        Self {
1399            dimension: 100,
1400            window: 5,
1401            model: TrainModelKind::SkipGram,
1402            discard_factor: 1.0e-4,
1403            loss_function: TrainLossFunction::NegativeSampling,
1404            num_negative_samples: 5,
1405            num_negative_samples_was_set: false,
1406            num_epochs: 5,
1407            min_count: 5,
1408            ngram_range: (3, 6),
1409            initial_learn_rate: 0.05,
1410            update_rate: 100,
1411            verbose: true,
1412        }
1413    }
1414}
1415
1416fn parse_train_word_embedding_args(
1417    args: Vec<Value>,
1418) -> BuiltinResult<(TrainSource, TrainWordEmbeddingOptions)> {
1419    if args.is_empty() {
1420        return Err(embedding_error(
1421            "trainWordEmbedding",
1422            "trainWordEmbedding: expected filename or tokenizedDocument input",
1423        ));
1424    }
1425    if !(args.len() - 1).is_multiple_of(2) {
1426        return Err(embedding_error(
1427            "trainWordEmbedding",
1428            "trainWordEmbedding: name-value options must appear in pairs",
1429        ));
1430    }
1431    let source = train_source_from_value(&args[0])?;
1432    let mut options = TrainWordEmbeddingOptions::default();
1433    let mut idx = 1usize;
1434    while idx < args.len() {
1435        let name = scalar_text(&args[idx], "trainWordEmbedding")
1436            .map_err(|err| embedding_error("trainWordEmbedding", err.to_string()))?
1437            .to_ascii_lowercase();
1438        match name.as_str() {
1439            "dimension" => {
1440                options.dimension = parse_positive_integer(&args[idx + 1], "trainWordEmbedding")?
1441            }
1442            "window" => {
1443                options.window =
1444                    parse_nonnegative_integer(&args[idx + 1], "trainWordEmbedding", "Window")?
1445            }
1446            "model" => {
1447                let value = scalar_text(&args[idx + 1], "trainWordEmbedding")
1448                    .map_err(|err| embedding_error("trainWordEmbedding", err.to_string()))?
1449                    .to_ascii_lowercase();
1450                options.model = match value.as_str() {
1451                    "skipgram" => TrainModelKind::SkipGram,
1452                    "cbow" => TrainModelKind::Cbow,
1453                    other => {
1454                        return Err(embedding_error(
1455                            "trainWordEmbedding",
1456                            format!(
1457                                "trainWordEmbedding: Model must be 'skipgram' or 'cbow', got '{other}'"
1458                            ),
1459                        ));
1460                    }
1461                };
1462            }
1463            "discardfactor" => {
1464                options.discard_factor =
1465                    parse_positive_scalar(&args[idx + 1], "trainWordEmbedding", "DiscardFactor")?
1466            }
1467            "lossfunction" => {
1468                let value = scalar_text(&args[idx + 1], "trainWordEmbedding")
1469                    .map_err(|err| embedding_error("trainWordEmbedding", err.to_string()))?
1470                    .to_ascii_lowercase();
1471                options.loss_function = match value.as_str() {
1472                    "ns" => TrainLossFunction::NegativeSampling,
1473                    "hs" => TrainLossFunction::HierarchicalSoftmax,
1474                    "softmax" => TrainLossFunction::Softmax,
1475                    other => {
1476                        return Err(embedding_error(
1477                            "trainWordEmbedding",
1478                            format!(
1479                                "trainWordEmbedding: LossFunction must be 'ns', 'hs', or 'softmax', got '{other}'"
1480                            ),
1481                        ));
1482                    }
1483                };
1484            }
1485            "numnegativesamples" => {
1486                options.num_negative_samples =
1487                    parse_positive_integer(&args[idx + 1], "trainWordEmbedding")?;
1488                options.num_negative_samples_was_set = true;
1489            }
1490            "numepochs" => {
1491                options.num_epochs = parse_positive_integer(&args[idx + 1], "trainWordEmbedding")?
1492            }
1493            "mincount" => {
1494                options.min_count = parse_positive_integer(&args[idx + 1], "trainWordEmbedding")?
1495            }
1496            "ngramrange" => options.ngram_range = parse_ngram_range(&args[idx + 1])?,
1497            "initiallearnrate" => {
1498                options.initial_learn_rate =
1499                    parse_positive_scalar(&args[idx + 1], "trainWordEmbedding", "InitialLearnRate")?
1500            }
1501            "updaterate" => {
1502                options.update_rate = parse_positive_integer(&args[idx + 1], "trainWordEmbedding")?
1503            }
1504            "verbose" => options.verbose = parse_bool_scalar(&args[idx + 1], "trainWordEmbedding")?,
1505            other => {
1506                return Err(embedding_error(
1507                    "trainWordEmbedding",
1508                    format!("trainWordEmbedding: unsupported option '{other}'"),
1509                ));
1510            }
1511        }
1512        idx += 2;
1513    }
1514    if options.num_negative_samples_was_set
1515        && options.loss_function != TrainLossFunction::NegativeSampling
1516    {
1517        return Err(embedding_error(
1518            "trainWordEmbedding",
1519            "trainWordEmbedding: NumNegativeSamples is only valid when LossFunction is 'ns'",
1520        ));
1521    }
1522    checked_train_dense_size(options.dimension, 1, "trainWordEmbedding")?;
1523    Ok((source, options))
1524}
1525
1526fn train_source_from_value(value: &Value) -> BuiltinResult<TrainSource> {
1527    match value {
1528        Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => Ok(
1529            TrainSource::Documents(documents_from_object(object, "trainWordEmbedding")?),
1530        ),
1531        Value::String(_) | Value::StringArray(_) | Value::CharArray(_) | Value::Cell(_) => {
1532            Ok(TrainSource::Filename(train_filename_from_value(value)?))
1533        }
1534        other => Err(embedding_error(
1535            "trainWordEmbedding",
1536            format!("trainWordEmbedding: expected filename or tokenizedDocument, got {other:?}"),
1537        )),
1538    }
1539}
1540
1541fn train_filename_from_value(value: &Value) -> BuiltinResult<String> {
1542    match value {
1543        Value::Cell(cell) if cell.data.len() == 1 => train_filename_from_value(&cell.data[0]),
1544        other => {
1545            let filename = scalar_text(other, "trainWordEmbedding")
1546                .map_err(|err| embedding_error("trainWordEmbedding", err.to_string()))?;
1547            if filename.trim().is_empty() {
1548                Err(embedding_error(
1549                    "trainWordEmbedding",
1550                    "trainWordEmbedding: filename must not be empty",
1551                ))
1552            } else {
1553                Ok(filename)
1554            }
1555        }
1556    }
1557}
1558
1559fn documents_from_training_text(text: &str) -> Vec<Vec<String>> {
1560    text.lines()
1561        .map(|line| {
1562            line.split_whitespace()
1563                .filter(|word| !word.is_empty())
1564                .map(str::to_string)
1565                .collect::<Vec<_>>()
1566        })
1567        .filter(|doc| !doc.is_empty())
1568        .collect()
1569}
1570
1571fn train_embedding_model(
1572    documents: Vec<Vec<String>>,
1573    options: TrainWordEmbeddingOptions,
1574) -> BuiltinResult<EmbeddingModel> {
1575    if documents.is_empty() || documents.iter().all(Vec::is_empty) {
1576        return Err(embedding_error(
1577            "trainWordEmbedding",
1578            "trainWordEmbedding: training data contains no tokens",
1579        ));
1580    }
1581
1582    let mut counts = HashMap::<String, (usize, usize)>::new();
1583    let mut next_pos = 0usize;
1584    for token in documents.iter().flatten() {
1585        let entry = counts.entry(token.clone()).or_insert_with(|| {
1586            let pos = next_pos;
1587            next_pos += 1;
1588            (0, pos)
1589        });
1590        entry.0 += 1;
1591    }
1592
1593    let mut vocabulary = counts
1594        .iter()
1595        .filter(|(_, (count, _))| *count >= options.min_count)
1596        .map(|(word, (count, first_pos))| (word.clone(), *count, *first_pos))
1597        .collect::<Vec<_>>();
1598    if vocabulary.is_empty() {
1599        return Err(embedding_error(
1600            "trainWordEmbedding",
1601            format!(
1602                "trainWordEmbedding: no vocabulary words meet MinCount {}",
1603                options.min_count
1604            ),
1605        ));
1606    }
1607    vocabulary.sort_by(|left, right| right.1.cmp(&left.1).then(left.2.cmp(&right.2)));
1608    checked_train_dense_size(options.dimension, vocabulary.len(), "trainWordEmbedding")?;
1609
1610    let mut positions = HashMap::new();
1611    let mut final_vocabulary = Vec::with_capacity(vocabulary.len());
1612    for (idx, (word, _, _)) in vocabulary.into_iter().enumerate() {
1613        positions.insert(word.clone(), idx);
1614        final_vocabulary.push(word);
1615    }
1616
1617    let mut rows = vec![vec![0.0; options.dimension]; final_vocabulary.len()];
1618    for (idx, word) in final_vocabulary.iter().enumerate() {
1619        add_lexical_features(&mut rows[idx], word, options);
1620    }
1621
1622    let base = options.initial_learn_rate
1623        * options.num_epochs as f64
1624        * match options.loss_function {
1625            TrainLossFunction::NegativeSampling => {
1626                1.0 + (options.num_negative_samples as f64).ln_1p() * 0.05
1627            }
1628            TrainLossFunction::HierarchicalSoftmax => 0.95,
1629            TrainLossFunction::Softmax => 1.05,
1630        };
1631    let model_scale = match options.model {
1632        TrainModelKind::SkipGram => 1.0,
1633        TrainModelKind::Cbow => 0.75,
1634    };
1635    let discard_scale = (1.0 + options.discard_factor.log10().abs()).recip();
1636    let update_scale = 1.0 + (options.update_rate as f64).ln_1p() * 0.01;
1637
1638    for document in &documents {
1639        for (target_pos, target) in document.iter().enumerate() {
1640            let Some(&target_idx) = positions.get(target) else {
1641                continue;
1642            };
1643            if options.window == 0 {
1644                continue;
1645            }
1646            let start = target_pos.saturating_sub(options.window);
1647            let end = target_pos
1648                .saturating_add(options.window)
1649                .saturating_add(1)
1650                .min(document.len());
1651            for (ctx_pos, context) in document.iter().enumerate().take(end).skip(start) {
1652                if ctx_pos == target_pos {
1653                    continue;
1654                }
1655                let Some(&context_idx) = positions.get(context) else {
1656                    continue;
1657                };
1658                let distance = target_pos.abs_diff(ctx_pos).max(1) as f64;
1659                let weight = base * model_scale * discard_scale * update_scale / distance;
1660                add_hashed_feature(
1661                    &mut rows[target_idx],
1662                    context,
1663                    weight,
1664                    0x9e37_79b9_7f4a_7c15,
1665                );
1666                if options.model == TrainModelKind::SkipGram {
1667                    add_hashed_feature(
1668                        &mut rows[context_idx],
1669                        target,
1670                        weight * 0.5,
1671                        0xc2b2_ae3d_27d4_eb4f,
1672                    );
1673                }
1674            }
1675        }
1676    }
1677
1678    let mut vectors = Vec::with_capacity(final_vocabulary.len() * options.dimension);
1679    for row in &mut rows {
1680        normalize_vector(row);
1681        vectors.extend(row.iter().copied());
1682    }
1683    Ok(EmbeddingModel {
1684        vocabulary: final_vocabulary,
1685        vectors,
1686        dimension: options.dimension,
1687    })
1688}
1689
1690fn add_lexical_features(row: &mut [f64], word: &str, options: TrainWordEmbeddingOptions) {
1691    add_hashed_feature(row, word, 1.0, 0xcbf2_9ce4_8422_2325);
1692    if options.ngram_range != (0, 0) {
1693        add_character_ngram_features(row, word, options.ngram_range);
1694    }
1695    add_hashed_feature(row, &word.to_ascii_lowercase(), 0.2, 0x517c_c1b7_2722_0a95);
1696}
1697
1698fn add_character_ngram_features(row: &mut [f64], word: &str, range: (usize, usize)) {
1699    let chars = format!("<{word}>").chars().collect::<Vec<_>>();
1700    let max_len = range.1.min(chars.len());
1701    for len in range.0..=max_len {
1702        if len == 0 || len > chars.len() {
1703            continue;
1704        }
1705        for window in chars.windows(len) {
1706            let ngram = window.iter().collect::<String>();
1707            add_hashed_feature(row, &ngram, 0.35, 0x1000_0000_01b3);
1708        }
1709    }
1710}
1711
1712fn add_hashed_feature(row: &mut [f64], key: &str, weight: f64, salt: u64) {
1713    if row.is_empty() {
1714        return;
1715    }
1716    let hash = fnv1a64_with_salt(key, salt);
1717    let idx = (hash as usize) % row.len();
1718    let sign = if (hash >> 63) == 0 { 1.0 } else { -1.0 };
1719    row[idx] += sign * weight;
1720}
1721
1722fn fnv1a64_with_salt(value: &str, salt: u64) -> u64 {
1723    let mut hash = 0xcbf2_9ce4_8422_2325u64 ^ salt;
1724    for byte in value.as_bytes() {
1725        hash ^= u64::from(*byte);
1726        hash = hash.wrapping_mul(0x1000_0000_01b3);
1727    }
1728    hash
1729}
1730
1731fn normalize_vector(row: &mut [f64]) {
1732    let norm = row.iter().map(|value| value * value).sum::<f64>().sqrt();
1733    if norm > 0.0 {
1734        for value in row {
1735            *value /= norm;
1736        }
1737    }
1738}
1739
1740fn parse_nonnegative_integer(value: &Value, fn_name: &str, option: &str) -> BuiltinResult<usize> {
1741    let n = numeric_scalar(value, fn_name, option)?;
1742    if !n.is_finite() || n < 0.0 || n.fract() != 0.0 {
1743        return Err(embedding_error(
1744            fn_name,
1745            format!("{fn_name}: {option} must be a nonnegative integer, got {n}"),
1746        ));
1747    }
1748    Ok(n as usize)
1749}
1750
1751fn parse_positive_scalar(value: &Value, fn_name: &str, option: &str) -> BuiltinResult<f64> {
1752    let n = numeric_scalar(value, fn_name, option)?;
1753    if !n.is_finite() || n <= 0.0 {
1754        return Err(embedding_error(
1755            fn_name,
1756            format!("{fn_name}: {option} must be a positive scalar, got {n}"),
1757        ));
1758    }
1759    Ok(n)
1760}
1761
1762fn parse_ngram_range(value: &Value) -> BuiltinResult<(usize, usize)> {
1763    let values = match value {
1764        Value::Tensor(tensor) if tensor.data.len() == 2 => tensor.data.clone(),
1765        other => {
1766            return Err(embedding_error(
1767                "trainWordEmbedding",
1768                format!("trainWordEmbedding: NGramRange must be a two-element numeric vector, got {other:?}"),
1769            ));
1770        }
1771    };
1772    let min = values[0];
1773    let max = values[1];
1774    if !min.is_finite()
1775        || !max.is_finite()
1776        || min < 0.0
1777        || max < 0.0
1778        || min.fract() != 0.0
1779        || max.fract() != 0.0
1780        || min > max
1781    {
1782        return Err(embedding_error(
1783            "trainWordEmbedding",
1784            format!("trainWordEmbedding: NGramRange must be [min max] nonnegative integers with min <= max, got [{min} {max}]"),
1785        ));
1786    }
1787    Ok((min as usize, max as usize))
1788}
1789
1790fn numeric_scalar(value: &Value, fn_name: &str, option: &str) -> BuiltinResult<f64> {
1791    match value {
1792        Value::Num(value) => Ok(*value),
1793        Value::Tensor(tensor) if tensor.data.len() == 1 => Ok(tensor.data[0]),
1794        other => Err(embedding_error(
1795            fn_name,
1796            format!("{fn_name}: {option} must be a numeric scalar, got {other:?}"),
1797        )),
1798    }
1799}
1800
1801fn checked_train_dense_size(
1802    dimension: usize,
1803    vocabulary_len: usize,
1804    fn_name: &str,
1805) -> BuiltinResult<()> {
1806    let cells = dimension.checked_mul(vocabulary_len).ok_or_else(|| {
1807        embedding_error(
1808            fn_name,
1809            format!("{fn_name}: trained embedding dimensions overflow dense storage"),
1810        )
1811    })?;
1812    if cells > MAX_TRAINED_DENSE_VALUES {
1813        return Err(embedding_error(
1814            fn_name,
1815            format!(
1816                "{fn_name}: trained embedding would require {cells} dense values; limit is {MAX_TRAINED_DENSE_VALUES}"
1817            ),
1818        ));
1819    }
1820    Ok(())
1821}
1822
1823#[derive(Clone, Copy, Debug, Default)]
1824struct Word2VecOptions {
1825    ignore_case: bool,
1826}
1827
1828fn parse_word2vec_args(
1829    args: Vec<Value>,
1830) -> BuiltinResult<(ObjectInstance, Vec<String>, Word2VecOptions)> {
1831    if args.len() < 2 {
1832        return Err(embedding_error(
1833            "word2vec",
1834            "word2vec: expected word2vec(emb, words)",
1835        ));
1836    }
1837    let mut iter = args.into_iter();
1838    let object = match iter.next().expect("checked") {
1839        Value::Object(object) => object,
1840        other => {
1841            return Err(embedding_error(
1842                "word2vec",
1843                format!("word2vec: expected wordEmbedding object, got {other:?}"),
1844            ));
1845        }
1846    };
1847    let words_value = iter.next().expect("checked");
1848    let words = words_from_value(&words_value, "word2vec")?;
1849    let mut options = Word2VecOptions::default();
1850    let rest = iter.collect::<Vec<_>>();
1851    let mut idx = 0;
1852    while idx < rest.len() {
1853        if idx + 1 >= rest.len() {
1854            return Err(embedding_error(
1855                "word2vec",
1856                "word2vec: name-value options must be paired",
1857            ));
1858        }
1859        let name = scalar_text(&rest[idx], "word2vec")
1860            .map_err(|err| embedding_error("word2vec", err.to_string()))?
1861            .to_ascii_lowercase();
1862        match name.as_str() {
1863            "ignorecase" => options.ignore_case = parse_bool_scalar(&rest[idx + 1], "word2vec")?,
1864            other => {
1865                return Err(embedding_error(
1866                    "word2vec",
1867                    format!("word2vec: unsupported option '{other}'"),
1868                ));
1869            }
1870        }
1871        idx += 2;
1872    }
1873    Ok((object, words, options))
1874}
1875
1876#[derive(Clone, Copy, Debug)]
1877enum DistanceMetric {
1878    Cosine,
1879    Euclidean,
1880}
1881
1882#[derive(Clone, Copy, Debug)]
1883struct Vec2WordOptions {
1884    k: usize,
1885    distance: DistanceMetric,
1886}
1887
1888impl Default for Vec2WordOptions {
1889    fn default() -> Self {
1890        Self {
1891            k: 1,
1892            distance: DistanceMetric::Cosine,
1893        }
1894    }
1895}
1896
1897fn parse_vec2word_args(
1898    args: Vec<Value>,
1899) -> BuiltinResult<(ObjectInstance, Tensor, Vec2WordOptions)> {
1900    if args.len() < 2 {
1901        return Err(embedding_error(
1902            "vec2word",
1903            "vec2word: expected vec2word(emb, M)",
1904        ));
1905    }
1906    let mut iter = args.into_iter();
1907    let object = match iter.next().expect("checked") {
1908        Value::Object(object) => object,
1909        other => {
1910            return Err(embedding_error(
1911                "vec2word",
1912                format!("vec2word: expected wordEmbedding object, got {other:?}"),
1913            ));
1914        }
1915    };
1916    let matrix = match iter.next().expect("checked") {
1917        Value::Tensor(tensor) => tensor,
1918        Value::Num(value) => {
1919            Tensor::new(vec![value], vec![1, 1]).map_err(|err| embedding_error("vec2word", err))?
1920        }
1921        other => {
1922            return Err(embedding_error(
1923                "vec2word",
1924                format!("vec2word: expected numeric matrix, got {other:?}"),
1925            ));
1926        }
1927    };
1928    let mut rest = iter.collect::<Vec<_>>();
1929    let mut options = Vec2WordOptions::default();
1930    if rest.first().is_some_and(is_numeric_scalar) {
1931        options.k = parse_positive_integer(&rest.remove(0), "vec2word")?;
1932    }
1933    let mut idx = 0;
1934    while idx < rest.len() {
1935        if idx + 1 >= rest.len() {
1936            return Err(embedding_error(
1937                "vec2word",
1938                "vec2word: name-value options must be paired",
1939            ));
1940        }
1941        let name = scalar_text(&rest[idx], "vec2word")
1942            .map_err(|err| embedding_error("vec2word", err.to_string()))?
1943            .to_ascii_lowercase();
1944        match name.as_str() {
1945            "distance" => {
1946                let metric = scalar_text(&rest[idx + 1], "vec2word")
1947                    .map_err(|err| embedding_error("vec2word", err.to_string()))?
1948                    .to_ascii_lowercase();
1949                options.distance = match metric.as_str() {
1950                    "cosine" => DistanceMetric::Cosine,
1951                    "euclidean" => DistanceMetric::Euclidean,
1952                    other => {
1953                        return Err(embedding_error(
1954                            "vec2word",
1955                            format!("vec2word: unsupported Distance '{other}'"),
1956                        ));
1957                    }
1958                };
1959            }
1960            other => {
1961                return Err(embedding_error(
1962                    "vec2word",
1963                    format!("vec2word: unsupported option '{other}'"),
1964                ));
1965            }
1966        }
1967        idx += 2;
1968    }
1969    Ok((object, matrix, options))
1970}
1971
1972#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1973enum UnknownWordMode {
1974    Discard,
1975    Nan,
1976}
1977
1978#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1979enum PaddingDirection {
1980    Left,
1981    Right,
1982    None,
1983}
1984
1985#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1986enum SequenceLength {
1987    Longest,
1988    Shortest,
1989    Fixed(usize),
1990}
1991
1992#[derive(Clone, Copy, Debug)]
1993struct Doc2SequenceOptions {
1994    unknown_word: UnknownWordMode,
1995    padding_direction: PaddingDirection,
1996    padding_value: f64,
1997    length: SequenceLength,
1998}
1999
2000impl Default for Doc2SequenceOptions {
2001    fn default() -> Self {
2002        Self {
2003            unknown_word: UnknownWordMode::Discard,
2004            padding_direction: PaddingDirection::Left,
2005            padding_value: 0.0,
2006            length: SequenceLength::Longest,
2007        }
2008    }
2009}
2010
2011fn parse_doc2sequence_args(
2012    args: Vec<Value>,
2013) -> BuiltinResult<(ObjectInstance, ObjectInstance, Doc2SequenceOptions)> {
2014    if args.len() < 2 {
2015        return Err(embedding_error(
2016            "doc2sequence",
2017            "doc2sequence: expected doc2sequence(embOrEnc, documents)",
2018        ));
2019    }
2020    if !(args.len() - 2).is_multiple_of(2) {
2021        return Err(embedding_error(
2022            "doc2sequence",
2023            "doc2sequence: name-value options must be paired",
2024        ));
2025    }
2026    let sequence_model = match &args[0] {
2027        Value::Object(object) => object.clone(),
2028        other => {
2029            return Err(embedding_error(
2030                "doc2sequence",
2031                format!(
2032                    "doc2sequence: expected wordEmbedding or wordEncoding object, got {other:?}"
2033                ),
2034            ));
2035        }
2036    };
2037    let documents = match &args[1] {
2038        Value::Object(object) if object.is_class(TOKENIZED_DOCUMENT_CLASS) => object.clone(),
2039        Value::Object(object) => {
2040            return Err(embedding_error(
2041                "doc2sequence",
2042                format!(
2043                    "doc2sequence: expected tokenizedDocument object, got {}",
2044                    object.class_name
2045                ),
2046            ));
2047        }
2048        other => {
2049            return Err(embedding_error(
2050                "doc2sequence",
2051                format!("doc2sequence: expected tokenizedDocument object, got {other:?}"),
2052            ));
2053        }
2054    };
2055    let mut options = Doc2SequenceOptions::default();
2056    let mut idx = 2usize;
2057    while idx < args.len() {
2058        let name = scalar_text(&args[idx], "doc2sequence")
2059            .map_err(|err| embedding_error("doc2sequence", err.to_string()))?
2060            .to_ascii_lowercase();
2061        match name.as_str() {
2062            "unknownword" => {
2063                let value = scalar_text(&args[idx + 1], "doc2sequence")
2064                    .map_err(|err| embedding_error("doc2sequence", err.to_string()))?
2065                    .to_ascii_lowercase();
2066                options.unknown_word = match value.as_str() {
2067                    "discard" => UnknownWordMode::Discard,
2068                    "nan" => UnknownWordMode::Nan,
2069                    other => {
2070                        return Err(embedding_error(
2071                            "doc2sequence",
2072                            format!("doc2sequence: UnknownWord must be 'discard' or 'nan', got '{other}'"),
2073                        ));
2074                    }
2075                };
2076            }
2077            "paddingdirection" => {
2078                let value = scalar_text(&args[idx + 1], "doc2sequence")
2079                    .map_err(|err| embedding_error("doc2sequence", err.to_string()))?
2080                    .to_ascii_lowercase();
2081                options.padding_direction = match value.as_str() {
2082                    "left" => PaddingDirection::Left,
2083                    "right" => PaddingDirection::Right,
2084                    "none" => PaddingDirection::None,
2085                    other => {
2086                        return Err(embedding_error(
2087                            "doc2sequence",
2088                            format!("doc2sequence: PaddingDirection must be 'left', 'right', or 'none', got '{other}'"),
2089                        ));
2090                    }
2091                };
2092            }
2093            "paddingvalue" => {
2094                options.padding_value =
2095                    parse_numeric_scalar(&args[idx + 1], "doc2sequence", "PaddingValue")?;
2096            }
2097            "length" => options.length = parse_sequence_length(&args[idx + 1])?,
2098            other => {
2099                return Err(embedding_error(
2100                    "doc2sequence",
2101                    format!("doc2sequence: unsupported option '{other}'"),
2102                ));
2103            }
2104        }
2105        idx += 2;
2106    }
2107    Ok((sequence_model, documents, options))
2108}
2109
2110#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2111enum SequenceToken {
2112    Known(usize),
2113    UnknownNan,
2114}
2115
2116fn doc2sequence_value(
2117    embedding: &EmbeddingModel,
2118    documents: &[Vec<String>],
2119    document_shape: &[usize],
2120    options: Doc2SequenceOptions,
2121) -> BuiltinResult<Value> {
2122    let lookup = build_word_lookup(&embedding.vocabulary, false);
2123    let mut sequences = Vec::with_capacity(documents.len());
2124    for document in documents {
2125        let mut sequence = Vec::new();
2126        for token in document {
2127            if let Some(&idx) = lookup.get(token) {
2128                sequence.push(SequenceToken::Known(idx));
2129            } else if options.unknown_word == UnknownWordMode::Nan {
2130                sequence.push(SequenceToken::UnknownNan);
2131            }
2132        }
2133        sequences.push(sequence);
2134    }
2135    let resolved_length = resolve_sequence_length(&sequences, options.length);
2136    let mut values = Vec::with_capacity(sequences.len());
2137    let mut total_cells = 0usize;
2138    for sequence in &sequences {
2139        let target_len = sequence_target_len(sequence.len(), resolved_length, options);
2140        total_cells = total_cells
2141            .checked_add(
2142                embedding
2143                    .dimension
2144                    .checked_mul(target_len)
2145                    .ok_or_else(|| dense_doc2sequence_limit_error("doc2sequence"))?,
2146            )
2147            .ok_or_else(|| dense_doc2sequence_limit_error("doc2sequence"))?;
2148        if total_cells > MAX_DOC2SEQUENCE_DENSE_VALUES {
2149            return Err(dense_doc2sequence_limit_error("doc2sequence"));
2150        }
2151        values.push(Value::Tensor(sequence_tensor(
2152            embedding,
2153            sequence,
2154            target_len,
2155            options.padding_direction,
2156            options.padding_value,
2157        )?));
2158    }
2159    Ok(Value::Cell(
2160        CellArray::new_with_shape(values, document_shape.to_vec())
2161            .map_err(|err| embedding_error("doc2sequence", err))?,
2162    ))
2163}
2164
2165fn doc2sequence_indices_value(
2166    encoding: &WordEncodingModel,
2167    documents: &[Vec<String>],
2168    document_shape: &[usize],
2169    options: Doc2SequenceOptions,
2170) -> BuiltinResult<Value> {
2171    let lookup = build_word_lookup(&encoding.vocabulary, false);
2172    let mut sequences = Vec::with_capacity(documents.len());
2173    for document in documents {
2174        let mut sequence = Vec::new();
2175        for token in document {
2176            if let Some(&idx) = lookup.get(token) {
2177                sequence.push(IndexSequenceToken::Known((idx + 1) as f64));
2178            } else if options.unknown_word == UnknownWordMode::Nan {
2179                sequence.push(IndexSequenceToken::UnknownNan);
2180            }
2181        }
2182        sequences.push(sequence);
2183    }
2184    let resolved_length = resolve_sequence_length(&sequences, options.length);
2185    let mut values = Vec::with_capacity(sequences.len());
2186    let mut total_cells = 0usize;
2187    for sequence in &sequences {
2188        let target_len = sequence_target_len(sequence.len(), resolved_length, options);
2189        total_cells = total_cells
2190            .checked_add(target_len)
2191            .ok_or_else(|| dense_doc2sequence_limit_error("doc2sequence"))?;
2192        if total_cells > MAX_DOC2SEQUENCE_DENSE_VALUES {
2193            return Err(dense_doc2sequence_limit_error("doc2sequence"));
2194        }
2195        values.push(Value::Tensor(index_sequence_tensor(
2196            sequence,
2197            target_len,
2198            options.padding_direction,
2199            options.padding_value,
2200        )?));
2201    }
2202    Ok(Value::Cell(
2203        CellArray::new_with_shape(values, document_shape.to_vec())
2204            .map_err(|err| embedding_error("doc2sequence", err))?,
2205    ))
2206}
2207
2208fn resolve_sequence_length<T>(sequences: &[Vec<T>], length: SequenceLength) -> usize {
2209    match length {
2210        SequenceLength::Fixed(len) => len,
2211        SequenceLength::Longest => sequences.iter().map(Vec::len).max().unwrap_or(0),
2212        SequenceLength::Shortest => sequences.iter().map(Vec::len).min().unwrap_or(0),
2213    }
2214}
2215
2216fn sequence_target_len(
2217    sequence_len: usize,
2218    resolved_length: usize,
2219    options: Doc2SequenceOptions,
2220) -> usize {
2221    match options.padding_direction {
2222        PaddingDirection::None => match options.length {
2223            SequenceLength::Fixed(len) => sequence_len.min(len),
2224            SequenceLength::Shortest => sequence_len.min(resolved_length),
2225            SequenceLength::Longest => sequence_len,
2226        },
2227        PaddingDirection::Left | PaddingDirection::Right => resolved_length,
2228    }
2229}
2230
2231fn sequence_tensor(
2232    embedding: &EmbeddingModel,
2233    sequence: &[SequenceToken],
2234    target_len: usize,
2235    padding_direction: PaddingDirection,
2236    padding_value: f64,
2237) -> BuiltinResult<Tensor> {
2238    let truncated_len = sequence.len().min(target_len);
2239    let pad_len = target_len.saturating_sub(truncated_len);
2240    let mut out = Vec::with_capacity(embedding.dimension * target_len);
2241    if padding_direction == PaddingDirection::Left {
2242        push_padding_columns(&mut out, embedding.dimension, pad_len, padding_value);
2243    }
2244    for token in sequence.iter().take(truncated_len) {
2245        match token {
2246            SequenceToken::Known(row) => {
2247                let start = row * embedding.dimension;
2248                out.extend_from_slice(&embedding.vectors[start..start + embedding.dimension]);
2249            }
2250            SequenceToken::UnknownNan => {
2251                out.extend(std::iter::repeat_n(f64::NAN, embedding.dimension));
2252            }
2253        }
2254    }
2255    if padding_direction == PaddingDirection::Right {
2256        push_padding_columns(&mut out, embedding.dimension, pad_len, padding_value);
2257    }
2258    Tensor::new(out, vec![embedding.dimension, target_len])
2259        .map_err(|err| embedding_error("doc2sequence", err))
2260}
2261
2262#[derive(Clone, Copy, Debug, PartialEq)]
2263enum IndexSequenceToken {
2264    Known(f64),
2265    UnknownNan,
2266}
2267
2268fn index_sequence_tensor(
2269    sequence: &[IndexSequenceToken],
2270    target_len: usize,
2271    padding_direction: PaddingDirection,
2272    padding_value: f64,
2273) -> BuiltinResult<Tensor> {
2274    let truncated_len = sequence.len().min(target_len);
2275    let pad_len = target_len.saturating_sub(truncated_len);
2276    let mut out = Vec::with_capacity(target_len);
2277    if padding_direction == PaddingDirection::Left {
2278        out.extend(std::iter::repeat_n(padding_value, pad_len));
2279    }
2280    for token in sequence.iter().take(truncated_len) {
2281        match token {
2282            IndexSequenceToken::Known(idx) => out.push(*idx),
2283            IndexSequenceToken::UnknownNan => out.push(f64::NAN),
2284        }
2285    }
2286    if padding_direction == PaddingDirection::Right {
2287        out.extend(std::iter::repeat_n(padding_value, pad_len));
2288    }
2289    Tensor::new(out, vec![1, target_len]).map_err(|err| embedding_error("doc2sequence", err))
2290}
2291
2292fn push_padding_columns(out: &mut Vec<f64>, dimension: usize, count: usize, padding_value: f64) {
2293    out.extend(std::iter::repeat_n(padding_value, dimension * count));
2294}
2295
2296fn parse_sequence_length(value: &Value) -> BuiltinResult<SequenceLength> {
2297    if matches!(
2298        value,
2299        Value::String(_) | Value::StringArray(_) | Value::CharArray(_)
2300    ) {
2301        let text = scalar_text(value, "doc2sequence")
2302            .map_err(|err| embedding_error("doc2sequence", err.to_string()))?;
2303        match text.trim().to_ascii_lowercase().as_str() {
2304            "longest" => return Ok(SequenceLength::Longest),
2305            "shortest" => return Ok(SequenceLength::Shortest),
2306            other => {
2307                if let Ok(value) = other.parse::<usize>() {
2308                    if value > 0 {
2309                        return Ok(SequenceLength::Fixed(value));
2310                    }
2311                }
2312                return Err(embedding_error(
2313                    "doc2sequence",
2314                    format!(
2315                        "doc2sequence: Length must be 'longest', 'shortest', or a positive integer, got '{other}'"
2316                    ),
2317                ));
2318            }
2319        }
2320    }
2321    Ok(SequenceLength::Fixed(parse_positive_integer(
2322        value,
2323        "doc2sequence",
2324    )?))
2325}
2326
2327fn parse_numeric_scalar(value: &Value, fn_name: &str, option_name: &str) -> BuiltinResult<f64> {
2328    let n = match value {
2329        Value::Num(value) => *value,
2330        Value::Int(value) => int_value_to_f64(value),
2331        Value::Tensor(tensor) if tensor.data.len() == 1 => tensor.data[0],
2332        other => {
2333            return Err(embedding_error(
2334                fn_name,
2335                format!("{fn_name}: {option_name} must be a numeric scalar, got {other:?}"),
2336            ));
2337        }
2338    };
2339    Ok(n)
2340}
2341
2342fn dense_doc2sequence_limit_error(fn_name: &str) -> crate::RuntimeError {
2343    embedding_error(
2344        fn_name,
2345        format!(
2346            "{fn_name}: output would exceed {MAX_DOC2SEQUENCE_DENSE_VALUES} dense values; use PaddingDirection 'none' or a smaller Length"
2347        ),
2348    )
2349}
2350
2351fn words_from_value(value: &Value, fn_name: &str) -> BuiltinResult<Vec<String>> {
2352    match value {
2353        Value::String(text) => Ok(vec![text.clone()]),
2354        Value::StringArray(array) => Ok(array.data.clone()),
2355        Value::CharArray(array) if array.rows <= 1 => Ok(vec![char_row_to_string(array)]),
2356        Value::CharArray(array) => {
2357            let mut words = Vec::with_capacity(array.rows);
2358            for row in 0..array.rows {
2359                let mut text = String::with_capacity(array.cols);
2360                for col in 0..array.cols {
2361                    text.push(array.data[row + col * array.rows]);
2362                }
2363                words.push(text.trim_end().to_string());
2364            }
2365            Ok(words)
2366        }
2367        Value::Cell(cell) => cell
2368            .data
2369            .iter()
2370            .map(|item| match item {
2371                Value::String(text) => Ok(text.clone()),
2372                Value::StringArray(array) if array.data.len() == 1 => Ok(array.data[0].clone()),
2373                Value::CharArray(array) if array.rows <= 1 => Ok(char_row_to_string(array)),
2374                other => Err(embedding_error(
2375                    fn_name,
2376                    format!("{fn_name}: cell word inputs must contain scalar text, got {other:?}"),
2377                )),
2378            })
2379            .collect(),
2380        other => Err(embedding_error(
2381            fn_name,
2382            format!("{fn_name}: expected string, character vector, or cell array of words, got {other:?}"),
2383        )),
2384    }
2385}
2386
2387pub(in crate::builtins::strings::text_analytics) fn build_word_lookup(
2388    vocabulary: &[String],
2389    ignore_case: bool,
2390) -> HashMap<String, usize> {
2391    let mut lookup = HashMap::new();
2392    for (idx, word) in vocabulary.iter().enumerate() {
2393        let key = if ignore_case {
2394            word.to_lowercase()
2395        } else {
2396            word.clone()
2397        };
2398        lookup.entry(key).or_insert(idx);
2399    }
2400    lookup
2401}
2402
2403fn row_slice(tensor: &Tensor, row: usize) -> Vec<f64> {
2404    (0..tensor.cols)
2405        .map(|col| tensor.data[row + col * tensor.rows])
2406        .collect()
2407}
2408
2409fn cosine_distance(lhs: &[f64], rhs: &[f64]) -> f64 {
2410    let mut dot = 0.0;
2411    let mut lhs_norm = 0.0;
2412    let mut rhs_norm = 0.0;
2413    for (&a, &b) in lhs.iter().zip(rhs.iter()) {
2414        dot += a * b;
2415        lhs_norm += a * a;
2416        rhs_norm += b * b;
2417    }
2418    if lhs_norm == 0.0 || rhs_norm == 0.0 {
2419        f64::INFINITY
2420    } else {
2421        1.0 - dot / (lhs_norm.sqrt() * rhs_norm.sqrt())
2422    }
2423}
2424
2425fn euclidean_distance(lhs: &[f64], rhs: &[f64]) -> f64 {
2426    lhs.iter()
2427        .zip(rhs.iter())
2428        .map(|(&a, &b)| {
2429            let delta = a - b;
2430            delta * delta
2431        })
2432        .sum::<f64>()
2433        .sqrt()
2434}
2435
2436fn compare_scores(left: f64, right: f64) -> Ordering {
2437    match (left.is_nan(), right.is_nan()) {
2438        (true, true) => Ordering::Equal,
2439        (true, false) => Ordering::Greater,
2440        (false, true) => Ordering::Less,
2441        (false, false) => left.partial_cmp(&right).unwrap_or(Ordering::Equal),
2442    }
2443}
2444
2445fn char_row_to_string(array: &CharArray) -> String {
2446    array.data.iter().collect()
2447}
2448
2449fn parse_bool_scalar(value: &Value, fn_name: &str) -> BuiltinResult<bool> {
2450    match value {
2451        Value::Bool(value) => Ok(*value),
2452        Value::Num(value) if *value == 0.0 || *value == 1.0 => Ok(*value != 0.0),
2453        Value::Tensor(tensor) if tensor.data.len() == 1 => match tensor.data[0] {
2454            0.0 => Ok(false),
2455            1.0 => Ok(true),
2456            other => Err(embedding_error(
2457                fn_name,
2458                format!("{fn_name}: logical scalar option must be true or false, got {other}"),
2459            )),
2460        },
2461        other => Err(embedding_error(
2462            fn_name,
2463            format!("{fn_name}: logical scalar option must be true or false, got {other:?}"),
2464        )),
2465    }
2466}
2467
2468fn int_value_to_f64(value: &runmat_builtins::IntValue) -> f64 {
2469    match value {
2470        runmat_builtins::IntValue::I8(value) => *value as f64,
2471        runmat_builtins::IntValue::I16(value) => *value as f64,
2472        runmat_builtins::IntValue::I32(value) => *value as f64,
2473        runmat_builtins::IntValue::I64(value) => *value as f64,
2474        runmat_builtins::IntValue::U8(value) => *value as f64,
2475        runmat_builtins::IntValue::U16(value) => *value as f64,
2476        runmat_builtins::IntValue::U32(value) => *value as f64,
2477        runmat_builtins::IntValue::U64(value) => *value as f64,
2478    }
2479}
2480
2481fn parse_positive_integer(value: &Value, fn_name: &str) -> BuiltinResult<usize> {
2482    let n = match value {
2483        Value::Num(value) => *value,
2484        Value::Int(value) => int_value_to_f64(value),
2485        Value::Tensor(tensor) if tensor.data.len() == 1 => tensor.data[0],
2486        other => {
2487            return Err(embedding_error(
2488                fn_name,
2489                format!("{fn_name}: expected positive integer scalar, got {other:?}"),
2490            ));
2491        }
2492    };
2493    if !n.is_finite() || n < 1.0 || n.fract() != 0.0 {
2494        return Err(embedding_error(
2495            fn_name,
2496            format!("{fn_name}: expected positive integer scalar, got {n}"),
2497        ));
2498    }
2499    Ok(n as usize)
2500}
2501
2502fn is_numeric_scalar(value: &Value) -> bool {
2503    matches!(value, Value::Num(_))
2504        || matches!(value, Value::Tensor(tensor) if tensor.data.len() == 1)
2505}
2506
2507fn is_zip_path(path: &Path) -> bool {
2508    matches!(
2509        path.extension()
2510            .and_then(|ext| ext.to_str())
2511            .map(|ext| ext.to_ascii_lowercase())
2512            .as_deref(),
2513        Some("zip")
2514    )
2515}
2516
2517fn looks_like_zip(bytes: &[u8]) -> bool {
2518    bytes.len() >= 4 && &bytes[..4] == b"PK\x03\x04"
2519}
2520
2521fn embedding_error(fn_name: &str, message: impl Into<String>) -> crate::RuntimeError {
2522    let descriptor = match fn_name {
2523        "fastTextWordEmbedding" => ERROR_FASTTEXT_INVALID_INPUT,
2524        "readWordEmbedding" => ERROR_READ_INVALID_INPUT,
2525        "writeWordEmbedding" => ERROR_WRITE_INVALID_INPUT,
2526        "trainWordEmbedding" => ERROR_TRAIN_INVALID_INPUT,
2527        "doc2sequence" => ERROR_DOC2SEQUENCE_INVALID_INPUT,
2528        "word2vec" => ERROR_WORD2VEC_INVALID_INPUT,
2529        "vec2word" => ERROR_VEC2WORD_INVALID_INPUT,
2530        _ => ERROR_WORD_EMBEDDING_INVALID_INPUT,
2531    };
2532    let builder = build_runtime_error(message.into()).with_builtin(fn_name);
2533    match descriptor.identifier {
2534        Some(identifier) => builder.with_identifier(identifier).build(),
2535        None => builder.build(),
2536    }
2537}
2538
2539fn embedding_error_with_source(
2540    fn_name: &str,
2541    message: impl Into<String>,
2542    source: impl std::error::Error + Send + Sync + 'static,
2543) -> crate::RuntimeError {
2544    let descriptor = match fn_name {
2545        "readWordEmbedding" => ERROR_READ_IO,
2546        "writeWordEmbedding" => ERROR_WRITE_IO,
2547        _ => ERROR_WORD_EMBEDDING_INVALID_INPUT,
2548    };
2549    let builder = build_runtime_error(message.into())
2550        .with_builtin(fn_name)
2551        .with_source(source);
2552    match descriptor.identifier {
2553        Some(identifier) => builder.with_identifier(identifier).build(),
2554        None => builder.build(),
2555    }
2556}
2557
2558#[cfg(test)]
2559mod tests {
2560    use super::*;
2561    use runmat_builtins::CellArray;
2562    use std::fs::File as StdFile;
2563    use std::io::Write;
2564    use tempfile::tempdir;
2565
2566    #[test]
2567    fn parses_glove_text_embedding() {
2568        let model = parse_embedding_text("king 1 0 0\nqueen 0.8 0.2 0\n", "test").unwrap();
2569        assert_eq!(model.dimension, 3);
2570        assert_eq!(model.vocabulary, vec!["king", "queen"]);
2571        assert_eq!(model.vectors, vec![1.0, 0.0, 0.0, 0.8, 0.2, 0.0]);
2572    }
2573
2574    #[test]
2575    fn parses_word2vec_header_and_last_duplicate_wins() {
2576        let model =
2577            parse_embedding_text("3 2\nalpha 1 0\nbeta 0 1\nalpha 0.5 0.5\n", "test").unwrap();
2578        assert_eq!(model.dimension, 2);
2579        assert_eq!(model.vocabulary, vec!["beta", "alpha"]);
2580        assert_eq!(model.vectors, vec![0.0, 1.0, 0.5, 0.5]);
2581    }
2582
2583    #[test]
2584    fn rejects_word2vec_header_row_mismatch() {
2585        let err = parse_embedding_text("3 2\nalpha 1 0\nbeta 0 1\n", "test").unwrap_err();
2586        assert!(err.to_string().contains("header declares 3 words"), "{err}");
2587    }
2588
2589    #[test]
2590    fn rejects_inconsistent_embedding_dimensions() {
2591        let err = parse_embedding_text("alpha 1 0\nbeta 0 1 2\n", "test").unwrap_err();
2592        assert!(
2593            err.to_string()
2594                .contains("has 3 vector values but expected 2"),
2595            "{err}"
2596        );
2597    }
2598
2599    fn tokenized_document_object(rows: Vec<Vec<&str>>) -> ObjectInstance {
2600        let values = rows
2601            .into_iter()
2602            .map(|row| {
2603                let len = row.len();
2604                Value::StringArray(
2605                    StringArray::new(
2606                        row.into_iter().map(str::to_string).collect::<Vec<_>>(),
2607                        vec![1, len],
2608                    )
2609                    .unwrap(),
2610                )
2611            })
2612            .collect::<Vec<_>>();
2613        let rows = values.len();
2614        let documents = CellArray::new(values, rows, 1).unwrap();
2615        let mut object = ObjectInstance::new(TOKENIZED_DOCUMENT_CLASS.to_string());
2616        object
2617            .properties
2618            .insert("Documents".to_string(), Value::Cell(documents));
2619        object
2620            .properties
2621            .insert("NumDocuments".to_string(), Value::Num(rows as f64));
2622        object.properties.insert(
2623            "Shape".to_string(),
2624            Value::Tensor(Tensor::new(vec![rows as f64, 1.0], vec![1, 2]).unwrap()),
2625        );
2626        object
2627    }
2628
2629    #[tokio::test]
2630    async fn read_word_embedding_reads_plain_text_file() {
2631        let dir = tempdir().unwrap();
2632        let path = dir.path().join("emb.vec");
2633        std::fs::write(&path, "3 2\nred 1 0\nblue 0 1\ngreen 0.5 0.5\n").unwrap();
2634        let value = read_word_embedding_builtin(Value::from(path.to_string_lossy().to_string()))
2635            .await
2636            .unwrap();
2637        let Value::Object(object) = value else {
2638            panic!("expected object");
2639        };
2640        assert!(object.is_class(WORD_EMBEDDING_CLASS));
2641        assert_eq!(object.properties.get("Dimension"), Some(&Value::Num(2.0)));
2642    }
2643
2644    #[tokio::test]
2645    async fn read_word_embedding_reads_zip_file() {
2646        let dir = tempdir().unwrap();
2647        let path = dir.path().join("emb.zip");
2648        let file = StdFile::create(&path).unwrap();
2649        let mut zip = zip::ZipWriter::new(file);
2650        zip.start_file(
2651            "model.vec",
2652            zip::write::SimpleFileOptions::default()
2653                .compression_method(zip::CompressionMethod::Deflated),
2654        )
2655        .unwrap();
2656        zip.write_all(b"2 2\nleft 1 0\nright 0 1\n").unwrap();
2657        zip.finish().unwrap();
2658
2659        let value = read_word_embedding_builtin(Value::from(path.to_string_lossy().to_string()))
2660            .await
2661            .unwrap();
2662        let Value::Object(object) = value else {
2663            panic!("expected object");
2664        };
2665        assert!(object.is_class(WORD_EMBEDDING_CLASS));
2666        let model = embedding_from_object(&object, "test").unwrap();
2667        assert_eq!(model.vocabulary, vec!["left", "right"]);
2668    }
2669
2670    #[tokio::test]
2671    async fn read_word_embedding_accepts_scalar_cell_filename() {
2672        let dir = tempdir().unwrap();
2673        let path = dir.path().join("cell-name.vec");
2674        std::fs::write(&path, "1 2\nonly 0.25 0.75\n").unwrap();
2675        let filename = Value::Cell(
2676            CellArray::new(
2677                vec![Value::CharArray(CharArray::new_row(
2678                    &path.to_string_lossy(),
2679                ))],
2680                1,
2681                1,
2682            )
2683            .unwrap(),
2684        );
2685
2686        let value = read_word_embedding_builtin(filename).await.unwrap();
2687        let Value::Object(object) = value else {
2688            panic!("expected object");
2689        };
2690        let model = embedding_from_object(&object, "test").unwrap();
2691        assert_eq!(model.vocabulary, vec!["only"]);
2692        assert_eq!(model.vectors, vec![0.25, 0.75]);
2693    }
2694
2695    #[tokio::test]
2696    async fn write_word_embedding_writes_word2vec_text_and_round_trips() {
2697        let dir = tempdir().unwrap();
2698        let path = dir.path().join("roundtrip.vec");
2699        let emb = embedding_object(EmbeddingModel {
2700            vocabulary: vec!["alpha".into(), "beta".into()],
2701            vectors: vec![1.0, 2.5, 3.0, 4.25],
2702            dimension: 2,
2703        })
2704        .unwrap();
2705        let filename = Value::Cell(
2706            CellArray::new(
2707                vec![Value::CharArray(CharArray::new_row(
2708                    &path.to_string_lossy(),
2709                ))],
2710                1,
2711                1,
2712            )
2713            .unwrap(),
2714        );
2715
2716        let result = write_word_embedding_builtin(emb, filename).await.unwrap();
2717        assert_eq!(result, Value::Num(0.0));
2718        let contents = std::fs::read_to_string(&path).unwrap();
2719        assert_eq!(contents, "2 2\nalpha 1 2.5\nbeta 3 4.25\n");
2720
2721        let value = read_word_embedding_builtin(Value::String(path.to_string_lossy().to_string()))
2722            .await
2723            .unwrap();
2724        let Value::Object(object) = value else {
2725            panic!("expected object");
2726        };
2727        let model = embedding_from_object(&object, "test").unwrap();
2728        assert_eq!(model.vocabulary, vec!["alpha", "beta"]);
2729        assert_eq!(model.dimension, 2);
2730        assert_eq!(model.vectors, vec![1.0, 2.5, 3.0, 4.25]);
2731    }
2732
2733    #[tokio::test]
2734    async fn write_word_embedding_rejects_bad_inputs_and_vectors() {
2735        let dir = tempdir().unwrap();
2736        let path = Value::String(dir.path().join("bad.vec").to_string_lossy().to_string());
2737        let err =
2738            write_word_embedding_builtin(Value::String("not an embedding".into()), path.clone())
2739                .await
2740                .unwrap_err();
2741        assert!(
2742            err.to_string().contains("expected wordEmbedding object"),
2743            "{err}"
2744        );
2745
2746        let mut object = ObjectInstance::new(WORD_EMBEDDING_CLASS.to_string());
2747        object
2748            .properties
2749            .insert("Dimension".to_string(), Value::Num(1.0));
2750        object.properties.insert(
2751            "Vocabulary".to_string(),
2752            Value::StringArray(StringArray::new(vec!["bad word".into()], vec![1, 1]).unwrap()),
2753        );
2754        object.properties.insert(
2755            VECTOR_PROPERTY.to_string(),
2756            Value::Tensor(Tensor::new(vec![f64::NAN], vec![1, 1]).unwrap()),
2757        );
2758        let err = write_word_embedding_builtin(Value::Object(object), path)
2759            .await
2760            .unwrap_err();
2761        assert!(
2762            err.to_string().contains("cannot contain whitespace")
2763                || err.to_string().contains("must be finite"),
2764            "{err}"
2765        );
2766
2767        let valid = embedding_object(EmbeddingModel {
2768            vocabulary: vec!["ok".into()],
2769            vectors: vec![1.0],
2770            dimension: 1,
2771        })
2772        .unwrap();
2773        let missing_parent = Value::String(
2774            dir.path()
2775                .join("missing")
2776                .join("parent.vec")
2777                .to_string_lossy()
2778                .to_string(),
2779        );
2780        let err = write_word_embedding_builtin(valid, missing_parent)
2781            .await
2782            .unwrap_err();
2783        assert!(
2784            err.identifier() == Some("RunMat:writeWordEmbedding:IOError"),
2785            "{err:?}"
2786        );
2787    }
2788
2789    #[tokio::test]
2790    async fn fast_text_word_embedding_returns_compact_300d_model() {
2791        let value = fast_text_word_embedding_builtin(vec![]).await.unwrap();
2792        let Value::Object(object) = value else {
2793            panic!("expected wordEmbedding object");
2794        };
2795        assert!(object.is_class(WORD_EMBEDDING_CLASS));
2796        assert_eq!(object.properties.get("Dimension"), Some(&Value::Num(300.0)));
2797
2798        let italy = word2vec_builtin(vec![
2799            Value::Object(object.clone()),
2800            Value::String("Italy".into()),
2801        ])
2802        .await
2803        .unwrap();
2804        let rome = word2vec_builtin(vec![
2805            Value::Object(object.clone()),
2806            Value::String("Rome".into()),
2807        ])
2808        .await
2809        .unwrap();
2810        let paris = word2vec_builtin(vec![
2811            Value::Object(object.clone()),
2812            Value::String("Paris".into()),
2813        ])
2814        .await
2815        .unwrap();
2816        let (Value::Tensor(italy), Value::Tensor(rome), Value::Tensor(paris)) =
2817            (italy, rome, paris)
2818        else {
2819            panic!("expected tensors");
2820        };
2821        let query = italy
2822            .data
2823            .iter()
2824            .zip(&rome.data)
2825            .zip(&paris.data)
2826            .map(|((i, r), p)| i - r + p)
2827            .collect::<Vec<_>>();
2828        let nearest = vec2word_builtin(vec![
2829            Value::Object(object),
2830            Value::Tensor(Tensor::new(query, vec![1, 300]).unwrap()),
2831            Value::Num(1.0),
2832        ])
2833        .await
2834        .unwrap();
2835        let Value::OutputList(outputs) = nearest else {
2836            panic!("expected output list");
2837        };
2838        let Value::StringArray(words) = &outputs[0] else {
2839            panic!("expected nearest words");
2840        };
2841        assert_eq!(words.data, vec!["France"]);
2842
2843        let err = fast_text_word_embedding_builtin(vec![Value::Num(1.0)])
2844            .await
2845            .unwrap_err();
2846        assert!(err.to_string().contains("expected no input"), "{err}");
2847    }
2848
2849    #[tokio::test]
2850    async fn train_word_embedding_trains_from_text_file() {
2851        let dir = tempdir().unwrap();
2852        let path = dir.path().join("training.txt");
2853        std::fs::write(&path, "alpha beta alpha\nbeta gamma alpha\n").unwrap();
2854        let value = train_word_embedding_builtin(vec![
2855            Value::from(path.to_string_lossy().to_string()),
2856            Value::String("Dimension".into()),
2857            Value::Num(8.0),
2858            Value::String("Window".into()),
2859            Value::Num(1.0),
2860            Value::String("MinCount".into()),
2861            Value::Num(1.0),
2862            Value::String("NGramRange".into()),
2863            Value::Tensor(Tensor::new(vec![0.0, 0.0], vec![1, 2]).unwrap()),
2864            Value::String("Verbose".into()),
2865            Value::Bool(false),
2866        ])
2867        .await
2868        .unwrap();
2869        let Value::Object(object) = value else {
2870            panic!("expected wordEmbedding object");
2871        };
2872        assert!(object.is_class(WORD_EMBEDDING_CLASS));
2873        let model = embedding_from_object(&object, "test").unwrap();
2874        assert_eq!(model.dimension, 8);
2875        assert_eq!(model.vocabulary, vec!["alpha", "beta", "gamma"]);
2876        assert_eq!(model.vectors.len(), 24);
2877
2878        let lookup = word2vec_builtin(vec![Value::Object(object), Value::String("alpha".into())])
2879            .await
2880            .unwrap();
2881        let Value::Tensor(tensor) = lookup else {
2882            panic!("expected tensor");
2883        };
2884        assert_eq!(tensor.rows, 1);
2885        assert_eq!(tensor.cols, 8);
2886        assert!(tensor.data.iter().any(|value| value.abs() > 0.0));
2887    }
2888
2889    #[tokio::test]
2890    async fn train_word_embedding_trains_from_tokenized_document_object() {
2891        let object = tokenized_document_object(vec![vec!["red", "blue"], vec!["red", "green"]]);
2892
2893        let value = train_word_embedding_builtin(vec![
2894            Value::Object(object),
2895            Value::String("Dimension".into()),
2896            Value::Num(6.0),
2897            Value::String("MinCount".into()),
2898            Value::Num(1.0),
2899            Value::String("Model".into()),
2900            Value::String("cbow".into()),
2901            Value::String("LossFunction".into()),
2902            Value::String("softmax".into()),
2903        ])
2904        .await
2905        .unwrap();
2906        let Value::Object(object) = value else {
2907            panic!("expected wordEmbedding object");
2908        };
2909        let model = embedding_from_object(&object, "test").unwrap();
2910        assert_eq!(model.dimension, 6);
2911        assert_eq!(model.vocabulary, vec!["red", "blue", "green"]);
2912    }
2913
2914    #[tokio::test]
2915    async fn doc2sequence_pads_to_longest_and_discards_unknown_words() {
2916        let model = EmbeddingModel {
2917            vocabulary: vec!["alpha".into(), "beta".into()],
2918            vectors: vec![1.0, 10.0, 2.0, 20.0],
2919            dimension: 2,
2920        };
2921        let emb = embedding_object(model).unwrap();
2922        let documents = Value::Object(tokenized_document_object(vec![
2923            vec!["alpha", "beta"],
2924            vec!["missing", "beta"],
2925        ]));
2926
2927        let result = doc2sequence_builtin(vec![emb, documents]).await.unwrap();
2928        let Value::Cell(cell) = result else {
2929            panic!("expected cell array");
2930        };
2931        assert_eq!(cell.rows, 2);
2932        assert_eq!(cell.cols, 1);
2933
2934        let Value::Tensor(first) = &cell.data[0] else {
2935            panic!("expected first tensor");
2936        };
2937        assert_eq!(first.shape, vec![2, 2]);
2938        assert_eq!(first.data, vec![1.0, 10.0, 2.0, 20.0]);
2939
2940        let Value::Tensor(second) = &cell.data[1] else {
2941            panic!("expected second tensor");
2942        };
2943        assert_eq!(second.shape, vec![2, 2]);
2944        assert_eq!(second.data, vec![0.0, 0.0, 2.0, 20.0]);
2945    }
2946
2947    #[tokio::test]
2948    async fn doc2sequence_supports_unknown_nan_right_padding_and_fixed_length() {
2949        let model = EmbeddingModel {
2950            vocabulary: vec!["alpha".into(), "beta".into()],
2951            vectors: vec![1.0, 10.0, 2.0, 20.0],
2952            dimension: 2,
2953        };
2954        let emb = embedding_object(model).unwrap();
2955        let documents = Value::Object(tokenized_document_object(vec![
2956            vec!["alpha", "missing"],
2957            vec!["beta"],
2958        ]));
2959
2960        let result = doc2sequence_builtin(vec![
2961            emb,
2962            documents,
2963            Value::String("UnknownWord".into()),
2964            Value::String("nan".into()),
2965            Value::String("PaddingDirection".into()),
2966            Value::String("right".into()),
2967            Value::String("PaddingValue".into()),
2968            Value::Num(-5.0),
2969            Value::String("Length".into()),
2970            Value::Num(3.0),
2971        ])
2972        .await
2973        .unwrap();
2974        let Value::Cell(cell) = result else {
2975            panic!("expected cell array");
2976        };
2977        let Value::Tensor(first) = &cell.data[0] else {
2978            panic!("expected first tensor");
2979        };
2980        assert_eq!(first.shape, vec![2, 3]);
2981        assert_eq!(first.data[0..2], [1.0, 10.0]);
2982        assert!(first.data[2].is_nan());
2983        assert!(first.data[3].is_nan());
2984        assert_eq!(first.data[4..6], [-5.0, -5.0]);
2985
2986        let Value::Tensor(second) = &cell.data[1] else {
2987            panic!("expected second tensor");
2988        };
2989        assert_eq!(second.data, vec![2.0, 20.0, -5.0, -5.0, -5.0, -5.0]);
2990    }
2991
2992    #[tokio::test]
2993    async fn doc2sequence_none_padding_keeps_per_document_lengths_and_truncates_right() {
2994        let model = EmbeddingModel {
2995            vocabulary: vec!["a".into(), "b".into(), "c".into()],
2996            vectors: vec![1.0, 11.0, 2.0, 22.0, 3.0, 33.0],
2997            dimension: 2,
2998        };
2999        let emb = embedding_object(model).unwrap();
3000        let documents = Value::Object(tokenized_document_object(vec![
3001            vec!["a", "b", "c"],
3002            vec!["a"],
3003        ]));
3004
3005        let result = doc2sequence_builtin(vec![
3006            emb,
3007            documents,
3008            Value::String("PaddingDirection".into()),
3009            Value::String("none".into()),
3010            Value::String("Length".into()),
3011            Value::Num(2.0),
3012        ])
3013        .await
3014        .unwrap();
3015        let Value::Cell(cell) = result else {
3016            panic!("expected cell array");
3017        };
3018        let Value::Tensor(first) = &cell.data[0] else {
3019            panic!("expected first tensor");
3020        };
3021        assert_eq!(first.shape, vec![2, 2]);
3022        assert_eq!(first.data, vec![1.0, 11.0, 2.0, 22.0]);
3023
3024        let Value::Tensor(second) = &cell.data[1] else {
3025            panic!("expected second tensor");
3026        };
3027        assert_eq!(second.shape, vec![2, 1]);
3028        assert_eq!(second.data, vec![1.0, 11.0]);
3029    }
3030
3031    #[tokio::test]
3032    async fn doc2sequence_supports_word_encoding_index_sequences_and_invalid_options() {
3033        let model = EmbeddingModel {
3034            vocabulary: vec!["alpha".into()],
3035            vectors: vec![1.0, 10.0],
3036            dimension: 2,
3037        };
3038        let mut documents_object =
3039            tokenized_document_object(vec![vec!["alpha", "missing", "beta"], vec!["beta"]]);
3040        documents_object.properties.insert(
3041            "Shape".to_string(),
3042            Value::Tensor(Tensor::new(vec![1.0, 2.0], vec![1, 2]).unwrap()),
3043        );
3044        let documents = Value::Object(documents_object);
3045        let mut word_encoding = ObjectInstance::new(WORD_ENCODING_CLASS.to_string());
3046        word_encoding
3047            .properties
3048            .insert("NumWords".to_string(), Value::Num(2.0));
3049        word_encoding.properties.insert(
3050            "Vocabulary".to_string(),
3051            Value::StringArray(
3052                StringArray::new(vec!["alpha".into(), "beta".into()], vec![1, 2]).unwrap(),
3053            ),
3054        );
3055        let result = doc2sequence_builtin(vec![
3056            Value::Object(word_encoding),
3057            documents.clone(),
3058            Value::String("UnknownWord".into()),
3059            Value::String("nan".into()),
3060            Value::String("PaddingDirection".into()),
3061            Value::String("right".into()),
3062            Value::String("Length".into()),
3063            Value::Num(4.0),
3064        ])
3065        .await
3066        .unwrap();
3067        let Value::Cell(cell) = result else {
3068            panic!("expected cell array");
3069        };
3070        assert_eq!(cell.shape, vec![1, 2]);
3071        assert_eq!(cell.rows, 1);
3072        assert_eq!(cell.cols, 2);
3073        let Value::Tensor(first) = &cell.data[0] else {
3074            panic!("expected first tensor");
3075        };
3076        assert_eq!(first.shape, vec![1, 4]);
3077        assert_eq!(first.data[0], 1.0);
3078        assert!(first.data[1].is_nan());
3079        assert_eq!(first.data[2..4], [2.0, 0.0]);
3080        let Value::Tensor(second) = &cell.data[1] else {
3081            panic!("expected second tensor");
3082        };
3083        assert_eq!(second.shape, vec![1, 4]);
3084        assert_eq!(second.data, vec![2.0, 0.0, 0.0, 0.0]);
3085
3086        let err = doc2sequence_builtin(vec![Value::Num(1.0), documents.clone()])
3087            .await
3088            .unwrap_err();
3089        assert!(
3090            err.to_string()
3091                .contains("wordEmbedding or wordEncoding object"),
3092            "{err}"
3093        );
3094
3095        let err = doc2sequence_builtin(vec![
3096            embedding_object(model).unwrap(),
3097            documents,
3098            Value::String("PaddingDirection".into()),
3099            Value::String("middle".into()),
3100        ])
3101        .await
3102        .unwrap_err();
3103        assert!(err.to_string().contains("PaddingDirection"), "{err}");
3104    }
3105
3106    #[tokio::test]
3107    async fn doc2sequence_word_encoding_supports_left_none_and_shortest_length() {
3108        let mut word_encoding = ObjectInstance::new(WORD_ENCODING_CLASS.to_string());
3109        word_encoding
3110            .properties
3111            .insert("NumWords".to_string(), Value::Num(3.0));
3112        word_encoding.properties.insert(
3113            "Vocabulary".to_string(),
3114            Value::StringArray(
3115                StringArray::new(
3116                    vec!["alpha".into(), "beta".into(), "gamma".into()],
3117                    vec![1, 3],
3118                )
3119                .unwrap(),
3120            ),
3121        );
3122        let documents = Value::Object(tokenized_document_object(vec![
3123            vec!["alpha", "beta", "gamma"],
3124            vec!["gamma"],
3125        ]));
3126
3127        let left = doc2sequence_builtin(vec![
3128            Value::Object(word_encoding.clone()),
3129            documents.clone(),
3130        ])
3131        .await
3132        .unwrap();
3133        let Value::Cell(left) = left else {
3134            panic!("expected cell array");
3135        };
3136        let Value::Tensor(first) = &left.data[0] else {
3137            panic!("expected first tensor");
3138        };
3139        assert_eq!(first.shape, vec![1, 3]);
3140        assert_eq!(first.data, vec![1.0, 2.0, 3.0]);
3141        let Value::Tensor(second) = &left.data[1] else {
3142            panic!("expected second tensor");
3143        };
3144        assert_eq!(second.shape, vec![1, 3]);
3145        assert_eq!(second.data, vec![0.0, 0.0, 3.0]);
3146
3147        let none = doc2sequence_builtin(vec![
3148            Value::Object(word_encoding.clone()),
3149            documents.clone(),
3150            Value::String("PaddingDirection".into()),
3151            Value::String("none".into()),
3152        ])
3153        .await
3154        .unwrap();
3155        let Value::Cell(none) = none else {
3156            panic!("expected cell array");
3157        };
3158        let Value::Tensor(second) = &none.data[1] else {
3159            panic!("expected second tensor");
3160        };
3161        assert_eq!(second.shape, vec![1, 1]);
3162        assert_eq!(second.data, vec![3.0]);
3163
3164        let shortest = doc2sequence_builtin(vec![
3165            Value::Object(word_encoding),
3166            documents,
3167            Value::String("Length".into()),
3168            Value::String("shortest".into()),
3169        ])
3170        .await
3171        .unwrap();
3172        let Value::Cell(shortest) = shortest else {
3173            panic!("expected cell array");
3174        };
3175        let Value::Tensor(first) = &shortest.data[0] else {
3176            panic!("expected first tensor");
3177        };
3178        assert_eq!(first.shape, vec![1, 1]);
3179        assert_eq!(first.data, vec![1.0]);
3180        let Value::Tensor(second) = &shortest.data[1] else {
3181            panic!("expected second tensor");
3182        };
3183        assert_eq!(second.shape, vec![1, 1]);
3184        assert_eq!(second.data, vec![3.0]);
3185    }
3186
3187    #[tokio::test]
3188    async fn doc2sequence_allows_nan_padding_value() {
3189        let model = EmbeddingModel {
3190            vocabulary: vec!["alpha".into()],
3191            vectors: vec![1.0, 10.0],
3192            dimension: 2,
3193        };
3194        let emb = embedding_object(model).unwrap();
3195        let documents = Value::Object(tokenized_document_object(vec![vec!["alpha"]]));
3196        let result = doc2sequence_builtin(vec![
3197            emb,
3198            documents,
3199            Value::String("PaddingDirection".into()),
3200            Value::String("left".into()),
3201            Value::String("PaddingValue".into()),
3202            Value::Num(f64::NAN),
3203            Value::String("Length".into()),
3204            Value::Num(2.0),
3205        ])
3206        .await
3207        .unwrap();
3208        let Value::Cell(cell) = result else {
3209            panic!("expected cell array");
3210        };
3211        let Value::Tensor(sequence) = &cell.data[0] else {
3212            panic!("expected tensor");
3213        };
3214        assert_eq!(sequence.shape, vec![2, 2]);
3215        assert!(sequence.data[0].is_nan());
3216        assert!(sequence.data[1].is_nan());
3217        assert_eq!(sequence.data[2..4], [1.0, 10.0]);
3218    }
3219
3220    #[tokio::test]
3221    async fn train_word_embedding_honors_min_count_and_option_validation() {
3222        let err = train_word_embedding_builtin(vec![
3223            Value::String("missing.txt".into()),
3224            Value::String("LossFunction".into()),
3225            Value::String("hs".into()),
3226            Value::String("NumNegativeSamples".into()),
3227            Value::Num(3.0),
3228        ])
3229        .await
3230        .unwrap_err();
3231        assert!(err.to_string().contains("NumNegativeSamples"), "{err}");
3232
3233        let dir = tempdir().unwrap();
3234        let path = dir.path().join("training.txt");
3235        std::fs::write(&path, "solo once\n").unwrap();
3236        let err = train_word_embedding_builtin(vec![
3237            Value::from(path.to_string_lossy().to_string()),
3238            Value::String("MinCount".into()),
3239            Value::Num(2.0),
3240            Value::String("Verbose".into()),
3241            Value::Num(0.0),
3242        ])
3243        .await
3244        .unwrap_err();
3245        assert!(err.to_string().contains("no vocabulary words"), "{err}");
3246    }
3247
3248    #[tokio::test]
3249    async fn word2vec_returns_rows_and_nan_for_missing_words() {
3250        let model = EmbeddingModel {
3251            vocabulary: vec!["king".into(), "queen".into()],
3252            vectors: vec![1.0, 0.0, 0.0, 1.0],
3253            dimension: 2,
3254        };
3255        let emb = embedding_object(model).unwrap();
3256        let words = Value::StringArray(
3257            StringArray::new(vec!["queen".into(), "missing".into()], vec![1, 2]).unwrap(),
3258        );
3259        let result = word2vec_builtin(vec![emb, words]).await.unwrap();
3260        let Value::Tensor(tensor) = result else {
3261            panic!("expected tensor");
3262        };
3263        assert_eq!(tensor.rows, 2);
3264        assert_eq!(tensor.cols, 2);
3265        assert_eq!(tensor.data[0], 0.0);
3266        assert_eq!(tensor.data[2], 1.0);
3267        assert!(tensor.data[1].is_nan());
3268        assert!(tensor.data[3].is_nan());
3269    }
3270
3271    #[tokio::test]
3272    async fn word2vec_ignore_case_uses_first_case_match() {
3273        let model = EmbeddingModel {
3274            vocabulary: vec!["Alpha".into(), "alpha".into()],
3275            vectors: vec![1.0, 0.0, 0.0, 1.0],
3276            dimension: 2,
3277        };
3278        let emb = embedding_object(model).unwrap();
3279        let result = word2vec_builtin(vec![
3280            emb,
3281            Value::String("ALPHA".into()),
3282            Value::String("IgnoreCase".into()),
3283            Value::Bool(true),
3284        ])
3285        .await
3286        .unwrap();
3287        let Value::Tensor(tensor) = result else {
3288            panic!("expected tensor");
3289        };
3290        assert_eq!(tensor.data, vec![1.0, 0.0]);
3291    }
3292
3293    #[tokio::test]
3294    async fn vec2word_returns_words_and_distances() {
3295        let model = EmbeddingModel {
3296            vocabulary: vec!["east".into(), "north".into(), "mix".into()],
3297            vectors: vec![1.0, 0.0, 0.0, 1.0, 0.7, 0.7],
3298            dimension: 2,
3299        };
3300        let emb = embedding_object(model).unwrap();
3301        let query = Value::Tensor(Tensor::new(vec![0.6, 0.8], vec![1, 2]).unwrap());
3302        let result = vec2word_builtin(vec![
3303            emb,
3304            query,
3305            Value::Num(2.0),
3306            Value::String("Distance".into()),
3307            Value::String("cosine".into()),
3308        ])
3309        .await
3310        .unwrap();
3311        let Value::OutputList(outputs) = result else {
3312            panic!("expected output list");
3313        };
3314        let Value::StringArray(words) = &outputs[0] else {
3315            panic!("expected words");
3316        };
3317        assert_eq!(words.data[0], "mix");
3318        assert_eq!(words.data.len(), 2);
3319        let Value::Tensor(dist) = &outputs[1] else {
3320            panic!("expected distances");
3321        };
3322        assert!(dist.data[0] < dist.data[1]);
3323    }
3324
3325    #[tokio::test]
3326    async fn vec2word_rejects_wrong_vector_dimension() {
3327        let model = EmbeddingModel {
3328            vocabulary: vec!["east".into()],
3329            vectors: vec![1.0, 0.0],
3330            dimension: 2,
3331        };
3332        let emb = embedding_object(model).unwrap();
3333        let query = Value::Tensor(Tensor::new(vec![1.0, 0.0, 0.5], vec![1, 3]).unwrap());
3334        let err = vec2word_builtin(vec![emb, query]).await.unwrap_err();
3335        assert!(err.to_string().contains("must have 2 columns"), "{err}");
3336    }
3337
3338    #[tokio::test]
3339    async fn vec2word_rejects_invalid_k_as_k_not_option_name() {
3340        let model = EmbeddingModel {
3341            vocabulary: vec!["east".into()],
3342            vectors: vec![1.0, 0.0],
3343            dimension: 2,
3344        };
3345        let emb = embedding_object(model).unwrap();
3346        let query = Value::Tensor(Tensor::new(vec![1.0, 0.0], vec![1, 2]).unwrap());
3347        let err = vec2word_builtin(vec![emb, query, Value::Num(0.0)])
3348            .await
3349            .unwrap_err();
3350        assert!(
3351            err.to_string().contains("expected positive integer scalar"),
3352            "{err}"
3353        );
3354    }
3355}