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