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> {
|
||||
|
||||
@@ -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<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 {
|
||||
|
||||
Reference in New Issue
Block a user