1use serde::{Deserialize, Serialize};
9use std::time::SystemTime;
10
11#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct ModelProvenance {
18 pub model: EmbeddingModel,
20 pub model_id: String,
22 pub hash: String,
24 pub loaded_at: SystemTime,
26 pub loaded_at_iso: String,
28}
29
30impl ModelProvenance {
31 pub fn new(model: EmbeddingModel, model_id: String) -> Self {
33 let loaded_at = SystemTime::now();
34 let loaded_at_iso = {
35 let dt: chrono::DateTime<chrono::Utc> = loaded_at.into();
36 dt.to_rfc3339()
37 };
38
39 let hash_input = format!("{model_id}:{loaded_at_iso}:{model:?}");
40 let hash = blake3::hash(hash_input.as_bytes()).to_hex().to_string();
41
42 Self {
43 model,
44 model_id,
45 hash,
46 loaded_at,
47 loaded_at_iso,
48 }
49 }
50
51 pub fn dimensions(&self) -> usize {
53 self.model.dimensions()
54 }
55
56 pub fn matches_model(&self, expected: EmbeddingModel) -> bool {
58 self.model == expected
59 }
60}
61
62#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
68#[serde(rename_all = "snake_case")]
69#[non_exhaustive]
70pub enum EmbeddingModel {
71 #[default]
73 #[serde(alias = "BgeSmallEnV15")]
74 BgeSmallEnV15,
75
76 #[serde(alias = "BgeBaseEnV15")]
78 BgeBaseEnV15,
79
80 #[serde(alias = "BgeLargeEnV15")]
82 BgeLargeEnV15,
83
84 #[serde(alias = "MultilingualE5Small")]
86 MultilingualE5Small,
87
88 #[serde(alias = "MultilingualE5Base")]
90 MultilingualE5Base,
91
92 #[serde(alias = "Qwen3Embedding0_6B")]
94 Qwen3Embedding0_6B,
95
96 #[serde(alias = "Qwen3Embedding4B")]
98 Qwen3Embedding4B,
99
100 #[serde(alias = "AllMiniLmL6V2")]
102 AllMiniLmL6V2,
103
104 #[serde(alias = "ParaphraseMultilingualMiniLmL12V2")]
106 ParaphraseMultilingualMiniLmL12V2,
107
108 #[serde(alias = "TextEmbedding3Small")]
110 TextEmbedding3Small,
111}
112
113impl EmbeddingModel {
114 #[inline]
119 pub const fn native_dimensions(&self) -> usize {
120 match self {
121 EmbeddingModel::BgeSmallEnV15
122 | EmbeddingModel::MultilingualE5Small
123 | EmbeddingModel::AllMiniLmL6V2
124 | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => 384,
125 EmbeddingModel::BgeBaseEnV15 | EmbeddingModel::MultilingualE5Base => 768,
126 EmbeddingModel::BgeLargeEnV15 | EmbeddingModel::Qwen3Embedding0_6B => 1024,
127 EmbeddingModel::Qwen3Embedding4B => 2560,
128 EmbeddingModel::TextEmbedding3Small => 1536,
129 }
130 }
131
132 #[inline]
136 pub const fn dimensions(&self) -> usize {
137 self.native_dimensions()
138 }
139
140 #[inline]
142 pub const fn is_local(&self) -> bool {
143 matches!(
144 self,
145 EmbeddingModel::BgeSmallEnV15
146 | EmbeddingModel::BgeBaseEnV15
147 | EmbeddingModel::BgeLargeEnV15
148 | EmbeddingModel::MultilingualE5Small
149 | EmbeddingModel::MultilingualE5Base
150 | EmbeddingModel::AllMiniLmL6V2
151 | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2
152 | EmbeddingModel::Qwen3Embedding0_6B
153 | EmbeddingModel::Qwen3Embedding4B
154 )
155 }
156
157 #[inline]
159 pub const fn is_remote(&self) -> bool {
160 matches!(self, EmbeddingModel::TextEmbedding3Small)
161 }
162
163 #[inline]
167 pub const fn max_input_tokens(&self) -> usize {
168 match self {
169 EmbeddingModel::BgeSmallEnV15 => 512,
170 EmbeddingModel::BgeBaseEnV15 => 512,
171 EmbeddingModel::BgeLargeEnV15 => 512,
172 EmbeddingModel::MultilingualE5Small => 512,
173 EmbeddingModel::MultilingualE5Base => 512,
174 EmbeddingModel::AllMiniLmL6V2 => 256,
175 EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => 128,
176 EmbeddingModel::Qwen3Embedding0_6B => 8192,
178 EmbeddingModel::Qwen3Embedding4B => 8192,
179 EmbeddingModel::TextEmbedding3Small => 8191,
180 }
181 }
182
183 #[inline]
187 pub const fn query_instruction(&self) -> Option<&'static str> {
188 match self {
189 EmbeddingModel::MultilingualE5Small | EmbeddingModel::MultilingualE5Base => {
190 Some("query: ")
191 }
192 EmbeddingModel::Qwen3Embedding0_6B | EmbeddingModel::Qwen3Embedding4B => Some(
193 "Instruct: Given a web search query, retrieve relevant passages that answer the query\nQuery: ",
194 ),
195 EmbeddingModel::BgeSmallEnV15
196 | EmbeddingModel::BgeBaseEnV15
197 | EmbeddingModel::BgeLargeEnV15 => {
198 Some("Represent this sentence for searching relevant passages: ")
199 }
200 _ => None,
201 }
202 }
203
204 #[inline]
208 pub const fn document_instruction(&self) -> Option<&'static str> {
209 match self {
210 EmbeddingModel::MultilingualE5Small | EmbeddingModel::MultilingualE5Base => {
211 Some("passage: ")
212 }
213 _ => None,
214 }
215 }
216
217 #[inline]
228 pub const fn max_instruction_bytes(&self) -> usize {
229 let q = match self.query_instruction() {
230 Some(s) => s.len(),
231 None => 0,
232 };
233 let d = match self.document_instruction() {
234 Some(s) => s.len(),
235 None => 0,
236 };
237 if q > d { q } else { d }
238 }
239
240 #[inline]
242 pub const fn model_id(&self) -> &'static str {
243 match self {
244 EmbeddingModel::BgeSmallEnV15 => "BAAI/bge-small-en-v1.5",
245 EmbeddingModel::BgeBaseEnV15 => "BAAI/bge-base-en-v1.5",
246 EmbeddingModel::BgeLargeEnV15 => "BAAI/bge-large-en-v1.5",
247 EmbeddingModel::MultilingualE5Small => "intfloat/multilingual-e5-small",
248 EmbeddingModel::MultilingualE5Base => "intfloat/multilingual-e5-base",
249 EmbeddingModel::AllMiniLmL6V2 => "sentence-transformers/all-MiniLM-L6-v2",
250 EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
251 "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
252 }
253 EmbeddingModel::Qwen3Embedding0_6B => "Qwen/Qwen3-Embedding-0.6B",
254 EmbeddingModel::Qwen3Embedding4B => "Qwen/Qwen3-Embedding-4B",
255 EmbeddingModel::TextEmbedding3Small => "text-embedding-3-small",
256 }
257 }
258
259 #[inline]
261 pub const fn supports_output_dim(&self) -> bool {
262 matches!(
263 self,
264 EmbeddingModel::Qwen3Embedding0_6B | EmbeddingModel::Qwen3Embedding4B
265 )
266 }
267
268 #[cfg(feature = "native")]
272 #[inline]
273 pub const fn bert_pooling(&self) -> Option<lattice_inference::BertPooling> {
274 match self {
275 EmbeddingModel::BgeSmallEnV15
276 | EmbeddingModel::BgeBaseEnV15
277 | EmbeddingModel::BgeLargeEnV15 => Some(lattice_inference::BertPooling::CLS),
278 EmbeddingModel::MultilingualE5Small | EmbeddingModel::MultilingualE5Base => {
279 Some(lattice_inference::BertPooling::Mean)
280 }
281 EmbeddingModel::AllMiniLmL6V2 | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
282 Some(lattice_inference::BertPooling::Mean)
283 }
284 EmbeddingModel::Qwen3Embedding0_6B
285 | EmbeddingModel::Qwen3Embedding4B
286 | EmbeddingModel::TextEmbedding3Small => None,
287 }
288 }
289
290 #[inline]
292 pub const fn key_version(&self) -> &'static str {
293 match self {
294 EmbeddingModel::TextEmbedding3Small
295 | EmbeddingModel::Qwen3Embedding0_6B
296 | EmbeddingModel::Qwen3Embedding4B => "v3",
297 EmbeddingModel::AllMiniLmL6V2 | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
298 "v2"
299 }
300 _ => "v1.5",
301 }
302 }
303}
304
305impl std::fmt::Display for EmbeddingModel {
306 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
307 match self {
308 EmbeddingModel::BgeSmallEnV15 => write!(f, "bge-small-en-v1.5"),
309 EmbeddingModel::BgeBaseEnV15 => write!(f, "bge-base-en-v1.5"),
310 EmbeddingModel::BgeLargeEnV15 => write!(f, "bge-large-en-v1.5"),
311 EmbeddingModel::MultilingualE5Small => write!(f, "multilingual-e5-small"),
312 EmbeddingModel::MultilingualE5Base => write!(f, "multilingual-e5-base"),
313 EmbeddingModel::Qwen3Embedding0_6B => write!(f, "qwen3-embedding-0.6b"),
314 EmbeddingModel::Qwen3Embedding4B => write!(f, "qwen3-embedding-4b"),
315 EmbeddingModel::AllMiniLmL6V2 => write!(f, "all-minilm-l6-v2"),
316 EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
317 write!(f, "paraphrase-multilingual-minilm-l12-v2")
318 }
319 EmbeddingModel::TextEmbedding3Small => write!(f, "text-embedding-3-small"),
320 }
321 }
322}
323
324impl std::str::FromStr for EmbeddingModel {
325 type Err = String;
326
327 fn from_str(s: &str) -> Result<Self, Self::Err> {
331 let lower = s.to_lowercase();
332 let normalized = lower.trim().replace("_", "-").replace("baai/", "");
333
334 match normalized.as_str() {
335 "bge-small-en-v1.5" | "bge-small-en" | "bge-small" | "small" => {
336 Ok(EmbeddingModel::BgeSmallEnV15)
337 }
338 "bge-base-en-v1.5" | "bge-base-en" | "bge-base" | "base" => {
339 Ok(EmbeddingModel::BgeBaseEnV15)
340 }
341 "bge-large-en-v1.5" | "bge-large-en" | "bge-large" | "large" => {
342 Ok(EmbeddingModel::BgeLargeEnV15)
343 }
344 "multilingual-e5-small" | "e5-small" | "intfloat/multilingual-e5-small" => {
345 Ok(EmbeddingModel::MultilingualE5Small)
346 }
347 "multilingual-e5-base" | "e5-base" | "intfloat/multilingual-e5-base" => {
348 Ok(EmbeddingModel::MultilingualE5Base)
349 }
350 "qwen3-embedding-0.6b" | "qwen3-embedding" | "qwen3" | "qwen/qwen3-embedding-0.6b" => {
351 Ok(EmbeddingModel::Qwen3Embedding0_6B)
352 }
353 "qwen3-embedding-4b" | "qwen3-4b" | "qwen/qwen3-embedding-4b" => {
354 Ok(EmbeddingModel::Qwen3Embedding4B)
355 }
356 "all-minilm-l6-v2"
357 | "minilm"
358 | "all-minilm"
359 | "sentence-transformers/all-minilm-l6-v2" => Ok(EmbeddingModel::AllMiniLmL6V2),
360 "paraphrase-multilingual-minilm-l12-v2"
361 | "paraphrase-multilingual"
362 | "multilingual-minilm"
363 | "sentence-transformers/paraphrase-multilingual-minilm-l12-v2" => {
364 Ok(EmbeddingModel::ParaphraseMultilingualMiniLmL12V2)
365 }
366 "text-embedding-3-small" | "openai-small" | "openai" => {
367 Ok(EmbeddingModel::TextEmbedding3Small)
368 }
369 _ => Err(format!(
370 "unknown embedding model: '{s}'. Valid: bge-small-en-v1.5, bge-base-en-v1.5, bge-large-en-v1.5, multilingual-e5-small, multilingual-e5-base, text-embedding-3-small"
371 )),
372 }
373 }
374}
375
376pub const MIN_MRL_OUTPUT_DIM: usize = 32;
382
383#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
387pub struct ModelConfig {
388 pub model: EmbeddingModel,
390 #[serde(default)]
392 pub output_dim: Option<usize>,
393}
394
395impl Default for ModelConfig {
396 fn default() -> Self {
397 Self::new(EmbeddingModel::default())
398 }
399}
400
401impl ModelConfig {
402 pub const fn new(model: EmbeddingModel) -> Self {
404 Self {
405 model,
406 output_dim: None,
407 }
408 }
409
410 pub fn try_new(
412 model: EmbeddingModel,
413 output_dim: Option<usize>,
414 ) -> std::result::Result<Self, crate::error::EmbedError> {
415 let config = Self { model, output_dim };
416 config.validate()?;
417 Ok(config)
418 }
419
420 pub fn validate(&self) -> std::result::Result<(), crate::error::EmbedError> {
422 let Some(dim) = self.output_dim else {
423 return Ok(());
424 };
425 if !self.model.supports_output_dim() {
426 return Err(crate::error::EmbedError::InvalidInput(format!(
427 "{} does not support configurable embedding dimensions",
428 self.model
429 )));
430 }
431 if dim < MIN_MRL_OUTPUT_DIM {
432 return Err(crate::error::EmbedError::InvalidInput(format!(
433 "embedding output dimension {dim} is below minimum {MIN_MRL_OUTPUT_DIM}"
434 )));
435 }
436 let native = self.model.native_dimensions();
437 if dim > native {
438 return Err(crate::error::EmbedError::InvalidInput(format!(
439 "embedding output dimension {dim} exceeds native dimension {native} for {}",
440 self.model
441 )));
442 }
443 Ok(())
444 }
445
446 pub fn dimensions(&self) -> usize {
448 self.output_dim
449 .unwrap_or_else(|| self.model.native_dimensions())
450 }
451}
452
453#[cfg(test)]
454mod tests {
455 use super::*;
456
457 #[test]
458 fn test_default_model() {
459 let model = EmbeddingModel::default();
460 assert_eq!(model, EmbeddingModel::BgeSmallEnV15);
461 }
462
463 #[test]
464 fn test_model_provenance_new() {
465 let provenance = ModelProvenance::new(
466 EmbeddingModel::BgeSmallEnV15,
467 "BAAI/bge-small-en-v1.5".into(),
468 );
469
470 assert_eq!(provenance.model, EmbeddingModel::BgeSmallEnV15);
471 assert_eq!(provenance.model_id, "BAAI/bge-small-en-v1.5");
472 assert!(!provenance.hash.is_empty());
473 assert_eq!(provenance.hash.len(), 64); assert!(!provenance.loaded_at_iso.is_empty());
475 }
476
477 #[test]
478 fn test_model_provenance_unique_hash() {
479 let p1 = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "model1".into());
480 std::thread::sleep(std::time::Duration::from_millis(10)); let p2 = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "model1".into());
482
483 assert_ne!(p1.hash, p2.hash);
485 }
486
487 #[test]
488 fn test_model_provenance_dimensions() {
489 let p1 = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "small".into());
490 assert_eq!(p1.dimensions(), 384);
491
492 let p2 = ModelProvenance::new(EmbeddingModel::BgeBaseEnV15, "base".into());
493 assert_eq!(p2.dimensions(), 768);
494
495 let p3 = ModelProvenance::new(EmbeddingModel::BgeLargeEnV15, "large".into());
496 assert_eq!(p3.dimensions(), 1024);
497 }
498
499 #[test]
500 fn test_model_provenance_matches_model() {
501 let provenance = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "test".into());
502
503 assert!(provenance.matches_model(EmbeddingModel::BgeSmallEnV15));
504 assert!(!provenance.matches_model(EmbeddingModel::BgeBaseEnV15));
505 assert!(!provenance.matches_model(EmbeddingModel::BgeLargeEnV15));
506 }
507
508 #[test]
509 fn test_model_provenance_serialization() {
510 let provenance = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "test-model".into());
511
512 let json = serde_json::to_string(&provenance).unwrap();
513 assert!(json.contains("bge_small_en_v15"), "json={json}");
516 assert!(json.contains("test-model"));
517 assert!(json.contains(&provenance.hash));
518
519 let parsed: ModelProvenance = serde_json::from_str(&json).unwrap();
520 assert_eq!(parsed.model, provenance.model);
521 assert_eq!(parsed.model_id, provenance.model_id);
522 assert_eq!(parsed.hash, provenance.hash);
523 }
524
525 #[test]
526 fn test_dimensions() {
527 assert_eq!(EmbeddingModel::BgeSmallEnV15.dimensions(), 384);
528 assert_eq!(EmbeddingModel::BgeBaseEnV15.dimensions(), 768);
529 assert_eq!(EmbeddingModel::BgeLargeEnV15.dimensions(), 1024);
530 assert_eq!(EmbeddingModel::Qwen3Embedding4B.dimensions(), 2560);
531 }
532
533 #[test]
534 fn test_model_config_native_dims() {
535 assert_eq!(
536 ModelConfig::new(EmbeddingModel::Qwen3Embedding4B).dimensions(),
537 2560
538 );
539 assert_eq!(
540 ModelConfig::new(EmbeddingModel::Qwen3Embedding0_6B).dimensions(),
541 1024
542 );
543 assert_eq!(
544 ModelConfig::new(EmbeddingModel::BgeSmallEnV15).dimensions(),
545 384
546 );
547 }
548
549 #[test]
550 fn test_model_config_configured_dim() {
551 let cfg = ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(1024)).unwrap();
552 assert_eq!(cfg.dimensions(), 1024);
553
554 let cfg = ModelConfig::try_new(EmbeddingModel::Qwen3Embedding0_6B, Some(512)).unwrap();
555 assert_eq!(cfg.dimensions(), 512);
556 }
557
558 #[test]
559 fn test_model_config_validation_below_min() {
560 assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(31)).is_err());
561 assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(0)).is_err());
562 }
563
564 #[test]
565 fn test_model_config_validation_above_native() {
566 assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(2561)).is_err());
567 assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding0_6B, Some(1025)).is_err());
568 }
569
570 #[test]
571 fn test_model_config_validation_non_mrl_model() {
572 assert!(ModelConfig::try_new(EmbeddingModel::BgeSmallEnV15, Some(128)).is_err());
573 assert!(ModelConfig::try_new(EmbeddingModel::BgeBaseEnV15, Some(512)).is_err());
574 }
575
576 #[test]
577 fn test_model_config_none_output_dim_ok_for_any_model() {
578 assert!(ModelConfig::try_new(EmbeddingModel::BgeSmallEnV15, None).is_ok());
579 assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, None).is_ok());
580 }
581
582 #[test]
583 fn test_is_local() {
584 assert!(EmbeddingModel::BgeSmallEnV15.is_local());
585 assert!(EmbeddingModel::BgeBaseEnV15.is_local());
586 assert!(EmbeddingModel::BgeLargeEnV15.is_local());
587 }
588
589 #[test]
590 fn test_display() {
591 assert_eq!(
592 EmbeddingModel::BgeSmallEnV15.to_string(),
593 "bge-small-en-v1.5"
594 );
595 assert_eq!(EmbeddingModel::BgeBaseEnV15.to_string(), "bge-base-en-v1.5");
596 assert_eq!(
597 EmbeddingModel::BgeLargeEnV15.to_string(),
598 "bge-large-en-v1.5"
599 );
600 }
601
602 #[test]
603 fn test_serialization_roundtrip() {
604 let model = EmbeddingModel::BgeSmallEnV15;
605 let json = serde_json::to_string(&model).unwrap();
606 let parsed: EmbeddingModel = serde_json::from_str(&json).unwrap();
607 assert_eq!(model, parsed);
608 }
609
610 #[test]
611 fn test_max_input_tokens() {
612 assert_eq!(EmbeddingModel::BgeSmallEnV15.max_input_tokens(), 512);
613 assert_eq!(EmbeddingModel::BgeBaseEnV15.max_input_tokens(), 512);
614 assert_eq!(EmbeddingModel::BgeLargeEnV15.max_input_tokens(), 512);
615 }
616
617 #[test]
618 fn test_from_str_display_names() {
619 assert_eq!(
620 "bge-small-en-v1.5".parse::<EmbeddingModel>().unwrap(),
621 EmbeddingModel::BgeSmallEnV15
622 );
623 assert_eq!(
624 "bge-base-en-v1.5".parse::<EmbeddingModel>().unwrap(),
625 EmbeddingModel::BgeBaseEnV15
626 );
627 assert_eq!(
628 "bge-large-en-v1.5".parse::<EmbeddingModel>().unwrap(),
629 EmbeddingModel::BgeLargeEnV15
630 );
631 }
632
633 #[test]
634 fn test_from_str_short_names() {
635 assert_eq!(
636 "small".parse::<EmbeddingModel>().unwrap(),
637 EmbeddingModel::BgeSmallEnV15
638 );
639 assert_eq!(
640 "bge-base".parse::<EmbeddingModel>().unwrap(),
641 EmbeddingModel::BgeBaseEnV15
642 );
643 assert_eq!(
644 "LARGE".parse::<EmbeddingModel>().unwrap(), EmbeddingModel::BgeLargeEnV15
646 );
647 }
648
649 #[test]
650 fn test_from_str_huggingface_ids() {
651 assert_eq!(
652 "BAAI/bge-small-en-v1.5".parse::<EmbeddingModel>().unwrap(),
653 EmbeddingModel::BgeSmallEnV15
654 );
655 }
656
657 #[test]
658 fn test_from_str_invalid() {
659 let result = "unknown-model".parse::<EmbeddingModel>();
660 assert!(result.is_err());
661 assert!(result.unwrap_err().contains("unknown embedding model"));
662 }
663
664 #[cfg(feature = "native")]
670 #[test]
671 fn test_bge_models_use_cls_pooling() {
672 use lattice_inference::BertPooling;
673
674 assert_eq!(
675 EmbeddingModel::BgeSmallEnV15.bert_pooling(),
676 Some(BertPooling::CLS),
677 "BgeSmallEnV15 must use CLS pooling"
678 );
679 assert_eq!(
680 EmbeddingModel::BgeBaseEnV15.bert_pooling(),
681 Some(BertPooling::CLS),
682 "BgeBaseEnV15 must use CLS pooling"
683 );
684 assert_eq!(
685 EmbeddingModel::BgeLargeEnV15.bert_pooling(),
686 Some(BertPooling::CLS),
687 "BgeLargeEnV15 must use CLS pooling"
688 );
689 }
690
691 #[cfg(feature = "native")]
693 #[test]
694 fn test_e5_models_use_mean_pooling() {
695 use lattice_inference::BertPooling;
696
697 assert_eq!(
698 EmbeddingModel::MultilingualE5Small.bert_pooling(),
699 Some(BertPooling::Mean),
700 "MultilingualE5Small must use mean pooling"
701 );
702 assert_eq!(
703 EmbeddingModel::MultilingualE5Base.bert_pooling(),
704 Some(BertPooling::Mean),
705 "MultilingualE5Base must use mean pooling"
706 );
707 }
708
709 #[cfg(feature = "native")]
711 #[test]
712 fn test_minilm_models_use_mean_pooling() {
713 use lattice_inference::BertPooling;
714
715 assert_eq!(
716 EmbeddingModel::AllMiniLmL6V2.bert_pooling(),
717 Some(BertPooling::Mean),
718 "AllMiniLmL6V2 must use mean pooling"
719 );
720 assert_eq!(
721 EmbeddingModel::ParaphraseMultilingualMiniLmL12V2.bert_pooling(),
722 Some(BertPooling::Mean),
723 "ParaphraseMultilingualMiniLmL12V2 must use mean pooling"
724 );
725 }
726
727 #[cfg(feature = "native")]
729 #[test]
730 fn test_non_bert_models_return_none_pooling() {
731 assert_eq!(
732 EmbeddingModel::Qwen3Embedding0_6B.bert_pooling(),
733 None,
734 "Qwen model must return None for bert_pooling()"
735 );
736 assert_eq!(
737 EmbeddingModel::Qwen3Embedding4B.bert_pooling(),
738 None,
739 "Qwen model must return None for bert_pooling()"
740 );
741 assert_eq!(
742 EmbeddingModel::TextEmbedding3Small.bert_pooling(),
743 None,
744 "Remote model must return None for bert_pooling()"
745 );
746 }
747
748 #[cfg(feature = "native")]
750 #[test]
751 fn test_bge_and_e5_use_different_pooling() {
752 assert_ne!(
753 EmbeddingModel::BgeSmallEnV15.bert_pooling(),
754 EmbeddingModel::MultilingualE5Small.bert_pooling(),
755 "BGE and E5 must use different pooling strategies"
756 );
757 }
758}