1use super::isq::UqffFullSer;
2use super::{
3 get_model_paths, get_xlora_paths, AdapterKind, AnyMoePipelineMixin, CacheManagerMixin,
4 EitherCache, ForwardInputsResult, GeneralMetadata, IsqPipelineMixin, Loader, MetadataMixin,
5 ModelCategory, ModelKind, ModelPaths, PreProcessingMixin, TokenSource,
6};
7use crate::attention::ATTENTION_CHUNK_SIZE;
8use crate::device_map::{self, DeviceMapper};
9use crate::distributed::{self, use_ring, WorkerTransferData};
10use crate::embedding_models::inputs_processor::{EmbeddingProcessor, ModelInputs};
11use crate::embedding_models::{Dense, DenseActivation, Normalize, Pooling};
12use crate::embedding_normal_model_loader;
13use crate::embedding_normal_model_loader_sharded;
14use crate::get_embedding_paths;
15use crate::paged_attention::AttentionImplementation;
16use crate::pipeline::loaders::auto_device_map;
17use crate::pipeline::loaders::QuantizationConfigShim;
18use crate::pipeline::sampling::sample_and_add_toks;
19use crate::pipeline::EmbeddingLoaderType;
20use crate::pipeline::EmbeddingModel;
21use crate::pipeline::EmbeddingModelLoader;
22use crate::pipeline::{AutoEmbeddingLoader, EmbeddingModulePaths};
23use crate::pipeline::{ChatTemplate, EmbeddingModelPaths, IsqOrganization, Processor};
24use crate::pipeline::{EmbeddingGemmaLoader, Qwen3EmbeddingLoader};
25use crate::prefix_cacher::PrefixCacheManagerV2;
26use crate::sequence::Sequence;
27use crate::utils::tokenizer::get_tokenizer;
28use crate::utils::{
29 progress::{new_multi_progress, ProgressScopeGuard},
30 tokens::get_token,
31 varbuilder_utils::from_mmaped_safetensors,
32};
33use crate::Modalities;
34use crate::SupportedModality;
35use crate::{
36 get_uqff_paths, DeviceMapSetting, PagedAttentionConfig, Pipeline, Topology, TryIntoDType,
37 GLOBAL_HF_CACHE,
38};
39use anyhow::Context;
40use anyhow::Result;
41use candle_core::{Device, Tensor};
42use candle_nn::{Linear, Module};
43use hf_hub::Cache;
44use hf_hub::{api::sync::ApiBuilder, Repo, RepoType};
45use mistralrs_quant::log::once_log_info;
46use mistralrs_quant::safetensors::MmapedSafetensors;
47use mistralrs_quant::{
48 AfqLayer, GgufMatMul, HqqLayer, ImmediateIsqOverride, IsqType, QuantizedSerdeType,
49};
50use rand_isaac::Isaac64Rng;
51use std::any::Any;
52use std::borrow::Cow;
53use std::env;
54use std::path::{Path, PathBuf};
55use std::str::FromStr;
56use std::sync::{Arc, RwLock};
57use tokenizers::Tokenizer;
58use tokio::sync::Mutex;
59use tracing::{info, warn};
60
61pub struct EmbeddingPipeline {
62 model: Box<dyn EmbeddingModel + Send + Sync>,
63 tokenizer: Arc<Tokenizer>,
64 model_id: String,
65 metadata: Arc<GeneralMetadata>,
66 topology: Option<Topology>,
67 silent: bool,
68 config: String,
69 modules_ser: String,
70 modules_manifest: Vec<EmbeddingModulePaths>,
71 mapper: Box<dyn DeviceMapper + Send + Sync>,
72 modules: Vec<Box<dyn Module + Send + Sync>>,
73 processor: Arc<dyn Processor + Send + Sync>,
74}
75
76pub struct EmbeddingLoader {
78 inner: Box<dyn EmbeddingModelLoader>,
79 model_id: String,
80 config: EmbeddingSpecificConfig,
81 kind: ModelKind,
82 tokenizer_json: Option<String>,
83 token_source: RwLock<Option<TokenSource>>,
84 revision: RwLock<Option<String>>,
85 from_uqff: RwLock<Option<Vec<PathBuf>>>,
86 hf_cache_path: Option<PathBuf>,
87 lora_adapter_ids: Option<Vec<String>>,
88}
89
90#[derive(Default)]
91pub struct EmbeddingLoaderBuilder {
93 model_id: Option<String>,
94 config: EmbeddingSpecificConfig,
95 kind: ModelKind,
96 tokenizer_json: Option<String>,
97 hf_cache_path: Option<PathBuf>,
98 lora_adapter_ids: Option<Vec<String>>,
99}
100
101#[derive(Clone, Default)]
102pub struct EmbeddingSpecificConfig {
104 pub topology: Option<Topology>,
105 pub write_uqff: Option<PathBuf>,
106 pub from_uqff: Option<Vec<PathBuf>>,
107 pub hf_cache_path: Option<PathBuf>,
108}
109
110impl EmbeddingLoaderBuilder {
111 pub fn new(
112 config: EmbeddingSpecificConfig,
113 tokenizer_json: Option<String>,
114 model_id: Option<String>,
115 ) -> Self {
116 Self {
117 config,
118 tokenizer_json,
119 model_id,
120 kind: ModelKind::Normal,
121 hf_cache_path: None,
122 ..Default::default()
123 }
124 }
125
126 pub fn hf_cache_path(mut self, hf_cache_path: PathBuf) -> Self {
127 self.hf_cache_path = Some(hf_cache_path);
128 self
129 }
130
131 pub fn with_lora(mut self, lora_adapter_ids: Vec<String>) -> Self {
132 self.kind = ModelKind::Adapter {
133 adapter: AdapterKind::Lora,
134 };
135 self.lora_adapter_ids = Some(lora_adapter_ids);
136 self
137 }
138
139 pub fn build(self, loader: Option<EmbeddingLoaderType>) -> Box<dyn Loader> {
140 let loader: Box<dyn EmbeddingModelLoader> = match loader {
141 Some(EmbeddingLoaderType::EmbeddingGemma) => Box::new(EmbeddingGemmaLoader),
142 Some(EmbeddingLoaderType::Qwen3Embedding) => Box::new(Qwen3EmbeddingLoader),
143 None => Box::new(AutoEmbeddingLoader),
144 };
145 Box::new(EmbeddingLoader {
146 inner: loader,
147 model_id: self.model_id.unwrap(),
148 config: self.config,
149 kind: self.kind,
150 tokenizer_json: self.tokenizer_json,
151 token_source: RwLock::new(None),
152 revision: RwLock::new(None),
153 from_uqff: RwLock::new(None),
154 hf_cache_path: self.hf_cache_path,
155 lora_adapter_ids: self.lora_adapter_ids,
156 })
157 }
158}
159
160impl Loader for EmbeddingLoader {
161 #[allow(clippy::type_complexity, clippy::too_many_arguments)]
162 fn load_model_from_hf(
163 &self,
164 revision: Option<String>,
165 token_source: TokenSource,
166 dtype: &dyn TryIntoDType,
167 device: &Device,
168 silent: bool,
169 mapper: DeviceMapSetting,
170 in_situ_quant: Option<IsqType>,
171 paged_attn_config: Option<PagedAttentionConfig>,
172 ) -> Result<Arc<Mutex<dyn Pipeline + Send + Sync>>> {
173 let _progress_guard = ProgressScopeGuard::new(silent);
174 let cache = self
175 .hf_cache_path
176 .clone()
177 .map(Cache::new)
178 .unwrap_or_default();
179 GLOBAL_HF_CACHE.get_or_init(|| cache);
180
181 let paths: anyhow::Result<Box<dyn ModelPaths>> = get_embedding_paths!(
182 EmbeddingModelPaths,
183 &token_source,
184 revision.clone(),
185 self,
186 None,
187 None,
188 silent,
189 self.config.from_uqff.is_some()
190 );
191 *self
192 .token_source
193 .write()
194 .expect("Failed to write to token source") = Some(token_source);
195 *self.revision.write().expect("Failed to write to revision") = revision.clone();
196 if let Some(from_uqff) = self.config.from_uqff.clone() {
197 *self.from_uqff.write().unwrap() = Some(get_uqff_paths!(&from_uqff, self, silent));
198 }
199 self.load_model_from_path(
200 &paths?,
201 dtype,
202 device,
203 silent,
204 mapper,
205 in_situ_quant,
206 paged_attn_config,
207 )
208 }
209
210 #[allow(clippy::type_complexity, clippy::too_many_arguments)]
211 fn load_model_from_path(
212 &self,
213 paths: &Box<dyn ModelPaths>,
214 dtype: &dyn TryIntoDType,
215 device: &Device,
216 silent: bool,
217 mut mapper: DeviceMapSetting,
218 in_situ_quant: Option<IsqType>,
219 mut paged_attn_config: Option<PagedAttentionConfig>,
220 ) -> Result<Arc<Mutex<dyn Pipeline + Send + Sync>>> {
221 let _progress_guard = ProgressScopeGuard::new(silent);
222 let config = std::fs::read_to_string(paths.get_config_filename())?;
223
224 if paged_attn_config.is_some() {
225 warn!("PagedAttention is not supported for embedding models, disabling it.");
226 paged_attn_config = None;
227 }
228
229 info!("Prompt chunk size is {ATTENTION_CHUNK_SIZE}.");
230
231 let use_nccl = mistralrs_quant::distributed::use_nccl();
232
233 let available_devices = if let Ok(payload) = env::var(distributed::IS_DAEMON_FLAG) {
234 let payload: WorkerTransferData = serde_json::from_str(&payload)?;
235 let WorkerTransferData::Init { id: _, worker_rank } = payload;
236 vec![candle_core::Device::new_cuda_with_stream(worker_rank + 1)?]
237 } else if use_nccl || use_ring() {
238 vec![candle_core::Device::new_cuda_with_stream(0)?]
239 } else {
240 device_map::get_all_similar_devices(device)?
241 };
242 #[cfg(feature = "cuda")]
243 for device in &available_devices {
244 if let Device::Cuda(dev) = device {
245 unsafe { dev.disable_event_tracking() };
246 }
247 }
248 let device = if use_nccl || use_ring() {
249 available_devices[0].clone()
250 } else {
251 device.clone()
252 };
253
254 if use_nccl || use_ring() {
256 mapper = DeviceMapSetting::DummyNccl {
257 nm_device: available_devices[0].clone(),
258 };
259 } else if let DeviceMapSetting::Auto(params) = mapper.clone() {
260 let dtype = dtype.try_into_dtype(&available_devices.iter().collect::<Vec<_>>())?;
262
263 let (layer_sizes_in_bytes, non_mapped_size_in_bytes, total_model_size_in_bytes) =
266 if let Some(serialized) = &*self.from_uqff.read().unwrap() {
267 let weight_pack_factor = {
268 let ser_artifacts = unsafe {
269 candle_core::safetensors::MmapedSafetensors::multi(serialized)?
270 };
271 let mut total_pack_factors = 0;
272 let total_tensors = ser_artifacts.tensors().len();
273 for (_, artifact) in ser_artifacts.tensors() {
274 let artifact = artifact.data();
275 let isq_type = artifact[mistralrs_quant::UQFF_QUANT_TYPE_OFFSET];
277 let pack_factor = match QuantizedSerdeType::try_from(isq_type as usize)?
278 {
279 QuantizedSerdeType::Hqq => {
280 HqqLayer::get_isq_type_from_uqff(Cow::Borrowed(artifact))?
281 .pack_factor(dtype)
282 }
283 QuantizedSerdeType::Gguf => {
284 GgufMatMul::get_isq_type_from_uqff(Cow::Borrowed(artifact))?
285 .pack_factor(dtype)
286 }
287 QuantizedSerdeType::Fp8 => IsqType::F8E4M3.pack_factor(dtype),
288 QuantizedSerdeType::Unquant => 1,
289 QuantizedSerdeType::Afq => {
290 AfqLayer::get_isq_type_from_uqff(Cow::Borrowed(artifact))?
291 .pack_factor(dtype)
292 }
293 QuantizedSerdeType::F8Q8 => IsqType::F8Q8.pack_factor(dtype),
294 QuantizedSerdeType::Mxfp4 => IsqType::MXFP4.pack_factor(dtype),
295 };
296 total_pack_factors += pack_factor;
297 }
298
299 total_pack_factors / total_tensors
300 };
301
302 let layer_sizes_in_bytes = self.inner.layer_sizes_in_bytes(
303 &config,
304 dtype,
305 weight_pack_factor,
306 None,
307 )?;
308 let non_mapped_size_in_bytes = self.inner.non_mapped_size_in_bytes(
309 &config,
310 dtype,
311 weight_pack_factor,
312 None,
313 )?;
314 let layer_sizes_sum = layer_sizes_in_bytes.iter().sum::<usize>();
315 (
316 layer_sizes_in_bytes,
317 non_mapped_size_in_bytes,
318 layer_sizes_sum + non_mapped_size_in_bytes,
319 )
320 } else if let Some(isq) = in_situ_quant {
321 let weight_pack_factor = isq.pack_factor(dtype);
322 let layer_sizes_in_bytes = self.inner.layer_sizes_in_bytes(
323 &config,
324 dtype,
325 weight_pack_factor,
326 None,
327 )?;
328 let non_mapped_size_in_bytes = self.inner.non_mapped_size_in_bytes(
329 &config,
330 dtype,
331 weight_pack_factor,
332 None,
333 )?;
334 let layer_sizes_sum = layer_sizes_in_bytes.iter().sum::<usize>();
335 (
336 layer_sizes_in_bytes,
337 non_mapped_size_in_bytes,
338 layer_sizes_sum + non_mapped_size_in_bytes,
339 )
340 } else {
341 let weight_pack_factor =
343 QuantizationConfigShim::get_quant_config_pack_factor(&config, dtype)?;
344 let layer_sizes_in_bytes = self.inner.layer_sizes_in_bytes(
345 &config,
346 dtype,
347 weight_pack_factor,
348 None,
349 )?;
350 let non_mapped_size_in_bytes = self.inner.non_mapped_size_in_bytes(
351 &config,
352 dtype,
353 weight_pack_factor,
354 None,
355 )?;
356 let layer_sizes_sum = layer_sizes_in_bytes.iter().sum::<usize>();
357 (
358 layer_sizes_in_bytes,
359 non_mapped_size_in_bytes,
360 layer_sizes_sum + non_mapped_size_in_bytes,
361 )
362 };
363
364 let new = auto_device_map::get_device_layers(
365 &*self.inner,
366 &config,
367 self.inner.num_layers(&config)?,
368 layer_sizes_in_bytes,
369 non_mapped_size_in_bytes,
370 total_model_size_in_bytes,
371 &available_devices,
372 dtype,
373 ¶ms,
374 paged_attn_config.as_ref(),
375 )?;
376 mapper = DeviceMapSetting::Map(new);
377 }
378
379 let pipeline_mapper = mapper.into_mapper(
380 self.inner.num_layers(&config)?,
381 &device,
382 self.config.topology.as_ref(),
383 &available_devices,
384 )?;
385 let mapper = mapper.into_mapper(
386 self.inner.num_layers(&config)?,
387 &device,
388 self.config.topology.as_ref(),
389 &available_devices,
390 )?;
391 let mut layer_devices = Vec::new();
392 for layer in 0..self.inner.num_layers(&config)? {
393 let device = mapper.device_for(layer, false).cloned();
394 layer_devices.push(device);
395 }
396 let dtype = mapper.get_min_dtype(dtype)?;
397
398 info!("Model config: {:?}", self.inner.get_config_repr(&config)?);
399 if crate::using_flash_attn() {
400 once_log_info("FlashAttention is enabled.");
401 }
402
403 let topology_overrides = self
404 .config
405 .topology
406 .as_ref()
407 .map(|topology| {
408 topology
409 .pattern_overrides()
410 .into_iter()
411 .map(|(regex, layer)| ImmediateIsqOverride {
412 predicate: regex,
413 ty: layer.isq,
414 device: layer.device.clone(),
415 })
416 .collect::<Vec<_>>()
417 })
418 .unwrap_or_default();
419 let has_override_isq = topology_overrides
420 .iter()
421 .any(|override_entry| override_entry.ty.is_some());
422 let topology_requires_post_quant = self
423 .config
424 .topology
425 .as_ref()
426 .is_some_and(|topology| topology.requires_post_quantization());
427
428 let allow_immediate_cli = in_situ_quant.is_some();
429
430 let mut immediate_ty = None;
431 let mut immediate_predicates = Vec::new();
432 if allow_immediate_cli {
433 immediate_ty = in_situ_quant;
434 immediate_predicates = self.inner.immediate_isq_predicates(&config)?;
435 info!("Applying ISQ to {in_situ_quant:?}");
436 if immediate_predicates.is_empty() {
437 warn!("No predicates for this model and ISQ setting detected. ISQ will not be applied to any weights!");
438 }
439 }
440
441 let use_immediate = allow_immediate_cli || has_override_isq;
442 if use_immediate {
443 let (pool, num_threads) = mistralrs_quant::create_isq_thread_pool(immediate_ty);
444 info!("Applying immediate ISQ in parallel on {num_threads} threads.");
445 mistralrs_quant::set_immediate_isq_with_pool(
446 immediate_ty,
447 immediate_predicates.clone(),
448 topology_overrides.clone(),
449 pool,
450 );
451 }
452
453 let mut loading_isq = if use_immediate {
455 false
456 } else {
457 in_situ_quant.is_some()
458 };
459 loading_isq |= topology_requires_post_quant;
460 loading_isq |= self.config.from_uqff.is_some();
461
462 let load_device = if !loading_isq {
468 loading_isq = false;
469 if use_immediate && !crate::utils::normal::is_integrated_gpu(&device) {
470 Device::Cpu
471 } else {
472 device.clone()
473 }
474 } else {
475 Device::Cpu
476 };
477
478 let attention_mechanism = if paged_attn_config.is_some() {
479 AttentionImplementation::PagedAttention
480 } else {
481 AttentionImplementation::Eager
482 };
483
484 let multi_progress = Arc::new(new_multi_progress());
485
486 let modules_config: Vec<_> = paths
487 .get_modules()
488 .context("Embedding models require the `modules.json` file.")?
489 .to_vec();
490 assert!(matches!(
491 modules_config.first(),
492 Some(EmbeddingModulePaths::Transformer { .. })
493 ));
494
495 let mut modules: Vec<Box<dyn Module + Send + Sync>> = Vec::new();
496 for module in &modules_config {
497 match module {
498 EmbeddingModulePaths::Transformer { .. } => (),
499 EmbeddingModulePaths::Pooling { config, .. } => {
500 let layer: Pooling = serde_json::from_str(&std::fs::read_to_string(config)?)?;
501 modules.push(Box::new(layer));
502 }
503 EmbeddingModulePaths::Dense { config, model, .. } => {
504 let config: Dense = serde_json::from_str(&std::fs::read_to_string(config)?)?;
505 let safetensors = unsafe { MmapedSafetensors::new(model)? };
506 let weight = safetensors.load("linear.weight", &device, Some(dtype))?;
507 let bias = if config.bias {
508 Some(safetensors.load("linear.bias", &device, Some(dtype))?)
509 } else {
510 None
511 };
512 let (out_f, in_f) = weight.dims2()?;
513 assert_eq!((out_f, in_f), (config.out_features, config.in_features));
514 if !matches!(config.activation_function, DenseActivation::Identity) {
515 anyhow::bail!("Expected Identity activation function.");
516 }
517
518 modules.push(Box::new(Linear::new(weight, bias)));
519 }
520 EmbeddingModulePaths::Normalize { .. } => {
521 modules.push(Box::new(Normalize));
522 }
523 }
524 }
525 let modules_ser = EmbeddingModulePaths::serialize_modules(&modules_config);
526
527 let mut model = if use_nccl || use_ring() {
528 let (mapper, sharded_vb) = distributed::prepare_distributed_mapper(
529 dtype,
530 &device,
531 &available_devices,
532 silent,
533 &config,
534 loading_isq,
535 self.config.from_uqff.is_some(),
536 IsqOrganization::Default,
537 &*self.inner,
538 paths.as_ref(),
539 )?;
540
541 match self.kind {
543 ModelKind::Normal => embedding_normal_model_loader_sharded!(
544 sharded_vb,
545 config,
546 self.inner,
547 mapper,
548 loading_isq,
549 device.clone(),
550 attention_mechanism,
551 multi_progress.clone(),
552 ),
553 _ => unreachable!(),
554 }
555 } else {
556 match self.kind {
557 ModelKind::Normal => embedding_normal_model_loader!(
558 paths,
559 Some(dtype),
560 &load_device,
561 layer_devices.clone(),
562 config,
563 self.inner,
564 silent,
565 mapper,
566 loading_isq,
567 self.config.from_uqff.is_some(),
568 device.clone(),
569 attention_mechanism,
570 multi_progress,
571 ),
572 _ => unreachable!(),
573 }
574 };
575
576 let tokenizer = get_tokenizer(paths.get_tokenizer_filename(), None)?;
577
578 let should_serialize = self.config.write_uqff.is_some();
579 let should_quantize_pass = loading_isq;
580
581 if (should_quantize_pass || should_serialize) && self.config.from_uqff.is_none() {
582 if should_quantize_pass {
583 info!("Applying ISQ to all ranks.");
584 } else {
585 info!("Serializing existing ISQ tensors without additional quantization.");
586 }
587 model.quantize(
588 in_situ_quant,
589 device.clone(),
590 self.config.topology.as_ref(),
591 silent,
592 None,
593 IsqOrganization::Default,
594 should_quantize_pass,
595 self.config.write_uqff.as_ref(),
596 UqffFullSer {
597 tokenizer: &tokenizer,
598 template_filename: paths.get_template_filename(),
599 generation_config: paths.get_gen_conf_filename(),
600 config: config.clone(),
601 processor_filename: paths.get_processor_config(),
602 preprocessor_filename: paths.get_preprocessor_config(),
603 modules: Some(&modules_ser),
604 module_paths: Some(&modules_config),
605 },
606 Arc::new(new_multi_progress()),
607 )?;
608 } else if let Some(from_uqff) = &*self.from_uqff.read().unwrap() {
609 model.load_from_artifacts(
610 device.clone(),
611 self.config.topology.as_ref(),
612 silent,
613 from_uqff,
614 )?;
615 }
616
617 let has_causal_attention = self.inner.has_causal_attention(&config)?;
618 let max_seq_len = self.inner.model_config(&config)?.max_seq_len();
619 Ok(Arc::new(Mutex::new(EmbeddingPipeline {
620 model,
621 tokenizer: tokenizer.into(),
622 model_id: self.model_id.clone(),
623 metadata: Arc::new(GeneralMetadata {
624 max_seq_len,
625 llg_factory: None,
626 is_xlora: false,
627 no_prefix_cache: false,
628 num_hidden_layers: 1, eos_tok: vec![],
630 kind: ModelKind::Normal,
631 no_kv_cache: true, activation_dtype: dtype,
633 sliding_window: None,
634 cache_config: None,
635 cache_engine: None,
636 model_metadata: None,
637 modalities: Modalities {
638 input: vec![SupportedModality::Text],
639 output: vec![SupportedModality::Embedding],
640 },
641 }),
642 topology: self.config.topology.clone(),
643 silent,
644 config,
645 modules_ser,
646 modules_manifest: modules_config,
647 mapper: pipeline_mapper,
648 modules,
649 processor: Arc::new(EmbeddingProcessor {
650 has_causal_attention,
651 }),
652 })))
653 }
654
655 fn get_id(&self) -> String {
656 self.model_id.to_string()
657 }
658
659 fn get_kind(&self) -> ModelKind {
660 self.kind.clone()
661 }
662}
663
664impl PreProcessingMixin for EmbeddingPipeline {
665 fn get_processor(&self) -> Arc<dyn Processor> {
666 self.processor.clone()
667 }
668 fn get_chat_template(&self) -> Option<Arc<ChatTemplate>> {
669 None
670 }
671 fn get_input_processor_config(&self) -> Option<Arc<dyn Any>> {
672 None
673 }
674}
675
676impl IsqPipelineMixin for EmbeddingPipeline {
677 fn re_isq_model(&mut self, dtype: IsqType) -> Result<()> {
678 let device = self.device().clone();
679 self.model
680 .quantize(
681 Some(dtype),
682 device,
683 self.topology.as_ref(),
684 self.silent,
685 None,
686 IsqOrganization::Default,
687 true,
688 None,
689 UqffFullSer {
690 tokenizer: &self.tokenizer,
691 template_filename: &None,
692 generation_config: None,
693 config: self.config.clone(),
694 processor_filename: &None,
695 preprocessor_filename: &None,
696 modules: Some(&self.modules_ser),
697 module_paths: Some(&self.modules_manifest),
698 },
699 Arc::new(new_multi_progress()),
700 )
701 .map_err(anyhow::Error::msg)
702 }
703}
704
705impl CacheManagerMixin for EmbeddingPipeline {
706 fn clone_in_cache(&self, _seqs: &mut [&mut Sequence]) {}
707 fn clone_out_cache(&self, _seqs: &mut [&mut Sequence]) {}
708 fn set_none_cache(
709 &self,
710 _seqs: &mut [&mut Sequence],
711 _reset_non_granular: bool,
712 _modify_draft_cache: bool,
713 _load_preallocated_cache: bool,
714 ) {
715 }
716 fn cache(&self) -> &EitherCache {
717 unreachable!()
718 }
719}
720
721impl MetadataMixin for EmbeddingPipeline {
722 fn device(&self) -> Device {
723 self.model.device().clone()
724 }
725 fn get_metadata(&self) -> Arc<GeneralMetadata> {
726 self.metadata.clone()
727 }
728 fn name(&self) -> String {
729 self.model_id.clone()
730 }
731 fn reset_non_granular_state(&self) {}
732 fn tokenizer(&self) -> Option<Arc<Tokenizer>> {
733 Some(self.tokenizer.clone())
734 }
735 fn device_mapper(&self) -> Option<&dyn DeviceMapper> {
736 Some(&*self.mapper)
737 }
738}
739
740#[async_trait::async_trait]
741impl Pipeline for EmbeddingPipeline {
742 fn forward_inputs(
743 &mut self,
744 inputs: Box<dyn Any>,
745 _return_raw_logits: bool,
746 ) -> candle_core::Result<ForwardInputsResult> {
747 let ModelInputs {
748 input_ids,
749 flash_meta,
750 } = *inputs.downcast::<ModelInputs>().expect("Downcast failed.");
751
752 let mut xs = self.model.forward(&input_ids, &flash_meta)?;
753 for module in &self.modules {
754 xs = module.forward(&xs)?;
755 }
756
757 Ok(ForwardInputsResult::Embeddings { embeddings: xs })
758 }
759 async fn sample_causal_gen(
760 &self,
761 seqs: &mut [&mut Sequence],
762 logits: Vec<Tensor>,
763 prefix_cacher: &mut PrefixCacheManagerV2,
764 disable_eos_stop: bool,
765 rng: Arc<std::sync::Mutex<Isaac64Rng>>,
766 ) -> Result<(), candle_core::Error> {
767 sample_and_add_toks(self, seqs, logits, prefix_cacher, disable_eos_stop, rng).await
768 }
769 fn category(&self) -> ModelCategory {
770 ModelCategory::Embedding
771 }
772}
773
774impl AnyMoePipelineMixin for EmbeddingPipeline {}