feat(embed): MiniCPM4 embed/embed_batch + MiniCPM4GenerateModel embed_text/embed_text_batch
- 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:
@@ -42,6 +42,25 @@ impl<'a> MiniCPM4GenerateModel<'a> {
|
||||
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> {
|
||||
|
||||
Reference in New Issue
Block a user