add MiniCPM5-1B
This commit is contained in:
@@ -16,7 +16,7 @@ use crate::models::minicpm4::model::MiniCPMModel;
|
||||
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||
|
||||
pub struct MiniCPMGenerateModel<'a> {
|
||||
pub struct MiniCPM4GenerateModel<'a> {
|
||||
chat_template: ChatTemplate<'a>,
|
||||
tokenizer: TokenizerModel,
|
||||
minicpm: MiniCPMModel,
|
||||
@@ -24,7 +24,7 @@ pub struct MiniCPMGenerateModel<'a> {
|
||||
model_name: String,
|
||||
}
|
||||
|
||||
impl<'a> MiniCPMGenerateModel<'a> {
|
||||
impl<'a> MiniCPM4GenerateModel<'a> {
|
||||
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||
let chat_template = ChatTemplate::init(path)?;
|
||||
let tokenizer = TokenizerModel::init(path)?;
|
||||
@@ -33,18 +33,15 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
||||
let device = &get_device(device);
|
||||
let cfg_dtype = cfg.torch_dtype.as_str();
|
||||
let dtype = get_dtype(dtype, cfg_dtype);
|
||||
let endoftext_id = cfg.eos_token_id[0];
|
||||
let im_end_id = cfg.eos_token_id[1];
|
||||
let model_list = find_type_files(path, "safetensors")?;
|
||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||
let eos_ids = vec![endoftext_id, im_end_id];
|
||||
let minicpm = MiniCPMModel::new(vb, cfg, eos_ids)?;
|
||||
let minicpm = MiniCPMModel::new(vb, cfg)?;
|
||||
let model_name = std::path::Path::new(path)
|
||||
.file_name()
|
||||
.and_then(|s| s.to_str())
|
||||
.unwrap_or("minicpm4")
|
||||
.to_string();
|
||||
Ok(MiniCPMGenerateModel {
|
||||
Ok(MiniCPM4GenerateModel {
|
||||
chat_template,
|
||||
tokenizer,
|
||||
minicpm,
|
||||
@@ -54,7 +51,7 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||
impl<'a> GenerateModel for MiniCPM4GenerateModel<'a> {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||
|
||||
@@ -214,7 +214,7 @@ pub struct MiniCPMModel {
|
||||
}
|
||||
|
||||
impl MiniCPMModel {
|
||||
pub fn new(vb: VarBuilder, cfg: MiniCPM4Config, eos_ids: Vec<u32>) -> Result<Self> {
|
||||
pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> Result<Self> {
|
||||
let vb = vb.pp("model");
|
||||
let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?;
|
||||
let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
|
||||
@@ -226,6 +226,7 @@ impl MiniCPMModel {
|
||||
let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?;
|
||||
let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device())?;
|
||||
let lm_head = Linear::new(embed_tokens.embeddings().clone(), None);
|
||||
let stop_token_ids = cfg.eos_token_id.clone();
|
||||
Ok(Self {
|
||||
cfg,
|
||||
embed_tokens,
|
||||
@@ -233,7 +234,7 @@ impl MiniCPMModel {
|
||||
norm,
|
||||
rope_emb,
|
||||
lm_head,
|
||||
stop_token_ids: eos_ids,
|
||||
stop_token_ids,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user