1mod epm_types;
13pub use epm_types::*;
14
15mod epm_processing;
16use epm_processing::{
17 add_positional_encoding, apply_ngram, apply_stop_word_filter, build_vocab,
18 inverse_document_frequencies, quantize_to_byte, reduce_dimensions, term_frequencies,
19 tokenize_text, CorpusStats, PipelineState,
20};
21pub use epm_processing::{l2_normalize, mean_pool, random_projection};
22
23use std::collections::HashMap;
24use std::sync::{Arc, Mutex};
25use std::time::Instant;
26
27pub struct EmbeddingPipelineManager {
37 config: EpmPipelineConfig,
38 state: Arc<Mutex<PipelineState>>,
39}
40
41impl EmbeddingPipelineManager {
42 pub fn new(config: EpmPipelineConfig) -> Result<Self, EpmPipelineError> {
44 let mgr = Self {
45 config,
46 state: Arc::new(Mutex::new(PipelineState::default())),
47 };
48 mgr.validate_config()?;
49 Ok(mgr)
50 }
51
52 pub fn process_text(
57 &self,
58 ids: Vec<String>,
59 texts: Vec<String>,
60 corpus: Option<&[String]>,
61 ) -> Result<EmbeddingBatch, EpmPipelineError> {
62 if ids.is_empty() || texts.is_empty() {
63 return Err(EpmPipelineError::EmptyInput);
64 }
65 if ids.len() != texts.len() {
66 return Err(EpmPipelineError::InvalidConfig(format!(
67 "ids.len() ({}) != texts.len() ({})",
68 ids.len(),
69 texts.len()
70 )));
71 }
72
73 let batch_start = Instant::now();
74 let n = texts.len();
75
76 let idf_corpus: Vec<String> = match corpus {
78 Some(c) => c.to_vec(),
79 None => texts.clone(),
80 };
81
82 let mut token_lists: Vec<Vec<String>> = texts.iter().map(|t| vec![t.clone()]).collect();
85
86 let mut stage_state = self
88 .state
89 .lock()
90 .map_err(|e| EpmPipelineError::ProcessingFailed(format!("mutex poisoned: {e}")))?;
91
92 let mut embeddings_opt: Option<Vec<Vec<f64>>> = None;
94
95 for stage in &self.config.stages {
96 let stage_start = Instant::now();
97
98 match stage {
99 EpmPipelineStage::Tokenize {
100 lowercase,
101 strip_punct,
102 } => {
103 token_lists = texts
104 .iter()
105 .map(|t| tokenize_text(t, *lowercase, *strip_punct))
106 .collect();
107 }
108 EpmPipelineStage::StopWordFilter(stop_words) => {
109 token_lists = token_lists
110 .into_iter()
111 .map(|toks| apply_stop_word_filter(toks, stop_words))
112 .collect();
113 }
114 EpmPipelineStage::NGram { n } => {
115 token_lists = token_lists
116 .iter()
117 .map(|toks| apply_ngram(toks, *n))
118 .collect();
119 }
120 EpmPipelineStage::TfIdfWeighting => {
121 let corpus_tokens: Vec<Vec<String>> = idf_corpus
123 .iter()
124 .map(|t| tokenize_text(t, true, true))
125 .collect();
126 let idf = inverse_document_frequencies(&corpus_tokens);
127 let tf_maps: Vec<HashMap<String, f64>> = token_lists
128 .iter()
129 .map(|toks| term_frequencies(toks))
130 .collect();
131 let vocab = build_vocab(&tf_maps);
132 if vocab.is_empty() {
133 return Err(EpmPipelineError::StageError {
134 stage: "TfIdfWeighting".to_string(),
135 reason: "empty vocabulary".to_string(),
136 });
137 }
138 embeddings_opt = Some(
139 tf_maps
140 .iter()
141 .map(|tf| epm_processing::tfidf_vector(tf, &idf, &vocab))
142 .collect(),
143 );
144 }
145 EpmPipelineStage::L2Normalize => {
147 let embs = embeddings_opt.get_or_insert_with(|| {
148 token_lists
149 .iter()
150 .map(|toks| toks.iter().map(|_| 1.0_f64).collect())
151 .collect()
152 });
153 for v in embs.iter_mut() {
154 l2_normalize(v);
155 }
156 }
157 EpmPipelineStage::DimensionReduce { target_dim, method } => {
158 let embs = embeddings_opt.get_or_insert_with(|| {
159 token_lists
160 .iter()
161 .map(|toks| toks.iter().map(|_| 1.0_f64).collect())
162 .collect()
163 });
164 let stats = CorpusStats::from_embeddings(embs);
165 let reduced: Vec<Vec<f64>> = embs
166 .iter()
167 .map(|v| reduce_dimensions(v, *target_dim, method, stats.as_ref()))
168 .collect();
169 *embs = reduced;
170 }
171 EpmPipelineStage::QuantizeToByte => {
172 let embs = embeddings_opt.get_or_insert_with(|| {
173 token_lists
174 .iter()
175 .map(|toks| toks.iter().map(|_| 1.0_f64).collect())
176 .collect()
177 });
178 for v in embs.iter_mut() {
179 *v = quantize_to_byte(v);
180 }
181 }
182 EpmPipelineStage::AddPositionalEncoding { max_len } => {
183 let embs = embeddings_opt.get_or_insert_with(|| {
184 token_lists
185 .iter()
186 .map(|toks| toks.iter().map(|_| 1.0_f64).collect())
187 .collect()
188 });
189 for (pos, v) in embs.iter_mut().enumerate() {
190 add_positional_encoding(v, pos, *max_len);
191 }
192 }
193 }
194
195 let elapsed_us = stage_start.elapsed().as_micros() as u64;
196 stage_state.record_stage(stage.name(), elapsed_us, n);
197 }
198
199 let output_embeddings = match embeddings_opt {
201 Some(e) => e,
202 None => {
203 let tf_maps: Vec<HashMap<String, f64>> = token_lists
205 .iter()
206 .map(|toks| term_frequencies(toks))
207 .collect();
208 let vocab = build_vocab(&tf_maps);
209 if vocab.is_empty() {
210 vec![vec![1.0_f64]; n]
212 } else {
213 tf_maps
214 .iter()
215 .map(|tf| {
216 vocab
217 .iter()
218 .map(|term| tf.get(term).copied().unwrap_or(0.0))
219 .collect()
220 })
221 .collect()
222 }
223 }
224 };
225
226 let batch_us = batch_start.elapsed().as_micros() as u64;
227 stage_state.record_batch(n, batch_us);
228
229 Ok(EmbeddingBatch {
230 ids,
231 texts: Some(texts),
232 raw_embeddings: None,
233 output_embeddings,
234 processing_time_us: batch_us,
235 })
236 }
237
238 pub fn process_embeddings(
240 &self,
241 ids: Vec<String>,
242 embeddings: Vec<Vec<f64>>,
243 ) -> Result<EmbeddingBatch, EpmPipelineError> {
244 if ids.is_empty() || embeddings.is_empty() {
245 return Err(EpmPipelineError::EmptyInput);
246 }
247 if ids.len() != embeddings.len() {
248 return Err(EpmPipelineError::InvalidConfig(format!(
249 "ids.len() ({}) != embeddings.len() ({})",
250 ids.len(),
251 embeddings.len()
252 )));
253 }
254
255 let batch_start = Instant::now();
256 let n = embeddings.len();
257 let raw = embeddings.clone();
258
259 let mut embs = embeddings;
260
261 let mut stage_state = self
262 .state
263 .lock()
264 .map_err(|e| EpmPipelineError::ProcessingFailed(format!("mutex poisoned: {e}")))?;
265
266 for stage in &self.config.stages {
267 if stage.requires_tokens() {
268 continue;
270 }
271 let stage_start = Instant::now();
272
273 match stage {
274 EpmPipelineStage::L2Normalize => {
275 for v in embs.iter_mut() {
276 l2_normalize(v);
277 }
278 }
279 EpmPipelineStage::DimensionReduce { target_dim, method } => {
280 let stats = CorpusStats::from_embeddings(&embs);
281 let reduced: Vec<Vec<f64>> = embs
282 .iter()
283 .map(|v| reduce_dimensions(v, *target_dim, method, stats.as_ref()))
284 .collect();
285 embs = reduced;
286 }
287 EpmPipelineStage::QuantizeToByte => {
288 for v in embs.iter_mut() {
289 *v = quantize_to_byte(v);
290 }
291 }
292 EpmPipelineStage::AddPositionalEncoding { max_len } => {
293 for (pos, v) in embs.iter_mut().enumerate() {
294 add_positional_encoding(v, pos, *max_len);
295 }
296 }
297 _ => {} }
299
300 let elapsed_us = stage_start.elapsed().as_micros() as u64;
301 stage_state.record_stage(stage.name(), elapsed_us, n);
302 }
303
304 let batch_us = batch_start.elapsed().as_micros() as u64;
305 stage_state.record_batch(n, batch_us);
306
307 Ok(EmbeddingBatch {
308 ids,
309 texts: None,
310 raw_embeddings: Some(raw),
311 output_embeddings: embs,
312 processing_time_us: batch_us,
313 })
314 }
315
316 pub fn add_stage(&mut self, stage: EpmPipelineStage) -> Result<(), EpmPipelineError> {
318 self.config.stages.push(stage);
319 self.validate_config()
320 }
321
322 pub fn remove_stage(&mut self, index: usize) -> Result<(), EpmPipelineError> {
324 if index >= self.config.stages.len() {
325 return Err(EpmPipelineError::InvalidConfig(format!(
326 "stage index {index} out of range (pipeline has {} stages)",
327 self.config.stages.len()
328 )));
329 }
330 self.config.stages.remove(index);
331 Ok(())
332 }
333
334 pub fn validate_config(&self) -> Result<(), EpmPipelineError> {
342 let mut seen_tokenize = false;
343 for stage in &self.config.stages {
344 match stage {
345 EpmPipelineStage::Tokenize { .. } => {
346 seen_tokenize = true;
347 }
348 EpmPipelineStage::NGram { n } => {
349 if *n == 0 {
350 return Err(EpmPipelineError::InvalidConfig(
351 "NGram n must be >= 1".to_string(),
352 ));
353 }
354 }
355 EpmPipelineStage::StopWordFilter(_) => {
356 }
358 EpmPipelineStage::TfIdfWeighting => {
359 if !seen_tokenize {
360 return Err(EpmPipelineError::InvalidConfig(
361 "TfIdfWeighting must be preceded by a Tokenize stage".to_string(),
362 ));
363 }
364 }
365 EpmPipelineStage::DimensionReduce { target_dim, .. } => {
366 if *target_dim == 0 {
367 return Err(EpmPipelineError::InvalidConfig(
368 "DimensionReduce target_dim must be > 0".to_string(),
369 ));
370 }
371 }
372 EpmPipelineStage::L2Normalize
373 | EpmPipelineStage::QuantizeToByte
374 | EpmPipelineStage::AddPositionalEncoding { .. } => {}
375 }
376 }
377 if self.config.output_dim == 0 {
378 return Err(EpmPipelineError::InvalidConfig(
379 "output_dim must be > 0".to_string(),
380 ));
381 }
382 Ok(())
383 }
384
385 pub fn benchmark(&self, texts: &[String], n_runs: usize) -> Vec<StageTiming> {
388 if texts.is_empty() || n_runs == 0 {
389 return vec![];
390 }
391 let mut timings: Vec<StageTiming> = Vec::new();
392
393 for stage in &self.config.stages {
394 let name = stage.name().to_string();
395 let mut total_us: u64 = 0;
396
397 let token_lists: Vec<Vec<String>> =
399 texts.iter().map(|t| tokenize_text(t, true, true)).collect();
400 let tf_maps: Vec<HashMap<String, f64>> = token_lists
401 .iter()
402 .map(|toks| term_frequencies(toks))
403 .collect();
404 let vocab = build_vocab(&tf_maps);
405 let idf = inverse_document_frequencies(&token_lists);
406 let base_embeddings: Vec<Vec<f64>> = tf_maps
407 .iter()
408 .map(|tf| epm_processing::tfidf_vector(tf, &idf, &vocab))
409 .collect();
410
411 for _ in 0..n_runs {
412 let start = Instant::now();
413 match stage {
414 EpmPipelineStage::Tokenize {
415 lowercase,
416 strip_punct,
417 } => {
418 for t in texts {
419 let _ = tokenize_text(t, *lowercase, *strip_punct);
420 }
421 }
422 EpmPipelineStage::StopWordFilter(sw) => {
423 for toks in &token_lists {
424 let _ = apply_stop_word_filter(toks.clone(), sw);
425 }
426 }
427 EpmPipelineStage::NGram { n } => {
428 for toks in &token_lists {
429 let _ = apply_ngram(toks, *n);
430 }
431 }
432 EpmPipelineStage::TfIdfWeighting => {
433 let corpus_tokens: Vec<Vec<String>> =
434 texts.iter().map(|t| tokenize_text(t, true, true)).collect();
435 let idf_b = inverse_document_frequencies(&corpus_tokens);
436 let tf_b: Vec<HashMap<String, f64>> = token_lists
437 .iter()
438 .map(|toks| term_frequencies(toks))
439 .collect();
440 let vocab_b = build_vocab(&tf_b);
441 for tf in &tf_b {
442 let _ = epm_processing::tfidf_vector(tf, &idf_b, &vocab_b);
443 }
444 }
445 EpmPipelineStage::L2Normalize => {
446 let mut embs = base_embeddings.clone();
447 for v in embs.iter_mut() {
448 l2_normalize(v);
449 }
450 }
451 EpmPipelineStage::DimensionReduce { target_dim, method } => {
452 let stats = CorpusStats::from_embeddings(&base_embeddings);
453 for v in &base_embeddings {
454 let _ = reduce_dimensions(v, *target_dim, method, stats.as_ref());
455 }
456 }
457 EpmPipelineStage::QuantizeToByte => {
458 for v in &base_embeddings {
459 let _ = quantize_to_byte(v);
460 }
461 }
462 EpmPipelineStage::AddPositionalEncoding { max_len } => {
463 let mut embs = base_embeddings.clone();
464 for (pos, v) in embs.iter_mut().enumerate() {
465 add_positional_encoding(v, pos, *max_len);
466 }
467 }
468 }
469 total_us += start.elapsed().as_micros() as u64;
470 }
471
472 let avg_time_us = total_us as f64 / n_runs as f64;
473 timings.push(StageTiming {
474 stage_name: name,
475 avg_time_us,
476 total_processed: (texts.len() * n_runs) as u64,
477 });
478 }
479 timings
480 }
481
482 pub fn stats(&self) -> EpmPipelineStats {
484 let state = match self.state.lock() {
485 Ok(s) => s,
486 Err(e) => e.into_inner(),
487 };
488 let stage_timings = state
489 .stage_time
490 .iter()
491 .map(|(name, (total_us, total_processed))| StageTiming {
492 stage_name: name.clone(),
493 avg_time_us: if *total_processed > 0 {
494 *total_us as f64 / *total_processed as f64
495 } else {
496 0.0
497 },
498 total_processed: *total_processed,
499 })
500 .collect();
501
502 EpmPipelineStats {
503 batches_processed: state.batches_processed,
504 total_inputs: state.total_inputs,
505 avg_batch_time_us: state.avg_batch_time_us(),
506 stage_timings,
507 output_dim: self.config.output_dim,
508 }
509 }
510
511 pub fn config(&self) -> &EpmPipelineConfig {
513 &self.config
514 }
515}
516
517#[cfg(test)]
522mod tests {
523 use super::epm_processing::{
524 add_positional_encoding, apply_ngram, apply_stop_word_filter, quantize_to_byte,
525 tokenize_text, xorshift64,
526 };
527 use super::*;
528
529 fn make_manager(stages: Vec<EpmPipelineStage>) -> EmbeddingPipelineManager {
534 let mut config = EpmPipelineConfig::new(32, 4);
535 config.stages = stages;
536 EmbeddingPipelineManager::new(config).expect("valid config")
537 }
538
539 fn text_ids(n: usize) -> Vec<String> {
540 (0..n).map(|i| format!("doc{i}")).collect()
541 }
542
543 fn sample_texts() -> Vec<String> {
544 vec![
545 "the quick brown fox jumps over the lazy dog".to_string(),
546 "a fast red cat leaps over a sleepy hound".to_string(),
547 "rust programming language is fast and safe".to_string(),
548 ]
549 }
550
551 fn sample_embeddings(n: usize, dim: usize) -> Vec<Vec<f64>> {
552 (0..n)
553 .map(|i| (0..dim).map(|j| (i * dim + j) as f64 / 100.0).collect())
554 .collect()
555 }
556
557 #[test]
562 fn test_xorshift64_nonzero() {
563 let mut state: u64 = 42;
564 let v = xorshift64(&mut state);
565 assert_ne!(v, 42);
566 }
567
568 #[test]
569 fn test_xorshift64_sequence_differs() {
570 let mut state: u64 = 1;
571 let a = xorshift64(&mut state);
572 let b = xorshift64(&mut state);
573 assert_ne!(a, b);
574 }
575
576 #[test]
577 fn test_xorshift64_deterministic() {
578 let mut s1 = 99u64;
579 let mut s2 = 99u64;
580 let a: Vec<u64> = (0..10).map(|_| xorshift64(&mut s1)).collect();
581 let b: Vec<u64> = (0..10).map(|_| xorshift64(&mut s2)).collect();
582 assert_eq!(a, b);
583 }
584
585 #[test]
590 fn test_l2_normalize_unit_result() {
591 let mut v = vec![3.0_f64, 4.0];
592 l2_normalize(&mut v);
593 let norm: f64 = v.iter().map(|x| x * x).sum::<f64>().sqrt();
594 assert!((norm - 1.0).abs() < 1e-10);
595 }
596
597 #[test]
598 fn test_l2_normalize_zero_vector() {
599 let mut v = vec![0.0_f64, 0.0, 0.0];
600 l2_normalize(&mut v);
601 assert_eq!(v, vec![0.0, 0.0, 0.0]);
602 }
603
604 #[test]
605 fn test_l2_normalize_single_element() {
606 let mut v = vec![5.0_f64];
607 l2_normalize(&mut v);
608 assert!((v[0] - 1.0).abs() < 1e-10);
609 }
610
611 #[test]
616 fn test_mean_pool_empty() {
617 assert_eq!(mean_pool(&[]), Vec::<f64>::new());
618 }
619
620 #[test]
621 fn test_mean_pool_single() {
622 let v = vec![1.0, 2.0, 3.0];
623 let result = mean_pool(std::slice::from_ref(&v));
624 assert_eq!(result, v);
625 }
626
627 #[test]
628 fn test_mean_pool_two_vectors() {
629 let a = vec![1.0_f64, 2.0];
630 let b = vec![3.0_f64, 4.0];
631 let result = mean_pool(&[a, b]);
632 assert!((result[0] - 2.0).abs() < 1e-10);
633 assert!((result[1] - 3.0).abs() < 1e-10);
634 }
635
636 #[test]
641 fn test_random_projection_output_dim() {
642 let v: Vec<f64> = (0..128).map(|i| i as f64).collect();
643 let out = random_projection(&v, 32, 42);
644 assert_eq!(out.len(), 32);
645 }
646
647 #[test]
648 fn test_random_projection_deterministic() {
649 let v: Vec<f64> = (0..64).map(|i| i as f64).collect();
650 let a = random_projection(&v, 16, 7);
651 let b = random_projection(&v, 16, 7);
652 assert_eq!(a, b);
653 }
654
655 #[test]
656 fn test_random_projection_different_seeds() {
657 let v: Vec<f64> = (0..64).map(|i| i as f64 / 64.0).collect();
658 let a = random_projection(&v, 16, 1);
659 let b = random_projection(&v, 16, 2);
660 assert_ne!(a, b);
662 }
663
664 #[test]
665 fn test_random_projection_zero_target() {
666 let v = vec![1.0_f64, 2.0];
667 let out = random_projection(&v, 0, 1);
668 assert_eq!(out.len(), 0);
669 }
670
671 #[test]
676 fn test_stage_tokenize_lowercase() {
677 let mgr = make_manager(vec![EpmPipelineStage::Tokenize {
678 lowercase: true,
679 strip_punct: false,
680 }]);
681 let batch = mgr
682 .process_text(text_ids(1), vec!["Hello World".to_string()], None)
683 .expect("test: tokenize lowercase stage");
684 assert_eq!(batch.output_embeddings.len(), 1);
685 }
686
687 #[test]
688 fn test_stage_tokenize_strip_punct() {
689 let tokens = tokenize_text("hello, world!", false, true);
690 assert!(tokens.contains(&"hello".to_string()));
691 assert!(tokens.contains(&"world".to_string()));
692 assert!(!tokens.iter().any(|t| t.contains(',')));
693 }
694
695 #[test]
696 fn test_stage_tokenize_no_lowercase() {
697 let tokens = tokenize_text("Hello World", false, false);
698 assert!(tokens.contains(&"Hello".to_string()));
699 }
700
701 #[test]
706 fn test_stage_stop_word_filter_removes_words() {
707 let stop_words = vec!["the".to_string(), "a".to_string(), "an".to_string()];
708 let tokens = vec!["the".to_string(), "quick".to_string(), "fox".to_string()];
709 let filtered = apply_stop_word_filter(tokens, &stop_words);
710 assert!(!filtered.contains(&"the".to_string()));
711 assert!(filtered.contains(&"quick".to_string()));
712 }
713
714 #[test]
715 fn test_stage_stop_word_filter_pipeline() {
716 let stop_words = vec!["the".to_string(), "over".to_string()];
717 let mgr = make_manager(vec![
718 EpmPipelineStage::Tokenize {
719 lowercase: true,
720 strip_punct: false,
721 },
722 EpmPipelineStage::StopWordFilter(stop_words),
723 ]);
724 let batch = mgr
725 .process_text(text_ids(1), vec!["the fox jumps over".to_string()], None)
726 .expect("test: stop word filter pipeline");
727 assert_eq!(batch.output_embeddings.len(), 1);
728 }
729
730 #[test]
735 fn test_ngram_bigrams() {
736 let tokens = vec!["a".to_string(), "b".to_string(), "c".to_string()];
737 let bigrams = apply_ngram(&tokens, 2);
738 assert_eq!(bigrams, vec!["a_b", "b_c"]);
739 }
740
741 #[test]
742 fn test_ngram_trigrams() {
743 let tokens: Vec<String> = vec!["a", "b", "c", "d"]
744 .into_iter()
745 .map(String::from)
746 .collect();
747 let trigrams = apply_ngram(&tokens, 3);
748 assert_eq!(trigrams, vec!["a_b_c", "b_c_d"]);
749 }
750
751 #[test]
752 fn test_ngram_unigram_passthrough() {
753 let tokens: Vec<String> = vec!["a", "b"].into_iter().map(String::from).collect();
754 let result = apply_ngram(&tokens, 1);
755 assert_eq!(result, tokens);
756 }
757
758 #[test]
759 fn test_ngram_too_few_tokens() {
760 let tokens = vec!["only".to_string()];
761 let result = apply_ngram(&tokens, 3);
762 assert_eq!(result, tokens);
764 }
765
766 #[test]
767 fn test_ngram_stage_pipeline() {
768 let mgr = make_manager(vec![
769 EpmPipelineStage::Tokenize {
770 lowercase: true,
771 strip_punct: false,
772 },
773 EpmPipelineStage::NGram { n: 2 },
774 ]);
775 let batch = mgr
776 .process_text(
777 text_ids(1),
778 vec!["alpha beta gamma delta".to_string()],
779 None,
780 )
781 .expect("test: ngram stage pipeline");
782 assert_eq!(batch.output_embeddings.len(), 1);
783 }
784
785 #[test]
790 fn test_tfidf_weighting_output_shape() {
791 let mgr = make_manager(vec![
792 EpmPipelineStage::Tokenize {
793 lowercase: true,
794 strip_punct: true,
795 },
796 EpmPipelineStage::TfIdfWeighting,
797 ]);
798 let texts = sample_texts();
799 let n = texts.len();
800 let batch = mgr
801 .process_text(text_ids(n), texts, None)
802 .expect("test: tfidf weighting output shape");
803 assert_eq!(batch.output_embeddings.len(), n);
804 let dim0 = batch.output_embeddings[0].len();
806 for v in &batch.output_embeddings {
807 assert_eq!(v.len(), dim0);
808 }
809 }
810
811 #[test]
812 fn test_tfidf_nonnegative_values() {
813 let mgr = make_manager(vec![
814 EpmPipelineStage::Tokenize {
815 lowercase: true,
816 strip_punct: true,
817 },
818 EpmPipelineStage::TfIdfWeighting,
819 ]);
820 let texts = sample_texts();
821 let n = texts.len();
822 let batch = mgr
823 .process_text(text_ids(n), texts, None)
824 .expect("test: tfidf nonnegative values");
825 for v in &batch.output_embeddings {
826 for &x in v {
827 assert!(x >= 0.0, "TF-IDF value should be non-negative");
828 }
829 }
830 }
831
832 #[test]
833 fn test_tfidf_with_external_corpus() {
834 let corpus = vec!["rust language".to_string(), "python language".to_string()];
835 let mgr = make_manager(vec![
836 EpmPipelineStage::Tokenize {
837 lowercase: true,
838 strip_punct: false,
839 },
840 EpmPipelineStage::TfIdfWeighting,
841 ]);
842 let batch = mgr
843 .process_text(
844 text_ids(1),
845 vec!["rust is great".to_string()],
846 Some(&corpus),
847 )
848 .expect("test: tfidf with external corpus");
849 assert_eq!(batch.output_embeddings.len(), 1);
850 }
851
852 #[test]
857 fn test_pipeline_l2_normalize() {
858 let mgr = make_manager(vec![
859 EpmPipelineStage::Tokenize {
860 lowercase: true,
861 strip_punct: true,
862 },
863 EpmPipelineStage::TfIdfWeighting,
864 EpmPipelineStage::L2Normalize,
865 ]);
866 let texts = sample_texts();
867 let n = texts.len();
868 let batch = mgr
869 .process_text(text_ids(n), texts, None)
870 .expect("test: pipeline l2 normalize");
871 for v in &batch.output_embeddings {
872 let norm: f64 = v.iter().map(|x| x * x).sum::<f64>().sqrt();
873 assert!((norm - 1.0).abs() < 1e-9 || norm < 1e-10, "norm={norm}");
874 }
875 }
876
877 #[test]
882 fn test_dimension_reduce_truncate() {
883 let mgr = make_manager(vec![
884 EpmPipelineStage::Tokenize {
885 lowercase: true,
886 strip_punct: true,
887 },
888 EpmPipelineStage::TfIdfWeighting,
889 EpmPipelineStage::DimensionReduce {
890 target_dim: 4,
891 method: EpmReductionMethod::TruncateDims,
892 },
893 ]);
894 let texts = sample_texts();
895 let n = texts.len();
896 let batch = mgr
897 .process_text(text_ids(n), texts, None)
898 .expect("test: dimension reduce truncate");
899 for v in &batch.output_embeddings {
900 assert_eq!(v.len(), 4);
901 }
902 }
903
904 #[test]
905 fn test_dimension_reduce_random_projection() {
906 let mgr = make_manager(vec![
907 EpmPipelineStage::Tokenize {
908 lowercase: true,
909 strip_punct: true,
910 },
911 EpmPipelineStage::TfIdfWeighting,
912 EpmPipelineStage::DimensionReduce {
913 target_dim: 8,
914 method: EpmReductionMethod::RandomProjection(42),
915 },
916 ]);
917 let texts = sample_texts();
918 let n = texts.len();
919 let batch = mgr
920 .process_text(text_ids(n), texts, None)
921 .expect("test: dimension reduce random projection");
922 for v in &batch.output_embeddings {
923 assert_eq!(v.len(), 8);
924 }
925 }
926
927 #[test]
928 fn test_dimension_reduce_mean_pooling() {
929 let mgr = make_manager(vec![
930 EpmPipelineStage::Tokenize {
931 lowercase: true,
932 strip_punct: true,
933 },
934 EpmPipelineStage::TfIdfWeighting,
935 EpmPipelineStage::DimensionReduce {
936 target_dim: 4,
937 method: EpmReductionMethod::MeanPooling,
938 },
939 ]);
940 let texts = sample_texts();
941 let n = texts.len();
942 let batch = mgr
943 .process_text(text_ids(n), texts, None)
944 .expect("test: dimension reduce mean pooling");
945 for v in &batch.output_embeddings {
946 assert_eq!(v.len(), 4);
947 }
948 }
949
950 #[test]
951 fn test_dimension_reduce_pca() {
952 let mgr = make_manager(vec![
953 EpmPipelineStage::Tokenize {
954 lowercase: true,
955 strip_punct: true,
956 },
957 EpmPipelineStage::TfIdfWeighting,
958 EpmPipelineStage::DimensionReduce {
959 target_dim: 4,
960 method: EpmReductionMethod::PCA,
961 },
962 ]);
963 let texts = sample_texts();
964 let n = texts.len();
965 let batch = mgr
966 .process_text(text_ids(n), texts, None)
967 .expect("test: dimension reduce pca");
968 for v in &batch.output_embeddings {
969 assert_eq!(v.len(), 4);
970 }
971 }
972
973 #[test]
978 fn test_quantize_to_byte_range() {
979 let v = vec![0.1_f64, 0.5, 1.0, -0.5, 2.0];
980 let q = quantize_to_byte(&v);
981 assert_eq!(q.len(), v.len());
982 for &val in &q {
983 assert!(
984 (0.0..=255.0).contains(&val),
985 "quantized value {val} out of [0,255]"
986 );
987 }
988 }
989
990 #[test]
991 fn test_quantize_to_byte_constant_vector() {
992 let v = vec![3.0_f64; 8];
994 let q = quantize_to_byte(&v);
995 assert!(q.iter().all(|&x| x == 0.0));
996 }
997
998 #[test]
999 fn test_quantize_pipeline_stage() {
1000 let mgr = make_manager(vec![
1001 EpmPipelineStage::Tokenize {
1002 lowercase: true,
1003 strip_punct: true,
1004 },
1005 EpmPipelineStage::TfIdfWeighting,
1006 EpmPipelineStage::QuantizeToByte,
1007 ]);
1008 let texts = sample_texts();
1009 let n = texts.len();
1010 let batch = mgr
1011 .process_text(text_ids(n), texts, None)
1012 .expect("test: quantize pipeline stage");
1013 for v in &batch.output_embeddings {
1014 for &val in v {
1015 assert!((0.0..=255.0).contains(&val));
1016 }
1017 }
1018 }
1019
1020 #[test]
1025 fn test_positional_encoding_changes_vector() {
1026 let mut v = vec![1.0_f64; 16];
1027 let original = v.clone();
1028 add_positional_encoding(&mut v, 0, 512);
1029 assert_ne!(v, original);
1030 }
1031
1032 #[test]
1033 fn test_positional_encoding_different_positions() {
1034 let mut a = vec![0.0_f64; 8];
1035 let mut b = vec![0.0_f64; 8];
1036 add_positional_encoding(&mut a, 0, 100);
1037 add_positional_encoding(&mut b, 1, 100);
1038 assert_ne!(a, b);
1039 }
1040
1041 #[test]
1042 fn test_positional_encoding_pipeline_stage() {
1043 let mgr = make_manager(vec![
1044 EpmPipelineStage::Tokenize {
1045 lowercase: true,
1046 strip_punct: true,
1047 },
1048 EpmPipelineStage::TfIdfWeighting,
1049 EpmPipelineStage::AddPositionalEncoding { max_len: 128 },
1050 ]);
1051 let texts = sample_texts();
1052 let n = texts.len();
1053 let batch = mgr
1054 .process_text(text_ids(n), texts, None)
1055 .expect("test: process text positional encoding stage");
1056 assert_eq!(batch.output_embeddings.len(), n);
1057 }
1058
1059 #[test]
1064 fn test_process_text_roundtrip() {
1065 let mgr = make_manager(vec![
1066 EpmPipelineStage::Tokenize {
1067 lowercase: true,
1068 strip_punct: true,
1069 },
1070 EpmPipelineStage::StopWordFilter(vec!["the".to_string(), "a".to_string()]),
1071 EpmPipelineStage::NGram { n: 2 },
1072 EpmPipelineStage::TfIdfWeighting,
1073 EpmPipelineStage::L2Normalize,
1074 EpmPipelineStage::DimensionReduce {
1075 target_dim: 8,
1076 method: EpmReductionMethod::RandomProjection(1337),
1077 },
1078 ]);
1079 let texts = sample_texts();
1080 let n = texts.len();
1081 let batch = mgr
1082 .process_text(text_ids(n), texts.clone(), None)
1083 .expect("test: process text roundtrip");
1084 assert_eq!(batch.ids.len(), n);
1085 assert_eq!(batch.output_embeddings.len(), n);
1086 assert!(batch.texts.is_some());
1087 assert_eq!(batch.processing_time_us, batch.processing_time_us); }
1089
1090 #[test]
1091 fn test_process_text_preserves_ids() {
1092 let ids = vec!["foo".to_string(), "bar".to_string(), "baz".to_string()];
1093 let texts = sample_texts();
1094 let mgr = make_manager(vec![EpmPipelineStage::Tokenize {
1095 lowercase: true,
1096 strip_punct: false,
1097 }]);
1098 let batch = mgr
1099 .process_text(ids.clone(), texts, None)
1100 .expect("test: process text preserves ids");
1101 assert_eq!(batch.ids, ids);
1102 }
1103
1104 #[test]
1105 fn test_process_text_empty_error() {
1106 let mgr = make_manager(vec![]);
1107 let result = mgr.process_text(vec![], vec![], None);
1108 assert!(matches!(result, Err(EpmPipelineError::EmptyInput)));
1109 }
1110
1111 #[test]
1112 fn test_process_text_mismatched_ids_error() {
1113 let mgr = make_manager(vec![EpmPipelineStage::Tokenize {
1114 lowercase: false,
1115 strip_punct: false,
1116 }]);
1117 let result = mgr.process_text(
1118 vec!["a".to_string()],
1119 vec!["hello".to_string(), "world".to_string()],
1120 None,
1121 );
1122 assert!(matches!(result, Err(EpmPipelineError::InvalidConfig(_))));
1123 }
1124
1125 #[test]
1130 fn test_process_embeddings_roundtrip() {
1131 let mgr = make_manager(vec![
1132 EpmPipelineStage::L2Normalize,
1133 EpmPipelineStage::DimensionReduce {
1134 target_dim: 4,
1135 method: EpmReductionMethod::TruncateDims,
1136 },
1137 ]);
1138 let embs = sample_embeddings(3, 16);
1139 let batch = mgr
1140 .process_embeddings(text_ids(3), embs.clone())
1141 .expect("test: process embeddings roundtrip");
1142 assert_eq!(batch.ids.len(), 3);
1143 assert_eq!(batch.output_embeddings.len(), 3);
1144 for v in &batch.output_embeddings {
1145 assert_eq!(v.len(), 4);
1146 }
1147 assert!(batch.raw_embeddings.is_some());
1148 assert!(batch.texts.is_none());
1149 }
1150
1151 #[test]
1152 fn test_process_embeddings_empty_error() {
1153 let mgr = make_manager(vec![]);
1154 let result = mgr.process_embeddings(vec![], vec![]);
1155 assert!(matches!(result, Err(EpmPipelineError::EmptyInput)));
1156 }
1157
1158 #[test]
1159 fn test_process_embeddings_skips_text_stages() {
1160 let mgr = make_manager(vec![
1162 EpmPipelineStage::Tokenize {
1163 lowercase: true,
1164 strip_punct: true,
1165 },
1166 EpmPipelineStage::L2Normalize,
1167 ]);
1168 let embs = sample_embeddings(2, 8);
1169 let batch = mgr
1170 .process_embeddings(text_ids(2), embs)
1171 .expect("test: process embeddings skips text stages");
1172 assert_eq!(batch.output_embeddings.len(), 2);
1173 }
1174
1175 #[test]
1176 fn test_process_embeddings_l2_unit_norm() {
1177 let mgr = make_manager(vec![EpmPipelineStage::L2Normalize]);
1178 let embs = sample_embeddings(4, 8);
1179 let batch = mgr
1180 .process_embeddings(text_ids(4), embs)
1181 .expect("test: process embeddings l2 unit norm");
1182 for v in &batch.output_embeddings {
1183 let norm: f64 = v.iter().map(|x| x * x).sum::<f64>().sqrt();
1184 assert!((norm - 1.0).abs() < 1e-9 || norm < 1e-10);
1185 }
1186 }
1187
1188 #[test]
1193 fn test_add_stage_appends() {
1194 let mut mgr = make_manager(vec![]);
1195 mgr.add_stage(EpmPipelineStage::L2Normalize)
1196 .expect("test: add stage");
1197 assert_eq!(mgr.config().stages.len(), 1);
1198 }
1199
1200 #[test]
1201 fn test_remove_stage_removes() {
1202 let mut mgr = make_manager(vec![
1203 EpmPipelineStage::L2Normalize,
1204 EpmPipelineStage::QuantizeToByte,
1205 ]);
1206 mgr.remove_stage(0).expect("test: remove stage");
1207 assert_eq!(mgr.config().stages.len(), 1);
1208 assert!(matches!(
1209 mgr.config().stages[0],
1210 EpmPipelineStage::QuantizeToByte
1211 ));
1212 }
1213
1214 #[test]
1215 fn test_remove_stage_out_of_bounds() {
1216 let mut mgr = make_manager(vec![]);
1217 let result = mgr.remove_stage(5);
1218 assert!(matches!(result, Err(EpmPipelineError::InvalidConfig(_))));
1219 }
1220
1221 #[test]
1226 fn test_validate_tfidf_without_tokenize_fails() {
1227 let config = EpmPipelineConfig {
1228 stages: vec![EpmPipelineStage::TfIdfWeighting],
1229 output_dim: 32,
1230 batch_size: 4,
1231 };
1232 let result = EmbeddingPipelineManager::new(config);
1233 assert!(matches!(result, Err(EpmPipelineError::InvalidConfig(_))));
1234 }
1235
1236 #[test]
1237 fn test_validate_zero_output_dim_fails() {
1238 let config = EpmPipelineConfig {
1239 stages: vec![],
1240 output_dim: 0,
1241 batch_size: 4,
1242 };
1243 let result = EmbeddingPipelineManager::new(config);
1244 assert!(matches!(result, Err(EpmPipelineError::InvalidConfig(_))));
1245 }
1246
1247 #[test]
1248 fn test_validate_zero_target_dim_fails() {
1249 let config = EpmPipelineConfig {
1250 stages: vec![EpmPipelineStage::DimensionReduce {
1251 target_dim: 0,
1252 method: EpmReductionMethod::TruncateDims,
1253 }],
1254 output_dim: 32,
1255 batch_size: 4,
1256 };
1257 let result = EmbeddingPipelineManager::new(config);
1258 assert!(matches!(result, Err(EpmPipelineError::InvalidConfig(_))));
1259 }
1260
1261 #[test]
1262 fn test_validate_zero_ngram_fails() {
1263 let config = EpmPipelineConfig {
1265 stages: vec![EpmPipelineStage::NGram { n: 0 }],
1266 output_dim: 32,
1267 batch_size: 4,
1268 };
1269 let result = EmbeddingPipelineManager::new(config);
1270 assert!(matches!(result, Err(EpmPipelineError::InvalidConfig(_))));
1271 }
1272
1273 #[test]
1274 fn test_validate_valid_config_ok() {
1275 let config = EpmPipelineConfig {
1276 stages: vec![
1277 EpmPipelineStage::Tokenize {
1278 lowercase: true,
1279 strip_punct: true,
1280 },
1281 EpmPipelineStage::TfIdfWeighting,
1282 EpmPipelineStage::L2Normalize,
1283 ],
1284 output_dim: 64,
1285 batch_size: 8,
1286 };
1287 assert!(EmbeddingPipelineManager::new(config).is_ok());
1288 }
1289
1290 #[test]
1295 fn test_benchmark_returns_timings() {
1296 let mgr = make_manager(vec![
1297 EpmPipelineStage::Tokenize {
1298 lowercase: true,
1299 strip_punct: true,
1300 },
1301 EpmPipelineStage::L2Normalize,
1302 ]);
1303 let texts = sample_texts();
1304 let timings = mgr.benchmark(&texts, 3);
1305 assert_eq!(timings.len(), 2);
1306 assert_eq!(timings[0].stage_name, "Tokenize");
1307 assert_eq!(timings[1].stage_name, "L2Normalize");
1308 }
1309
1310 #[test]
1311 fn test_benchmark_empty_returns_empty() {
1312 let mgr = make_manager(vec![EpmPipelineStage::L2Normalize]);
1313 let timings = mgr.benchmark(&[], 5);
1314 assert!(timings.is_empty());
1315 }
1316
1317 #[test]
1318 fn test_benchmark_zero_runs_returns_empty() {
1319 let mgr = make_manager(vec![EpmPipelineStage::L2Normalize]);
1320 let timings = mgr.benchmark(&sample_texts(), 0);
1321 assert!(timings.is_empty());
1322 }
1323
1324 #[test]
1325 fn test_benchmark_nonnegative_times() {
1326 let mgr = make_manager(vec![
1327 EpmPipelineStage::Tokenize {
1328 lowercase: true,
1329 strip_punct: false,
1330 },
1331 EpmPipelineStage::TfIdfWeighting,
1332 ]);
1333 let timings = mgr.benchmark(&sample_texts(), 2);
1334 for t in &timings {
1335 assert!(t.avg_time_us >= 0.0);
1336 }
1337 }
1338
1339 #[test]
1344 fn test_stats_initial_zero() {
1345 let mgr = make_manager(vec![]);
1346 let stats = mgr.stats();
1347 assert_eq!(stats.batches_processed, 0);
1348 assert_eq!(stats.total_inputs, 0);
1349 }
1350
1351 #[test]
1352 fn test_stats_increments_after_process() {
1353 let mgr = make_manager(vec![EpmPipelineStage::L2Normalize]);
1354 let embs = sample_embeddings(3, 8);
1355 mgr.process_embeddings(text_ids(3), embs)
1356 .expect("test: process embeddings stats increment");
1357 let stats = mgr.stats();
1358 assert_eq!(stats.batches_processed, 1);
1359 assert_eq!(stats.total_inputs, 3);
1360 }
1361
1362 #[test]
1363 fn test_stats_output_dim() {
1364 let mgr = make_manager(vec![]);
1365 assert_eq!(mgr.stats().output_dim, 32);
1366 }
1367
1368 #[test]
1369 fn test_stats_multiple_batches() {
1370 let mgr = make_manager(vec![EpmPipelineStage::L2Normalize]);
1371 for _ in 0..5 {
1372 let embs = sample_embeddings(2, 4);
1373 mgr.process_embeddings(text_ids(2), embs)
1374 .expect("test: process embeddings multiple batches");
1375 }
1376 let stats = mgr.stats();
1377 assert_eq!(stats.batches_processed, 5);
1378 assert_eq!(stats.total_inputs, 10);
1379 }
1380
1381 #[test]
1386 fn test_error_display_empty_input() {
1387 let e = EpmPipelineError::EmptyInput;
1388 assert_eq!(e.to_string(), "empty input");
1389 }
1390
1391 #[test]
1392 fn test_error_display_dimension() {
1393 let e = EpmPipelineError::DimensionError {
1394 expected: 128,
1395 got: 64,
1396 };
1397 assert!(e.to_string().contains("128"));
1398 assert!(e.to_string().contains("64"));
1399 }
1400
1401 #[test]
1402 fn test_error_display_stage() {
1403 let e = EpmPipelineError::StageError {
1404 stage: "TfIdfWeighting".to_string(),
1405 reason: "empty vocabulary".to_string(),
1406 };
1407 assert!(e.to_string().contains("TfIdfWeighting"));
1408 }
1409
1410 #[test]
1411 fn test_error_invalid_config_clone() {
1412 let e = EpmPipelineError::InvalidConfig("bad config".to_string());
1413 let cloned = e.clone();
1414 assert_eq!(e.to_string(), cloned.to_string());
1415 }
1416
1417 #[test]
1422 fn test_full_text_pipeline() {
1423 let config = EpmPipelineConfig {
1424 stages: vec![
1425 EpmPipelineStage::Tokenize {
1426 lowercase: true,
1427 strip_punct: true,
1428 },
1429 EpmPipelineStage::StopWordFilter(vec![
1430 "the".to_string(),
1431 "a".to_string(),
1432 "is".to_string(),
1433 ]),
1434 EpmPipelineStage::NGram { n: 2 },
1435 EpmPipelineStage::TfIdfWeighting,
1436 EpmPipelineStage::L2Normalize,
1437 EpmPipelineStage::DimensionReduce {
1438 target_dim: 16,
1439 method: EpmReductionMethod::RandomProjection(99),
1440 },
1441 EpmPipelineStage::QuantizeToByte,
1442 EpmPipelineStage::AddPositionalEncoding { max_len: 256 },
1443 ],
1444 output_dim: 16,
1445 batch_size: 8,
1446 };
1447 let mgr = EmbeddingPipelineManager::new(config)
1448 .expect("test: create pipeline manager for full text pipeline");
1449 let texts = vec![
1450 "the quick brown fox".to_string(),
1451 "rust is a systems language".to_string(),
1452 "semantic search embeddings".to_string(),
1453 "machine learning vectors".to_string(),
1454 ];
1455 let n = texts.len();
1456 let batch = mgr
1457 .process_text(text_ids(n), texts, None)
1458 .expect("test: process text full pipeline");
1459 assert_eq!(batch.output_embeddings.len(), n);
1460 for v in &batch.output_embeddings {
1461 assert_eq!(v.len(), 16);
1462 }
1463 }
1464
1465 #[test]
1466 fn test_full_embedding_pipeline() {
1467 let config = EpmPipelineConfig {
1468 stages: vec![
1469 EpmPipelineStage::L2Normalize,
1470 EpmPipelineStage::DimensionReduce {
1471 target_dim: 8,
1472 method: EpmReductionMethod::MeanPooling,
1473 },
1474 EpmPipelineStage::QuantizeToByte,
1475 EpmPipelineStage::AddPositionalEncoding { max_len: 64 },
1476 ],
1477 output_dim: 8,
1478 batch_size: 4,
1479 };
1480 let mgr = EmbeddingPipelineManager::new(config)
1481 .expect("test: create pipeline manager for full embedding pipeline");
1482 let embs = sample_embeddings(5, 32);
1483 let batch = mgr
1484 .process_embeddings(text_ids(5), embs)
1485 .expect("test: process embeddings full pipeline");
1486 assert_eq!(batch.output_embeddings.len(), 5);
1487 for v in &batch.output_embeddings {
1488 assert_eq!(v.len(), 8);
1489 }
1490 }
1491}