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,
|
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> {
|
||||||
|
|||||||
@@ -320,6 +320,60 @@ impl MiniCPMModel {
|
|||||||
layer.clear_kv_cache()
|
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 {
|
impl InferenceModel for MiniCPMModel {
|
||||||
|
|||||||
Reference in New Issue
Block a user