feat(embed): MiniCPM4 embed/embed_batch + MiniCPM4GenerateModel embed_text/embed_text_batch
ci / cargo fmt (push) Failing after 2m43s
ci / cargo clippy (push) Failing after 5m9s
ci / build and test (push) Successful in 8m32s

- MiniCPMModel::embed: run forward through layers + norm, mean pool, L2 normalize
- MiniCPMModel::embed_batch: batch version
- MiniCPM4GenerateModel::embed_text: tokenize + embed + return Vec<f32>
- MiniCPM4GenerateModel::embed_text_batch: batch version
This commit is contained in:
dengxuan
2026-07-25 13:03:25 +08:00
parent e29ddc589d
commit 2df0929440
2 changed files with 73 additions and 0 deletions
+19
View File
@@ -42,6 +42,25 @@ impl<'a> MiniCPM4GenerateModel<'a> {
model_name, model_name,
}) })
} }
/// 文本向量嵌入。
///
/// 将输入文本编码为固定长度向量(mean pool + L2 normalize)。
pub fn embed_text(&mut self, text: &str) -> Result<Vec<f32>> {
let input_ids = self.tokenizer.text_encode(text.to_string(), &self.device)?;
let embedding = self.model.embed(&input_ids)?;
let vec: Vec<f32> = embedding.flatten_all()?.to_vec1()?;
Ok(vec)
}
/// 批量文本向量嵌入。
pub fn embed_text_batch(&mut self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
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> { impl<'a> GenerationDataProvider for MiniCPM4GenerateModel<'a> {
+54
View File
@@ -320,6 +320,60 @@ impl MiniCPMModel {
layer.clear_kv_cache() layer.clear_kv_cache()
} }
} }
/// 文本向量嵌入 — mean pool + L2 normalize。
///
/// 复用完整 forward passembed → layers → norm),
/// 在 lm_head 之前截取 hidden states,做均值池化后 L2 归一化。
pub fn embed(&mut self, input_ids: &Tensor) -> Result<Tensor> {
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<Tensor> = {
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<Vec<Tensor>> {
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 { impl InferenceModel for MiniCPMModel {