refactor generate code

This commit is contained in:
jhqxxx
2026-04-02 22:22:52 +08:00
parent b254d21efc
commit bd8eee6520
59 changed files with 1745 additions and 1800 deletions
+40 -79
View File
@@ -1,6 +1,8 @@
use crate::models::common::generate::get_logit_processor;
use crate::models::common::MultiModalData;
use crate::models::common::generate::{
GenerationContext, generate_generic, generate_stream_generic,
};
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
use crate::utils::build_completion_chunk_response;
use crate::{
chat_template::ChatTemplate,
models::{
@@ -11,19 +13,17 @@ use crate::{
},
},
tokenizer::TokenizerModel,
utils::{build_completion_response, find_type_files, get_device, get_dtype},
utils::{find_type_files, get_device, get_dtype},
};
use anyhow::Result;
use candle_core::{DType, Device, Tensor};
use candle_core::{DType, Device};
use candle_nn::VarBuilder;
use rocket::async_stream::stream;
pub struct Lfm2GenerateModel<'a> {
chat_template: ChatTemplate<'a>,
tokenizer: TokenizerModel,
device: Device,
model: Lfm2Model,
eos_token_id: u32,
model_name: String,
}
impl<'a> Lfm2GenerateModel<'a> {
@@ -45,8 +45,8 @@ impl<'a> Lfm2GenerateModel<'a> {
};
let dtype = get_dtype(dtype, &cfg_dtype);
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? };
let model = Lfm2Model::new(vb, &cfg)?;
let eos_token_id = gen_cfg.eos_token_id;
let eos_ids = vec![gen_cfg.eos_token_id];
let model = Lfm2Model::new(vb, &cfg, eos_ids)?;
let model_name = std::path::Path::new(path)
.file_name()
.and_then(|s| s.to_str())
@@ -57,7 +57,6 @@ impl<'a> Lfm2GenerateModel<'a> {
tokenizer,
device,
model,
eos_token_id,
model_name,
})
}
@@ -66,40 +65,28 @@ impl<'a> Lfm2GenerateModel<'a> {
impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mut logits = get_logit_processor(
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
let seed = mes.seed.unwrap_or(34562) as u64;
let mut ctx = GenerationContext::new(
mes.temperature,
mes.top_p,
None,
mes.seed.unwrap_or(34562) as u64,
seed,
input_ids.dim(1)?,
sample_len,
self.device.clone(),
);
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0;
let mut generate = vec![];
let sample_len = mes.max_tokens.unwrap_or(1024);
for _ in 0..sample_len {
let logit = self.model.forward(&input_ids, seqlen_offset)?;
let logit = logit.squeeze(0)?.squeeze(0)?;
let next_token = logits.sample(&logit)?;
generate.push(next_token);
if next_token == self.eos_token_id {
break;
}
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
seqlen_offset += seq_len;
seq_len = 1;
}
self.model.clear_cache();
let completion_tokens = generate.len() as u32;
let decode = self.tokenizer.token_decode(generate)?;
let mes = build_completion_response(
decode,
let data = MultiModalData::new(vec![]);
generate_generic(
&mut self.model,
&self.tokenizer,
input_ids,
data,
&mut ctx,
&self.model_name,
Some(completion_tokens),
Some(prompt_tokens),
);
Ok(mes)
)
}
fn generate_stream(
@@ -115,50 +102,24 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
>,
> {
let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mut logits = get_logit_processor(
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
let data = MultiModalData::new(vec![]);
let seed = mes.seed.unwrap_or(34562) as u64;
let stream = generate_stream_generic(
&mut self.model,
&self.tokenizer,
input_ids,
data,
mes.temperature,
mes.top_p,
None,
mes.seed.unwrap_or(34562) as u64,
);
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let sample_len = mes.max_tokens.unwrap_or(1024);
let stream = stream! {
let mut err_tokens = vec![];
for _ in 0..sample_len {
let logit = self.model.forward(&input_ids, seqlen_offset)?;
let logit = logit.squeeze(0)?.squeeze(0)?;
let next_token = logits.sample(&logit)?;
let mut decode_ids = vec![];
if !err_tokens.is_empty() {
decode_ids.extend_from_slice(&err_tokens);
}
decode_ids.push(next_token);
let decode = self.tokenizer.token_decode(decode_ids)?;
if decode.contains("") {
err_tokens.push(next_token);
if err_tokens.len() > 3 {
err_tokens.clear();
}
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
seqlen_offset += seq_len;
seq_len = 1;
continue;
}
err_tokens.clear();
let chunk = build_completion_chunk_response(decode, &self.model_name, None, None);
yield Ok(chunk);
if next_token == self.eos_token_id {
break;
}
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
seqlen_offset += seq_len;
seq_len = 1;
}
self.model.clear_cache();
};
seed,
sample_len,
false,
&self.device,
&self.model_name,
)?;
Ok(Box::new(Box::pin(stream)))
}
}