diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index f99b0e7..628e532 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -42,6 +42,25 @@ impl<'a> MiniCPM4GenerateModel<'a> { model_name, }) } + + /// 文本向量嵌入。 + /// + /// 将输入文本编码为固定长度向量(mean pool + L2 normalize)。 + pub fn embed_text(&mut self, text: &str) -> Result> { + let input_ids = self.tokenizer.text_encode(text.to_string(), &self.device)?; + let embedding = self.model.embed(&input_ids)?; + let vec: Vec = embedding.flatten_all()?.to_vec1()?; + Ok(vec) + } + + /// 批量文本向量嵌入。 + pub fn embed_text_batch(&mut self, texts: &[&str]) -> Result>> { + let mut results = Vec::with_capacity(texts.len()); + for text in texts { + results.push(self.embed_text(text)?); + } + Ok(results) + } } impl<'a> GenerationDataProvider for MiniCPM4GenerateModel<'a> { diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index c00343c..2ec974a 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -320,6 +320,60 @@ impl MiniCPMModel { layer.clear_kv_cache() } } + + /// 文本向量嵌入 — mean pool + L2 normalize。 + /// + /// 复用完整 forward pass(embed → layers → norm), + /// 在 lm_head 之前截取 hidden states,做均值池化后 L2 归一化。 + pub fn embed(&mut self, input_ids: &Tensor) -> Result { + let (bs, seq_len) = input_ids.dims2()?; + let input_embeds = self + .embed_tokens + .forward(input_ids)? + .affine(self.cfg.scale_emb, 0.0)?; + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + bs, + seq_len, + 0, + input_ids.device(), + )?) + } + }; + + let (cos, sin) = self.rope_emb.forward(0, seq_len)?; + let mut hidden_states = input_embeds; + for layer in &self.layers { + hidden_states = + layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?; + } + hidden_states = self.norm.forward(&hidden_states)?; + + // Mean pool across sequence dimension + let embedding = hidden_states.mean(1)?; + let embedding = embedding.affine( + 1.0 / (self.cfg.hidden_size / self.cfg.dim_model_base) as f64, + 0.0, + )?; + + // L2 normalize + let norm = embedding.sqr()?.sum_keepdim(1)?.sqrt()?; + let embedding = embedding.broadcast_div(&norm)?; + + Ok(embedding) + } + + /// 批量文本向量嵌入。 + pub fn embed_batch(&mut self, input_ids: &[&Tensor]) -> Result> { + let mut results = Vec::with_capacity(input_ids.len()); + for ids in input_ids { + results.push(self.embed(ids)?); + } + Ok(results) + } } impl InferenceModel for MiniCPMModel {