Skip to main content

mistralrs_core/pipeline/
embedding.rs

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
76/// A loader for an embedding (non-quantized) model.
77pub 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)]
91/// A builder for a loader for an embedding (non-quantized) model.
92pub 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)]
102/// Config specific to loading an embedding model.
103pub 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 auto, convert to Map if not using nccl
255        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            // Initial dtype
261            let dtype = dtype.try_into_dtype(&available_devices.iter().collect::<Vec<_>>())?;
262
263            // ISQ or UQFF: quantized path
264            // Match logic below where UQFF has priority
265            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                            // NOTE(EricLBuehler): isq type is ALWAYS byte 4 (5th) of the tensor.
276                            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                    // Be sure to get the weight pack factor here; we might be loading a prequantized model.
342                    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                &params,
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        // Logic for ISQ here: if no calibration (i.e imatrix), then allow immediate ISQ. Otherwise, back to normal.
454        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        // Load onto the regular device if not using isq.
463        // For immediate ISQ on discrete GPUs, load to CPU: the mapper will set the correct target
464        // device per-layer, and linear constructors will override to CPU for ISQ-targeted weights.
465        // On integrated/unified memory systems (e.g. Grace Blackwell), CPU and GPU share memory,
466        // so we load directly to the device.
467        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            // Special case for where things can be more optimially loaded.
542            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, // FIXME(EricLBuehler): we know this is only for caching, so its OK.
629                eos_tok: vec![],
630                kind: ModelKind::Normal,
631                no_kv_cache: true, // NOTE(EricLBuehler): no cache for these.
632                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 {}