qwen3.5 add generate text methods
This commit is contained in:
@@ -1,7 +1,9 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::common::{
|
models::common::{
|
||||||
MultiModalData,
|
MultiModalData,
|
||||||
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
generate::{
|
||||||
|
GenerationContext, generate_generic, generate_generic_text, generate_stream_generic,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
@@ -156,6 +158,54 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
repeat_last_n: 64,
|
repeat_last_n: 64,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn generate_text(&mut self, mes: ChatCompletionParameters) -> Result<String> {
|
||||||
|
let seed = mes.seed.unwrap_or(32768) as u64;
|
||||||
|
let temperature = mes.temperature.unwrap_or(0.4);
|
||||||
|
let top_p = mes.top_p.unwrap_or(0.95);
|
||||||
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
|
let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) =
|
||||||
|
if let Some(processor) = &self.pre_processor {
|
||||||
|
let input = processor.process_info(&mes, &mes_render)?;
|
||||||
|
(
|
||||||
|
input.replace_text,
|
||||||
|
input.pixel_values,
|
||||||
|
input.image_grid_thw,
|
||||||
|
input.pixel_values_video,
|
||||||
|
input.video_grid_thw,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
(mes_render, None, None, None, None)
|
||||||
|
};
|
||||||
|
let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
|
||||||
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
let mut ctx = GenerationContext::new(
|
||||||
|
temperature.into(),
|
||||||
|
top_p.into(),
|
||||||
|
Some(20),
|
||||||
|
self.repeat_penalty.into(),
|
||||||
|
self.repeat_last_n.into(),
|
||||||
|
seed,
|
||||||
|
input_ids.dim(1)?,
|
||||||
|
sample_len,
|
||||||
|
self.device.clone(),
|
||||||
|
);
|
||||||
|
let data_vec = vec![
|
||||||
|
pixel_values,
|
||||||
|
image_grid_thw,
|
||||||
|
pixel_values_video,
|
||||||
|
video_grid_thw,
|
||||||
|
];
|
||||||
|
let data = MultiModalData::new(data_vec);
|
||||||
|
|
||||||
|
generate_generic_text(
|
||||||
|
&mut self.qwen3_5,
|
||||||
|
&self.tokenizer,
|
||||||
|
input_ids,
|
||||||
|
data,
|
||||||
|
&mut ctx,
|
||||||
|
)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
||||||
|
|||||||
+8
-12
@@ -19,12 +19,7 @@ fn qwen3_5_generate_no_visual() -> Result<()> {
|
|||||||
"messages": [
|
"messages": [
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
"content": [
|
"content": "你好啊"
|
||||||
{
|
|
||||||
"type": "text",
|
|
||||||
"text": "你好啊"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
@@ -35,12 +30,13 @@ fn qwen3_5_generate_no_visual() -> Result<()> {
|
|||||||
let mut qwen3_5 = Qwen3_5GenerateModel::init_without_visual(&model_path, None, None)?;
|
let mut qwen3_5 = Qwen3_5GenerateModel::init_without_visual(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
let text = qwen3_5.generate_text(mes)?;
|
||||||
let res = qwen3_5.generate(mes)?;
|
println!("text: {}", text);
|
||||||
println!("generate: \n {:?}", res);
|
// let res = qwen3_5.generate(mes)?;
|
||||||
if let Some(usage) = &res.usage {
|
// println!("generate: \n {:?}", res);
|
||||||
println!("usage: \n {:?}", usage);
|
// if let Some(usage) = &res.usage {
|
||||||
}
|
// println!("usage: \n {:?}", usage);
|
||||||
|
// }
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user