generate code refactoring progress 1/3
This commit is contained in:
Generated
+1
-1
@@ -21,7 +21,7 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "aha"
|
name = "aha"
|
||||||
version = "0.2.5"
|
version = "0.2.6"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"ahash",
|
"ahash",
|
||||||
"anyhow",
|
"anyhow",
|
||||||
|
|||||||
+2
-4
@@ -1,10 +1,10 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "aha"
|
name = "aha"
|
||||||
version = "0.2.5"
|
version = "0.2.6"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
repository = "https://github.com/jhqxxx/aha"
|
repository = "https://github.com/jhqxxx/aha"
|
||||||
license = "Apache-2.0"
|
license = "Apache-2.0"
|
||||||
description = "aha model inference library, now supports Qwen(2.5VL/3/3VL/3.5/ASR/3Embedding/3Reranker), MiniCPM4, VoxCPM(0.5B/1.5/2), DeepSeek-OCR/2, Hunyuan-OCR, PaddleOCR-VL/1.5, RMBG2.0, GLM(ASR-Nano-2512/OCR), Fun-ASR-Nano-2512, LFM(2/2.5/2VL/2.5VL)"
|
description = "aha model inference library, now supports Qwen(2.5VL/3/3VL/3.5/ASR/3Embedding/3Reranker), MiniCPM(4/5), VoxCPM(0.5B/1.5/2), DeepSeek-OCR/2, Hunyuan-OCR, PaddleOCR-VL/1.5, RMBG2.0, GLM(ASR-Nano-2512/OCR), Fun-ASR-Nano-2512, LFM(2/2.5/2VL/2.5VL)"
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
candle-core = { version = "0.9.2" }
|
candle-core = { version = "0.9.2" }
|
||||||
@@ -21,9 +21,7 @@ base64 = "0.22.1"
|
|||||||
num = "0.4.3"
|
num = "0.4.3"
|
||||||
minijinja = "2.12.0"
|
minijinja = "2.12.0"
|
||||||
tokenizers = "0.22.1"
|
tokenizers = "0.22.1"
|
||||||
# aha_openai_dive = { version = "1.4", features = ["stream"] }
|
|
||||||
uuid = { version = "1.18.1", features = ["v4"] }
|
uuid = { version = "1.18.1", features = ["v4"] }
|
||||||
# chrono = "0.4"
|
|
||||||
rocket = { version = "0.5.1", features = ["serde_json", "json"] }
|
rocket = { version = "0.5.1", features = ["serde_json", "json"] }
|
||||||
tokio = "1.47.1"
|
tokio = "1.47.1"
|
||||||
hound = "3.5.1"
|
hound = "3.5.1"
|
||||||
|
|||||||
@@ -39,6 +39,9 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
|||||||
| **Reranker** | Qwen3-Reranker |
|
| **Reranker** | Qwen3-Reranker |
|
||||||
|
|
||||||
## Changelog
|
## Changelog
|
||||||
|
### 2026-05-28
|
||||||
|
- generate code refactoring progress 1/3
|
||||||
|
|
||||||
### 2026-05-27
|
### 2026-05-27
|
||||||
- add MiniCPM5
|
- add MiniCPM5
|
||||||
|
|
||||||
@@ -54,11 +57,6 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
|||||||
### 2026-04-25
|
### 2026-04-25
|
||||||
- VoxCPM update stream
|
- VoxCPM update stream
|
||||||
|
|
||||||
### 2026-04-17
|
|
||||||
- Qwen3ASR add vad data recognition
|
|
||||||
|
|
||||||
### 2026-04-16
|
|
||||||
- fix FireRedVAD fsmn cache bug
|
|
||||||
|
|
||||||
**[View full changelog](docs/changelog.md)** →
|
**[View full changelog](docs/changelog.md)** →
|
||||||
|
|
||||||
|
|||||||
+3
-5
@@ -38,6 +38,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
|
|||||||
| **重排序** | Qwen3-Reranker |
|
| **重排序** | Qwen3-Reranker |
|
||||||
|
|
||||||
## 更新日志
|
## 更新日志
|
||||||
|
### 2026-05-28
|
||||||
|
- generate代码重构进度 1/3
|
||||||
|
|
||||||
### 2026-05-27
|
### 2026-05-27
|
||||||
- 新增 MiniCPM5
|
- 新增 MiniCPM5
|
||||||
|
|
||||||
@@ -53,11 +56,6 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
|
|||||||
### 2026-04-25
|
### 2026-04-25
|
||||||
- VoxCPM 更新流式生成
|
- VoxCPM 更新流式生成
|
||||||
|
|
||||||
### 2026-04-17
|
|
||||||
- Qwen3ASR 增加 vad 数据识别
|
|
||||||
|
|
||||||
### 2026-04-16
|
|
||||||
- 修复 FireRedVAD fsmn 缓存问题
|
|
||||||
|
|
||||||
**[查看完整更新日志](docs/changelog.zh-CN.md)** →
|
**[查看完整更新日志](docs/changelog.zh-CN.md)** →
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,9 @@ All notable changes to aha will be documented in this file.
|
|||||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
|
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
|
||||||
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||||
|
|
||||||
|
### 2026-05-28
|
||||||
|
- generate code refactoring progress 1/3
|
||||||
|
|
||||||
### 2026-05-27
|
### 2026-05-27
|
||||||
- add MiniCPM5
|
- add MiniCPM5
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,9 @@
|
|||||||
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/),
|
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/),
|
||||||
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
|
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
|
||||||
|
|
||||||
|
### 2026-05-28
|
||||||
|
- generate代码重构进度 1/3
|
||||||
|
|
||||||
### 2026-05-27
|
### 2026-05-27
|
||||||
- 新增 MiniCPM5
|
- 新增 MiniCPM5
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ use crate::{
|
|||||||
InferenceModel, MultiModalData,
|
InferenceModel, MultiModalData,
|
||||||
sample::{get_logit_processor, use_repeat_penalty},
|
sample::{get_logit_processor, use_repeat_penalty},
|
||||||
},
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::response_utils::{
|
utils::response_utils::{
|
||||||
build_chunk_response_with_reasoning, build_chunk_response_with_usage,
|
build_chunk_response_with_reasoning, build_chunk_response_with_usage,
|
||||||
@@ -366,3 +366,116 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
|||||||
};
|
};
|
||||||
Ok(stream)
|
Ok(stream)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub struct PrepareData {
|
||||||
|
pub in_reasoning: bool,
|
||||||
|
pub input_ids: Tensor,
|
||||||
|
pub multi_model_data: MultiModalData,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub trait GenerationDataProvider {
|
||||||
|
fn get_temperature(&self, req_temp: Option<f32>) -> Option<f32> {
|
||||||
|
req_temp
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
|
||||||
|
req_top_p
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
|
||||||
|
top_k
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_in_reasoning(&self, text: &str) -> bool {
|
||||||
|
text.ends_with("<think>\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_multi_model_data(&self) -> MultiModalData {
|
||||||
|
MultiModalData::new(vec![])
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData>;
|
||||||
|
}
|
||||||
|
|
||||||
|
#[macro_export]
|
||||||
|
macro_rules! impl_generate_model {
|
||||||
|
($struct_name: ty) => {
|
||||||
|
impl<'a> $crate::models::GenerateModel for $struct_name {
|
||||||
|
fn generate(
|
||||||
|
&mut self,
|
||||||
|
mes: $crate::params::chat::ChatCompletionParameters,
|
||||||
|
) -> anyhow::Result<$crate::params::chat::ChatCompletionResponse> {
|
||||||
|
let seed = mes.seed.unwrap_or(299792458) as u64;
|
||||||
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
let temperature = self.get_temperature(mes.temperature);
|
||||||
|
let top_p = self.get_top_p(mes.top_p);
|
||||||
|
let top_k = self.get_top_k(mes.top_k);
|
||||||
|
let prepare_data = self.get_data(&mes)?;
|
||||||
|
let input_ids = prepare_data.input_ids;
|
||||||
|
let data = prepare_data.multi_model_data;
|
||||||
|
let mut ctx = $crate::models::common::generate::GenerationContext::new(
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
mes.repeat_penalty,
|
||||||
|
mes.repeat_last_n,
|
||||||
|
seed,
|
||||||
|
input_ids.dim(1)?,
|
||||||
|
sample_len,
|
||||||
|
self.device.clone(),
|
||||||
|
);
|
||||||
|
|
||||||
|
$crate::models::common::generate::generate_generic(
|
||||||
|
&mut self.model,
|
||||||
|
&self.tokenizer,
|
||||||
|
input_ids,
|
||||||
|
data,
|
||||||
|
&mut ctx,
|
||||||
|
&self.model_name,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn generate_stream(
|
||||||
|
&mut self,
|
||||||
|
mes: $crate::params::chat::ChatCompletionParameters,
|
||||||
|
) -> anyhow::Result<
|
||||||
|
Box<
|
||||||
|
dyn rocket::futures::Stream<
|
||||||
|
Item = anyhow::Result<
|
||||||
|
$crate::params::chat::ChatCompletionChunkResponse,
|
||||||
|
>,
|
||||||
|
> + Send
|
||||||
|
+ Unpin
|
||||||
|
+ '_,
|
||||||
|
>,
|
||||||
|
> {
|
||||||
|
let seed = mes.seed.unwrap_or(299792458) as u64;
|
||||||
|
let prepare_data = self.get_data(&mes)?;
|
||||||
|
let input_ids = prepare_data.input_ids;
|
||||||
|
let data = prepare_data.multi_model_data;
|
||||||
|
let in_reasoning = prepare_data.in_reasoning;
|
||||||
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
let temperature = self.get_temperature(mes.temperature);
|
||||||
|
let top_p = self.get_top_p(mes.top_p);
|
||||||
|
let top_k = self.get_top_k(mes.top_k);
|
||||||
|
let stream = $crate::models::common::generate::generate_stream_generic(
|
||||||
|
&mut self.model,
|
||||||
|
&self.tokenizer,
|
||||||
|
input_ids,
|
||||||
|
data,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
mes.repeat_penalty,
|
||||||
|
mes.repeat_last_n,
|
||||||
|
seed,
|
||||||
|
sample_len,
|
||||||
|
in_reasoning,
|
||||||
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
|
)?;
|
||||||
|
Ok(Box::new(Box::pin(stream)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,21 +1,14 @@
|
|||||||
use crate::{
|
use crate::models::common::{
|
||||||
models::common::{
|
MultiModalData,
|
||||||
MultiModalData,
|
generate::{GenerationDataProvider, PrepareData},
|
||||||
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
|
||||||
},
|
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::futures::Stream;
|
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::deepseek_ocr::{
|
||||||
GenerateModel,
|
config::DeepseekOCRConfig, model::DeepseekOCRModel, processor::DeepseekOCRProcessor,
|
||||||
deepseek_ocr::{
|
|
||||||
config::DeepseekOCRConfig, model::DeepseekOCRModel, processor::DeepseekOCRProcessor,
|
|
||||||
},
|
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{extract_metadata_value, find_type_files, get_device, get_dtype},
|
utils::{extract_metadata_value, find_type_files, get_device, get_dtype},
|
||||||
@@ -24,9 +17,7 @@ use crate::{
|
|||||||
pub struct DeepseekOCRGenerateModel {
|
pub struct DeepseekOCRGenerateModel {
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
processor: DeepseekOCRProcessor,
|
processor: DeepseekOCRProcessor,
|
||||||
deepseekocr_model: DeepseekOCRModel,
|
model: DeepseekOCRModel,
|
||||||
// bos_token_id: u32,
|
|
||||||
// eos_token_id: u32,
|
|
||||||
device: Device,
|
device: Device,
|
||||||
size: Vec<u32>,
|
size: Vec<u32>,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -52,19 +43,15 @@ impl DeepseekOCRGenerateModel {
|
|||||||
1usize
|
1usize
|
||||||
};
|
};
|
||||||
let processor = DeepseekOCRProcessor::new(device, dtype, version)?;
|
let processor = DeepseekOCRProcessor::new(device, dtype, version)?;
|
||||||
// let eos_token_id = cfg.eos_token_id;
|
|
||||||
// let bos_token_id = cfg.bos_token_id;
|
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let deepseekocr_model = DeepseekOCRModel::new(vb, cfg, version)?;
|
let model = DeepseekOCRModel::new(vb, cfg, version)?;
|
||||||
let size = vec![512u32, 640, 1024, 1280];
|
let size = vec![512u32, 640, 1024, 1280];
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
deepseekocr_model,
|
model,
|
||||||
// bos_token_id,
|
|
||||||
// eos_token_id,
|
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
size,
|
size,
|
||||||
model_name: model_name.to_string(),
|
model_name: model_name.to_string(),
|
||||||
@@ -73,8 +60,8 @@ impl DeepseekOCRGenerateModel {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GenerateModel for DeepseekOCRGenerateModel {
|
impl GenerationDataProvider for DeepseekOCRGenerateModel {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result<PrepareData> {
|
||||||
let base_size = extract_metadata_value::<u32>(&mes.metadata, "base_size").unwrap_or(640);
|
let base_size = extract_metadata_value::<u32>(&mes.metadata, "base_size").unwrap_or(640);
|
||||||
let base_size = if self.size.contains(&base_size) {
|
let base_size = if self.size.contains(&base_size) {
|
||||||
base_size
|
base_size
|
||||||
@@ -92,93 +79,20 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
|||||||
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
|
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
|
||||||
let (input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
let (input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
||||||
.processor
|
.processor
|
||||||
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
.process_info(mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
||||||
let max_tokens = mes.max_tokens.unwrap_or(1024);
|
|
||||||
let mut ctx = GenerationContext::new(
|
|
||||||
mes.temperature,
|
|
||||||
mes.top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
mes.seed.unwrap_or(34562) as u64,
|
|
||||||
input_ids.dim(1)?,
|
|
||||||
max_tokens,
|
|
||||||
self.device.clone(),
|
|
||||||
);
|
|
||||||
let data_vec = vec![
|
let data_vec = vec![
|
||||||
Some(images_ori),
|
Some(images_ori),
|
||||||
Some(image_crop),
|
Some(image_crop),
|
||||||
Some(images_seq_mask),
|
Some(images_seq_mask),
|
||||||
Some(images_spatial_crop_t),
|
Some(images_spatial_crop_t),
|
||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let multi_model_data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
Ok(PrepareData {
|
||||||
&mut self.deepseekocr_model,
|
in_reasoning: false,
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
multi_model_data,
|
||||||
&mut ctx,
|
})
|
||||||
&self.model_name,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn generate_stream(
|
|
||||||
&mut self,
|
|
||||||
mes: ChatCompletionParameters,
|
|
||||||
) -> Result<
|
|
||||||
Box<
|
|
||||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
|
||||||
+ Send
|
|
||||||
+ Unpin
|
|
||||||
+ '_,
|
|
||||||
>,
|
|
||||||
> {
|
|
||||||
let base_size = extract_metadata_value::<u32>(&mes.metadata, "base_size").unwrap_or(640);
|
|
||||||
let base_size = if self.size.contains(&base_size) {
|
|
||||||
base_size
|
|
||||||
} else {
|
|
||||||
640
|
|
||||||
};
|
|
||||||
let image_size = extract_metadata_value::<u32>(&mes.metadata, "image_size").unwrap_or(640);
|
|
||||||
let image_size = if self.size.contains(&image_size) {
|
|
||||||
image_size
|
|
||||||
} else {
|
|
||||||
640
|
|
||||||
};
|
|
||||||
let base_size = if self.version == 2 { 1024 } else { base_size };
|
|
||||||
let image_size = if self.version == 2 { 768 } else { image_size };
|
|
||||||
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
|
|
||||||
let (input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
|
||||||
.processor
|
|
||||||
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
|
||||||
let data_vec = vec![
|
|
||||||
images_ori.into(),
|
|
||||||
image_crop.into(),
|
|
||||||
images_seq_mask.into(),
|
|
||||||
images_spatial_crop_t.into(),
|
|
||||||
];
|
|
||||||
let data = MultiModalData::new(data_vec);
|
|
||||||
|
|
||||||
let temperature = mes.temperature;
|
|
||||||
let top_p = mes.top_p;
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let max_tokens = mes.max_tokens.unwrap_or(1024);
|
|
||||||
let stream = generate_stream_generic(
|
|
||||||
&mut self.deepseekocr_model,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
temperature,
|
|
||||||
top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
max_tokens,
|
|
||||||
false,
|
|
||||||
&self.device,
|
|
||||||
&self.model_name,
|
|
||||||
)?;
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
crate::impl_generate_model!(DeepseekOCRGenerateModel);
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ use crate::{
|
|||||||
pub struct FunAsrNanoGenerateModel {
|
pub struct FunAsrNanoGenerateModel {
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
processor: FunAsrNanoProcessor,
|
processor: FunAsrNanoProcessor,
|
||||||
fun_asr_nano: FunAsrNanoModel,
|
model: FunAsrNanoModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
@@ -75,7 +75,7 @@ impl FunAsrNanoGenerateModel {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device);
|
let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device);
|
||||||
let fun_asr_nano =
|
let model =
|
||||||
FunAsrNanoModel::new(vb, &cfg, &llm_cfg, generation_config.eos_token_id.clone())?;
|
FunAsrNanoModel::new(vb, &cfg, &llm_cfg, generation_config.eos_token_id.clone())?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
@@ -85,7 +85,7 @@ impl FunAsrNanoGenerateModel {
|
|||||||
Ok(Self {
|
Ok(Self {
|
||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
fun_asr_nano,
|
model,
|
||||||
device,
|
device,
|
||||||
dtype,
|
dtype,
|
||||||
generation_config,
|
generation_config,
|
||||||
@@ -120,7 +120,7 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
|||||||
let data_vec = vec![speech.into(), fbank_mask.into()];
|
let data_vec = vec![speech.into(), fbank_mask.into()];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
generate_generic(
|
||||||
&mut self.fun_asr_nano,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -152,7 +152,7 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
|||||||
let data_vec = vec![speech.into(), fbank_mask.into()];
|
let data_vec = vec![speech.into(), fbank_mask.into()];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let stream = generate_stream_generic(
|
let stream = generate_stream_generic(
|
||||||
&mut self.fun_asr_nano,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ pub struct GlmAsrNanoGenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
processor: GlmAsrNanoProcessor,
|
processor: GlmAsrNanoProcessor,
|
||||||
glm_asr_nano: GlmAsrNanoModel,
|
model: GlmAsrNanoModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -45,7 +45,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
|
|||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let eos_ids = vec![59246u32, 59253, 59255];
|
let eos_ids = vec![59246u32, 59253, 59255];
|
||||||
let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg, eos_ids)?;
|
let model = GlmAsrNanoModel::new(vb, cfg, eos_ids)?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -55,7 +55,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
glm_asr_nano,
|
model,
|
||||||
device,
|
device,
|
||||||
dtype,
|
dtype,
|
||||||
model_name,
|
model_name,
|
||||||
@@ -88,7 +88,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
|
|||||||
let data_vec = vec![input_features.into(), audio_token_lengths.into()];
|
let data_vec = vec![input_features.into(), audio_token_lengths.into()];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
generate_generic(
|
||||||
&mut self.glm_asr_nano,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -119,7 +119,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
|
|||||||
let data_vec = vec![input_features.into(), audio_token_lengths.into()];
|
let data_vec = vec![input_features.into(), audio_token_lengths.into()];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let stream = generate_stream_generic(
|
let stream = generate_stream_generic(
|
||||||
&mut self.glm_asr_nano,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ pub struct HunyuanOCRGenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
pre_processor: HunyuanVLProcessor,
|
pre_processor: HunyuanVLProcessor,
|
||||||
hunyuan_vl: HunyuanVLModel,
|
model: HunyuanVLModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
generation_config: HunyuanOCRGenerationConfig,
|
generation_config: HunyuanOCRGenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -49,8 +49,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
|
|||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: HunyuanOCRGenerationConfig =
|
let generation_config: HunyuanOCRGenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
let hunyuan_vl =
|
let model = HunyuanVLModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?;
|
||||||
HunyuanVLModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?;
|
|
||||||
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
@@ -61,7 +60,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor,
|
pre_processor,
|
||||||
hunyuan_vl,
|
model,
|
||||||
device,
|
device,
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
@@ -103,7 +102,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
|||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
generate_generic(
|
||||||
&mut self.hunyuan_vl,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -144,7 +143,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
|||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let stream = generate_stream_generic(
|
let stream = generate_stream_generic(
|
||||||
&mut self.hunyuan_vl,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
|
|||||||
+15
-72
@@ -1,16 +1,9 @@
|
|||||||
use crate::models::common::MultiModalData;
|
use crate::models::common::generate::{GenerationDataProvider, PrepareData};
|
||||||
use crate::models::common::generate::{
|
|
||||||
GenerationContext, generate_generic, generate_stream_generic,
|
|
||||||
};
|
|
||||||
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
|
|
||||||
use crate::{
|
use crate::{
|
||||||
chat_template::ChatTemplate,
|
chat_template::ChatTemplate,
|
||||||
models::{
|
models::lfm2::{
|
||||||
GenerateModel,
|
config::{Lfm2Config, Lfm2GenerateConfig},
|
||||||
lfm2::{
|
model::Lfm2Model,
|
||||||
config::{Lfm2Config, Lfm2GenerateConfig},
|
|
||||||
model::Lfm2Model,
|
|
||||||
},
|
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{find_type_files, get_device, get_dtype},
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
@@ -62,68 +55,18 @@ impl<'a> Lfm2GenerateModel<'a> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
|
impl<'a> GenerationDataProvider for Lfm2GenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result<PrepareData> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(mes)?;
|
||||||
|
let in_reasoning = self.is_in_reasoning(&mes_render);
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let multi_model_data = self.get_multi_model_data();
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
Ok(PrepareData {
|
||||||
let mut ctx = GenerationContext::new(
|
in_reasoning,
|
||||||
mes.temperature,
|
|
||||||
mes.top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
input_ids.dim(1)?,
|
|
||||||
sample_len,
|
|
||||||
self.device.clone(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let data = MultiModalData::new(vec![]);
|
|
||||||
generate_generic(
|
|
||||||
&mut self.model,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
multi_model_data,
|
||||||
&mut ctx,
|
})
|
||||||
&self.model_name,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn generate_stream(
|
|
||||||
&mut self,
|
|
||||||
mes: ChatCompletionParameters,
|
|
||||||
) -> Result<
|
|
||||||
Box<
|
|
||||||
dyn rocket::futures::Stream<
|
|
||||||
Item = Result<crate::params::chat::ChatCompletionChunkResponse, anyhow::Error>,
|
|
||||||
> + Send
|
|
||||||
+ Unpin
|
|
||||||
+ '_,
|
|
||||||
>,
|
|
||||||
> {
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
|
||||||
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.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
sample_len,
|
|
||||||
false,
|
|
||||||
&self.device,
|
|
||||||
&self.model_name,
|
|
||||||
)?;
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
crate::impl_generate_model!(Lfm2GenerateModel<'a>);
|
||||||
|
|||||||
@@ -1,25 +1,18 @@
|
|||||||
use crate::models::common::MultiModalData;
|
use crate::models::common::generate::{GenerationDataProvider, PrepareData};
|
||||||
use crate::models::common::generate::{
|
|
||||||
GenerationContext, generate_generic, generate_stream_generic,
|
|
||||||
};
|
|
||||||
use crate::params::chat::{
|
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
|
||||||
};
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::futures::Stream;
|
|
||||||
|
|
||||||
use crate::models::minicpm4::config::MiniCPM4Config;
|
use crate::models::minicpm4::config::MiniCPM4Config;
|
||||||
use crate::models::minicpm4::model::MiniCPMModel;
|
use crate::models::minicpm4::model::MiniCPMModel;
|
||||||
// use crate::models::GenerateStream;
|
|
||||||
use crate::utils::{find_type_files, get_device, get_dtype};
|
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
pub struct MiniCPM4GenerateModel<'a> {
|
pub struct MiniCPM4GenerateModel<'a> {
|
||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
minicpm: MiniCPMModel,
|
model: MiniCPMModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
@@ -35,7 +28,7 @@ impl<'a> MiniCPM4GenerateModel<'a> {
|
|||||||
let dtype = get_dtype(dtype, cfg_dtype);
|
let dtype = get_dtype(dtype, cfg_dtype);
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let minicpm = MiniCPMModel::new(vb, cfg)?;
|
let model = MiniCPMModel::new(vb, cfg)?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -44,73 +37,25 @@ impl<'a> MiniCPM4GenerateModel<'a> {
|
|||||||
Ok(MiniCPM4GenerateModel {
|
Ok(MiniCPM4GenerateModel {
|
||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
minicpm,
|
model,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for MiniCPM4GenerateModel<'a> {
|
impl<'a> GenerationDataProvider for MiniCPM4GenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result<PrepareData> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(mes)?;
|
||||||
|
let in_reasoning = self.is_in_reasoning(&mes_render);
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let multi_model_data = self.get_multi_model_data();
|
||||||
let sample_len = mes.max_tokens.unwrap_or(2048);
|
Ok(PrepareData {
|
||||||
let mut ctx = GenerationContext::new(
|
in_reasoning,
|
||||||
mes.temperature,
|
|
||||||
mes.top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
input_ids.dim(1)?,
|
|
||||||
sample_len,
|
|
||||||
self.device.clone(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let data = MultiModalData::new(vec![]);
|
|
||||||
generate_generic(
|
|
||||||
&mut self.minicpm,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
multi_model_data,
|
||||||
&mut ctx,
|
})
|
||||||
&self.model_name,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
fn generate_stream(
|
|
||||||
&mut self,
|
|
||||||
mes: ChatCompletionParameters,
|
|
||||||
) -> Result<
|
|
||||||
Box<
|
|
||||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
|
||||||
+ Send
|
|
||||||
+ Unpin
|
|
||||||
+ '_,
|
|
||||||
>,
|
|
||||||
> {
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
|
||||||
let data = MultiModalData::new(vec![]);
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
|
||||||
let stream = generate_stream_generic(
|
|
||||||
&mut self.minicpm,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
mes.temperature,
|
|
||||||
mes.top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
sample_len,
|
|
||||||
false,
|
|
||||||
&self.device,
|
|
||||||
&self.model_name,
|
|
||||||
)?;
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
crate::impl_generate_model!(MiniCPM4GenerateModel<'a>);
|
||||||
|
|||||||
@@ -1,19 +1,13 @@
|
|||||||
use crate::models::common::MultiModalData;
|
use crate::models::common::generate::{GenerationDataProvider, PrepareData};
|
||||||
use crate::models::common::generate::{
|
|
||||||
GenerationContext, generate_generic, generate_stream_generic,
|
|
||||||
};
|
|
||||||
use crate::models::llama::LlamaForCausalLM;
|
use crate::models::llama::LlamaForCausalLM;
|
||||||
use crate::models::minicpm5::config::MiniCPM5Config;
|
use crate::models::minicpm5::config::MiniCPM5Config;
|
||||||
use crate::params::chat::{
|
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
|
||||||
};
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::futures::Stream;
|
|
||||||
|
|
||||||
use crate::utils::{find_type_files, get_device, get_dtype};
|
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
pub struct MiniCPM5GenerateModel<'a> {
|
pub struct MiniCPM5GenerateModel<'a> {
|
||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
@@ -70,66 +64,18 @@ impl<'a> MiniCPM5GenerateModel<'a> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for MiniCPM5GenerateModel<'a> {
|
impl<'a> GenerationDataProvider for MiniCPM5GenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result<PrepareData> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(mes)?;
|
||||||
|
let in_reasoning = self.is_in_reasoning(&mes_render);
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let multi_model_data = self.get_multi_model_data();
|
||||||
let sample_len = mes.max_tokens.unwrap_or(2048);
|
Ok(PrepareData {
|
||||||
let mut ctx = GenerationContext::new(
|
in_reasoning,
|
||||||
mes.temperature,
|
|
||||||
mes.top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
input_ids.dim(1)?,
|
|
||||||
sample_len,
|
|
||||||
self.device.clone(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let data = MultiModalData::new(vec![]);
|
|
||||||
generate_generic(
|
|
||||||
&mut self.model,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
multi_model_data,
|
||||||
&mut ctx,
|
})
|
||||||
&self.model_name,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
fn generate_stream(
|
|
||||||
&mut self,
|
|
||||||
mes: ChatCompletionParameters,
|
|
||||||
) -> Result<
|
|
||||||
Box<
|
|
||||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
|
||||||
+ Send
|
|
||||||
+ Unpin
|
|
||||||
+ '_,
|
|
||||||
>,
|
|
||||||
> {
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
|
||||||
let data = MultiModalData::new(vec![]);
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
|
||||||
let stream = generate_stream_generic(
|
|
||||||
&mut self.model,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
mes.temperature,
|
|
||||||
mes.top_p,
|
|
||||||
None,
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
sample_len,
|
|
||||||
false,
|
|
||||||
&self.device,
|
|
||||||
&self.model_name,
|
|
||||||
)?;
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
crate::impl_generate_model!(MiniCPM5GenerateModel<'a>);
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ pub struct PaddleOCRVLGenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
pre_processor: PaddleOCRVLProcessor,
|
pre_processor: PaddleOCRVLProcessor,
|
||||||
paddleocr_vl: PaddleOCRVLModel,
|
model: PaddleOCRVLModel,
|
||||||
cfg: PaddleOCRVLConfig,
|
cfg: PaddleOCRVLConfig,
|
||||||
device: Device,
|
device: Device,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -42,7 +42,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
|
|||||||
let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?;
|
let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?;
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb, vec![2])?;
|
let model = PaddleOCRVLModel::new(cfg.clone(), vb, vec![2])?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -52,7 +52,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor,
|
pre_processor,
|
||||||
paddleocr_vl,
|
model,
|
||||||
cfg,
|
cfg,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
model_name,
|
model_name,
|
||||||
@@ -94,7 +94,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
|
|||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
generate_generic(
|
||||||
&mut self.paddleocr_vl,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -136,7 +136,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
|
|||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let stream = generate_stream_generic(
|
let stream = generate_stream_generic(
|
||||||
&mut self.paddleocr_vl,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ pub struct Qwen2_5VLGenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
pre_processor: Qwen2_5VLProcessor,
|
pre_processor: Qwen2_5VLProcessor,
|
||||||
qwen2_5_vl: Qwen2_5VLModel,
|
model: Qwen2_5VLModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
endoftext_id: u32,
|
endoftext_id: u32,
|
||||||
im_end_id: u32,
|
im_end_id: u32,
|
||||||
@@ -52,7 +52,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
|
|||||||
// let model_list = find_safetensors_files(&path)?;
|
// let model_list = find_safetensors_files(&path)?;
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?;
|
let model = Qwen2_5VLModel::new(cfg, vb)?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -62,7 +62,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor,
|
pre_processor,
|
||||||
qwen2_5_vl,
|
model,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
endoftext_id,
|
endoftext_id,
|
||||||
im_end_id,
|
im_end_id,
|
||||||
@@ -102,7 +102,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
let mut completion_secs = 0.0f64;
|
let mut completion_secs = 0.0f64;
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let logits = self.qwen2_5_vl.forward(
|
let logits = self.model.forward(
|
||||||
&input_ids,
|
&input_ids,
|
||||||
pixel_values,
|
pixel_values,
|
||||||
image_grid_thw,
|
image_grid_thw,
|
||||||
@@ -136,7 +136,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
}
|
}
|
||||||
let num_token = generate.len() as u32;
|
let num_token = generate.len() as u32;
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
let res = self.tokenizer.token_decode(generate)?;
|
||||||
self.qwen2_5_vl.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
let response = build_completion_response_with_time(
|
let response = build_completion_response_with_time(
|
||||||
res,
|
res,
|
||||||
&self.model_name,
|
&self.model_name,
|
||||||
@@ -190,7 +190,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
let mut tool_call_content = String::new();
|
let mut tool_call_content = String::new();
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let logits = self.qwen2_5_vl.forward(
|
let logits = self.model.forward(
|
||||||
&input_ids,
|
&input_ids,
|
||||||
pixel_values,
|
pixel_values,
|
||||||
image_grid_thw,
|
image_grid_thw,
|
||||||
@@ -293,7 +293,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
pixel_values = None;
|
pixel_values = None;
|
||||||
pixel_values_video = None;
|
pixel_values_video = None;
|
||||||
}
|
}
|
||||||
self.qwen2_5_vl.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
};
|
};
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,24 +1,18 @@
|
|||||||
use crate::models::common::MultiModalData;
|
use crate::models::common::generate::{GenerationDataProvider, PrepareData};
|
||||||
use crate::models::common::generate::{
|
|
||||||
GenerationContext, generate_generic, generate_stream_generic,
|
|
||||||
};
|
|
||||||
use crate::params::chat::{
|
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
|
||||||
};
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::futures::Stream;
|
|
||||||
|
|
||||||
use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
|
use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
|
||||||
use crate::models::qwen3::model::Qwen3Model;
|
use crate::models::qwen3::model::Qwen3Model;
|
||||||
use crate::utils::{find_type_files, get_device, get_dtype};
|
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
pub struct Qwen3GenerateModel<'a> {
|
pub struct Qwen3GenerateModel<'a> {
|
||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
qwen3: Qwen3Model,
|
model: Qwen3Model,
|
||||||
device: Device,
|
device: Device,
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -38,7 +32,7 @@ impl<'a> Qwen3GenerateModel<'a> {
|
|||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: Qwen3GenerationConfig =
|
let generation_config: Qwen3GenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
let qwen3 = Qwen3Model::new(&cfg, vb, generation_config.eos_token_id.clone())?;
|
let model = Qwen3Model::new(&cfg, vb, generation_config.eos_token_id.clone())?;
|
||||||
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
@@ -48,7 +42,7 @@ impl<'a> Qwen3GenerateModel<'a> {
|
|||||||
Ok(Qwen3GenerateModel {
|
Ok(Qwen3GenerateModel {
|
||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
qwen3,
|
model,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
@@ -56,77 +50,30 @@ impl<'a> Qwen3GenerateModel<'a> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
|
impl<'a> GenerationDataProvider for Qwen3GenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn get_temperature(&self, req_temp: Option<f32>) -> Option<f32> {
|
||||||
let temperature = mes
|
Some(req_temp.unwrap_or(self.generation_config.temperature))
|
||||||
.temperature
|
|
||||||
.unwrap_or(self.generation_config.temperature);
|
|
||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
|
||||||
let top_k = self.generation_config.top_k;
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(2048);
|
|
||||||
let mut ctx = GenerationContext::new(
|
|
||||||
temperature.into(),
|
|
||||||
top_p.into(),
|
|
||||||
top_k.into(),
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
input_ids.dim(1)?,
|
|
||||||
sample_len,
|
|
||||||
self.device.clone(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let data = MultiModalData::new(vec![]);
|
|
||||||
generate_generic(
|
|
||||||
&mut self.qwen3,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
&mut ctx,
|
|
||||||
&self.model_name,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
fn generate_stream(
|
|
||||||
&mut self,
|
fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
|
||||||
mes: ChatCompletionParameters,
|
Some(req_top_p.unwrap_or(self.generation_config.top_p))
|
||||||
) -> Result<
|
}
|
||||||
Box<
|
|
||||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
|
||||||
+ Send
|
Some(top_k.unwrap_or(self.generation_config.top_k))
|
||||||
+ Unpin
|
}
|
||||||
+ '_,
|
|
||||||
>,
|
fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result<PrepareData> {
|
||||||
> {
|
let mes_render = self.chat_template.apply_chat_template(mes)?;
|
||||||
let temperature = mes
|
let in_reasoning = self.is_in_reasoning(&mes_render);
|
||||||
.temperature
|
|
||||||
.unwrap_or(self.generation_config.temperature);
|
|
||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
|
||||||
let top_k = self.generation_config.top_k;
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
|
||||||
let in_reasoning = mes_render.ends_with("<think>\n");
|
|
||||||
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let data = MultiModalData::new(vec![]);
|
let multi_model_data = self.get_multi_model_data();
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
Ok(PrepareData {
|
||||||
let stream = generate_stream_generic(
|
|
||||||
&mut self.qwen3,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
temperature.into(),
|
|
||||||
top_p.into(),
|
|
||||||
top_k.into(),
|
|
||||||
mes.repeat_penalty,
|
|
||||||
mes.repeat_last_n,
|
|
||||||
seed,
|
|
||||||
sample_len,
|
|
||||||
in_reasoning,
|
in_reasoning,
|
||||||
&self.device,
|
input_ids,
|
||||||
&self.model_name,
|
multi_model_data,
|
||||||
)?;
|
})
|
||||||
Ok(Box::new(Box::pin(stream)))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
crate::impl_generate_model!(Qwen3GenerateModel<'a>);
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ pub struct Qwen3_5GenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
pre_processor: Option<Qwen3VLProcessor>,
|
pre_processor: Option<Qwen3VLProcessor>,
|
||||||
qwen3_5: Qwen3_5Model,
|
model: Qwen3_5Model,
|
||||||
device: Device,
|
device: Device,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
repeat_penalty: f32,
|
repeat_penalty: f32,
|
||||||
@@ -53,13 +53,13 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let eos_ids = vec![cfg.text_config.eos_token_id];
|
let eos_ids = vec![cfg.text_config.eos_token_id];
|
||||||
let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?;
|
let model = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?;
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor: Some(pre_processor),
|
pre_processor: Some(pre_processor),
|
||||||
qwen3_5,
|
model,
|
||||||
device,
|
device,
|
||||||
model_name: model_name.to_string(),
|
model_name: model_name.to_string(),
|
||||||
repeat_penalty: 1.0,
|
repeat_penalty: 1.0,
|
||||||
@@ -89,13 +89,13 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let eos_ids = vec![cfg.text_config.eos_token_id];
|
let eos_ids = vec![cfg.text_config.eos_token_id];
|
||||||
// let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?;
|
// let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?;
|
||||||
let qwen3_5 = Qwen3_5Model::new_from_vb_without_visual(vb, cfg, eos_ids)?;
|
let model = Qwen3_5Model::new_from_vb_without_visual(vb, cfg, eos_ids)?;
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor,
|
pre_processor,
|
||||||
qwen3_5,
|
model,
|
||||||
device,
|
device,
|
||||||
model_name: model_name.to_string(),
|
model_name: model_name.to_string(),
|
||||||
repeat_penalty: 1.0,
|
repeat_penalty: 1.0,
|
||||||
@@ -142,7 +142,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
.get_matedata("tokenizer.ggml.eos_token_id")?
|
.get_matedata("tokenizer.ggml.eos_token_id")?
|
||||||
.to_u32()?;
|
.to_u32()?;
|
||||||
let eos_ids = vec![eos_token_id];
|
let eos_ids = vec![eos_token_id];
|
||||||
let qwen3_5 =
|
let model =
|
||||||
Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device, eos_ids)?;
|
Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device, eos_ids)?;
|
||||||
let stem = std::path::Path::new(model_file)
|
let stem = std::path::Path::new(model_file)
|
||||||
.file_stem() // 获取文件名主干(不含扩展名)
|
.file_stem() // 获取文件名主干(不含扩展名)
|
||||||
@@ -152,7 +152,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor,
|
pre_processor,
|
||||||
qwen3_5,
|
model,
|
||||||
device,
|
device,
|
||||||
model_name: stem.to_string(),
|
model_name: stem.to_string(),
|
||||||
repeat_penalty: 1.2,
|
repeat_penalty: 1.2,
|
||||||
@@ -199,13 +199,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
|
|
||||||
generate_generic_text(
|
generate_generic_text(&mut self.model, &self.tokenizer, input_ids, data, &mut ctx)
|
||||||
&mut self.qwen3_5,
|
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
&mut ctx,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn generate_stream_text(
|
pub fn generate_stream_text(
|
||||||
@@ -237,7 +231,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
generate_stream_generic_text(
|
generate_stream_generic_text(
|
||||||
&mut self.qwen3_5,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -293,7 +287,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
|||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
generate_generic(
|
||||||
&mut self.qwen3_5,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -339,7 +333,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
|||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let stream = generate_stream_generic(
|
let stream = generate_stream_generic(
|
||||||
&mut self.qwen3_5,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
|
|||||||
@@ -37,7 +37,7 @@ pub struct Qwen3AsrGenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
processor: Qwen3AsrProcessor,
|
processor: Qwen3AsrProcessor,
|
||||||
qwen3_asr: Qwen3ASRModel,
|
model: Qwen3ASRModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
eos_token_id1: u32,
|
eos_token_id1: u32,
|
||||||
@@ -65,7 +65,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
|
|||||||
let dtype = get_dtype(dtype, cfg_dtype);
|
let dtype = get_dtype(dtype, cfg_dtype);
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let qwen3_asr = Qwen3ASRModel::new(vb, &cfg, generation_config.eos_token_id.clone())?;
|
let model = Qwen3ASRModel::new(vb, &cfg, generation_config.eos_token_id.clone())?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -75,7 +75,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
qwen3_asr,
|
model,
|
||||||
device,
|
device,
|
||||||
dtype,
|
dtype,
|
||||||
eos_token_id1: generation_config.eos_token_id[0] as u32,
|
eos_token_id1: generation_config.eos_token_id[0] as u32,
|
||||||
@@ -116,13 +116,8 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
|
|||||||
);
|
);
|
||||||
let data_vec = vec![input_features];
|
let data_vec = vec![input_features];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let mut text = generate_generic_text(
|
let mut text =
|
||||||
&mut self.qwen3_asr,
|
generate_generic_text(&mut self.model, &self.tokenizer, input_ids, data, &mut ctx)?;
|
||||||
&self.tokenizer,
|
|
||||||
input_ids,
|
|
||||||
data,
|
|
||||||
&mut ctx,
|
|
||||||
)?;
|
|
||||||
if text.contains("<asr_text>") {
|
if text.contains("<asr_text>") {
|
||||||
let mut split: Vec<&str> = text.split("<asr_text>").collect();
|
let mut split: Vec<&str> = text.split("<asr_text>").collect();
|
||||||
text = split.pop().unwrap_or(&text).to_string();
|
text = split.pop().unwrap_or(&text).to_string();
|
||||||
@@ -156,7 +151,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let logits =
|
let logits =
|
||||||
self.qwen3_asr
|
self.model
|
||||||
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
@@ -175,7 +170,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
input_features = None;
|
input_features = None;
|
||||||
}
|
}
|
||||||
self.qwen3_asr.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
let num_token = generate.len() as u32;
|
let num_token = generate.len() as u32;
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
let res = self.tokenizer.token_decode(generate)?;
|
||||||
@@ -226,7 +221,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let logits =
|
let logits =
|
||||||
self.qwen3_asr
|
self.model
|
||||||
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
@@ -266,7 +261,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
input_features = None;
|
input_features = None;
|
||||||
}
|
}
|
||||||
self.qwen3_asr.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ pub struct Qwen3VLGenerateModel<'a> {
|
|||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
pre_processor: Qwen3VLProcessor,
|
pre_processor: Qwen3VLProcessor,
|
||||||
qwen3_vl: Qwen3VLModel,
|
model: Qwen3VLModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -46,7 +46,7 @@ impl<'a> Qwen3VLGenerateModel<'a> {
|
|||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: Qwen3GenerationConfig =
|
let generation_config: Qwen3GenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
let qwen3_vl = Qwen3VLModel::new(cfg, vb, generation_config.eos_token_id.clone())?;
|
let model = Qwen3VLModel::new(cfg, vb, generation_config.eos_token_id.clone())?;
|
||||||
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
@@ -57,7 +57,7 @@ impl<'a> Qwen3VLGenerateModel<'a> {
|
|||||||
chat_template,
|
chat_template,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
pre_processor,
|
pre_processor,
|
||||||
qwen3_vl,
|
model,
|
||||||
device,
|
device,
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
@@ -101,7 +101,7 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
|||||||
];
|
];
|
||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
generate_generic(
|
generate_generic(
|
||||||
&mut self.qwen3_vl,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
@@ -145,7 +145,7 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
|||||||
let data = MultiModalData::new(data_vec);
|
let data = MultiModalData::new(data_vec);
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let stream = generate_stream_generic(
|
let stream = generate_stream_generic(
|
||||||
&mut self.qwen3_vl,
|
&mut self.model,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
input_ids,
|
input_ids,
|
||||||
data,
|
data,
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
pub struct VoxCPMGenerate {
|
pub struct VoxCPMGenerate {
|
||||||
voxcpm: VoxCPMModel,
|
model: VoxCPMModel,
|
||||||
prompt_cache: Option<HashMap<String, Tensor>>,
|
prompt_cache: Option<HashMap<String, Tensor>>,
|
||||||
out_sample_rate: usize,
|
out_sample_rate: usize,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -106,12 +106,12 @@ impl VoxCPMGenerate {
|
|||||||
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
|
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
|
||||||
};
|
};
|
||||||
let tokenizer = SingleChineseTokenizer::new(path)?;
|
let tokenizer = SingleChineseTokenizer::new(path)?;
|
||||||
let voxcpm = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?;
|
let model = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?;
|
||||||
let out_sample_rate = audio_config
|
let out_sample_rate = audio_config
|
||||||
.out_sample_rate
|
.out_sample_rate
|
||||||
.unwrap_or(audio_config.sample_rate);
|
.unwrap_or(audio_config.sample_rate);
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
voxcpm,
|
model,
|
||||||
prompt_cache: None,
|
prompt_cache: None,
|
||||||
out_sample_rate,
|
out_sample_rate,
|
||||||
model_name,
|
model_name,
|
||||||
@@ -124,7 +124,7 @@ impl VoxCPMGenerate {
|
|||||||
prompt_wav_path: String,
|
prompt_wav_path: String,
|
||||||
) -> Result<()> {
|
) -> Result<()> {
|
||||||
let cache = self
|
let cache = self
|
||||||
.voxcpm
|
.model
|
||||||
.build_prompt_cache(prompt_text, prompt_wav_path)?;
|
.build_prompt_cache(prompt_text, prompt_wav_path)?;
|
||||||
self.prompt_cache = Some(cache);
|
self.prompt_cache = Some(cache);
|
||||||
Ok(())
|
Ok(())
|
||||||
@@ -143,7 +143,7 @@ impl VoxCPMGenerate {
|
|||||||
let audio = match &self.prompt_cache {
|
let audio = match &self.prompt_cache {
|
||||||
Some(cache) => {
|
Some(cache) => {
|
||||||
let prompt_cache = cache.clone();
|
let prompt_cache = cache.clone();
|
||||||
self.voxcpm.generate_with_prompt_cache(
|
self.model.generate_with_prompt_cache(
|
||||||
target_text,
|
target_text,
|
||||||
prompt_cache,
|
prompt_cache,
|
||||||
min_len,
|
min_len,
|
||||||
@@ -156,7 +156,7 @@ impl VoxCPMGenerate {
|
|||||||
}
|
}
|
||||||
None => self.generate_simple(target_text)?,
|
None => self.generate_simple(target_text)?,
|
||||||
};
|
};
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
Ok(audio)
|
Ok(audio)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -196,7 +196,7 @@ impl VoxCPMGenerate {
|
|||||||
// retry_badcase: bool,
|
// retry_badcase: bool,
|
||||||
retry_badcase_ratio_threshold: f64,
|
retry_badcase_ratio_threshold: f64,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let audio = self.voxcpm.generate(
|
let audio = self.model.generate(
|
||||||
target_text,
|
target_text,
|
||||||
prompt_text,
|
prompt_text,
|
||||||
prompt_wav_path,
|
prompt_wav_path,
|
||||||
@@ -207,7 +207,7 @@ impl VoxCPMGenerate {
|
|||||||
// retry_badcase,
|
// retry_badcase,
|
||||||
retry_badcase_ratio_threshold,
|
retry_badcase_ratio_threshold,
|
||||||
)?;
|
)?;
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
Ok(audio)
|
Ok(audio)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -250,7 +250,7 @@ impl GenerateModel for VoxCPMGenerate {
|
|||||||
target_text = format!("({instruction}){target_text}");
|
target_text = format!("({instruction}){target_text}");
|
||||||
}
|
}
|
||||||
let audio = self
|
let audio = self
|
||||||
.voxcpm
|
.model
|
||||||
.generate(
|
.generate(
|
||||||
target_text,
|
target_text,
|
||||||
prompt_text,
|
prompt_text,
|
||||||
@@ -262,13 +262,13 @@ impl GenerateModel for VoxCPMGenerate {
|
|||||||
retry_badcase_ratio_threshold,
|
retry_badcase_ratio_threshold,
|
||||||
)
|
)
|
||||||
.inspect_err(|_| {
|
.inspect_err(|_| {
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
})?;
|
})?;
|
||||||
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
|
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
|
||||||
// let wave_u8_str = String::from_utf8(wav_u8)?;
|
// let wave_u8_str = String::from_utf8(wav_u8)?;
|
||||||
let base64_audio = BASE64_STANDARD.encode(wav_u8);
|
let base64_audio = BASE64_STANDARD.encode(wav_u8);
|
||||||
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
Ok(response)
|
Ok(response)
|
||||||
}
|
}
|
||||||
#[allow(unused_variables)]
|
#[allow(unused_variables)]
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
pub struct VoxCPMGenerateRefact {
|
pub struct VoxCPMGenerateRefact {
|
||||||
voxcpm: VoxCPMModelRefact,
|
model: VoxCPMModelRefact,
|
||||||
tokenizer: SingleChineseTokenizer,
|
tokenizer: SingleChineseTokenizer,
|
||||||
audio_vae: AudioVAE,
|
audio_vae: AudioVAE,
|
||||||
processor: VoxCPMProcessor,
|
processor: VoxCPMProcessor,
|
||||||
@@ -113,13 +113,13 @@ impl VoxCPMGenerateRefact {
|
|||||||
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
|
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
|
||||||
};
|
};
|
||||||
let tokenizer = SingleChineseTokenizer::new(path)?;
|
let tokenizer = SingleChineseTokenizer::new(path)?;
|
||||||
let voxcpm =
|
let model =
|
||||||
VoxCPMModelRefact::new(vb_voxcpm, config, audio_vae.latent_dim, decode_chunk_size)?;
|
VoxCPMModelRefact::new(vb_voxcpm, config, audio_vae.latent_dim, decode_chunk_size)?;
|
||||||
let out_sample_rate = audio_config
|
let out_sample_rate = audio_config
|
||||||
.out_sample_rate
|
.out_sample_rate
|
||||||
.unwrap_or(audio_config.sample_rate);
|
.unwrap_or(audio_config.sample_rate);
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
voxcpm,
|
model,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
audio_vae,
|
audio_vae,
|
||||||
processor,
|
processor,
|
||||||
@@ -162,7 +162,7 @@ impl VoxCPMGenerateRefact {
|
|||||||
} else {
|
} else {
|
||||||
max_len
|
max_len
|
||||||
};
|
};
|
||||||
let audio = self.voxcpm.inference(
|
let audio = self.model.inference(
|
||||||
&text_token,
|
&text_token,
|
||||||
audio_feat.as_ref(),
|
audio_feat.as_ref(),
|
||||||
audio_mask.as_ref(),
|
audio_mask.as_ref(),
|
||||||
@@ -172,7 +172,7 @@ impl VoxCPMGenerateRefact {
|
|||||||
cfg_value,
|
cfg_value,
|
||||||
&self.audio_vae,
|
&self.audio_vae,
|
||||||
)?;
|
)?;
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
Ok(audio)
|
Ok(audio)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -240,7 +240,7 @@ impl VoxCPMGenerateRefact {
|
|||||||
} else {
|
} else {
|
||||||
max_len
|
max_len
|
||||||
};
|
};
|
||||||
self.voxcpm.inference(
|
self.model.inference(
|
||||||
&text_token,
|
&text_token,
|
||||||
audio_feat.as_ref(),
|
audio_feat.as_ref(),
|
||||||
audio_mask.as_ref(),
|
audio_mask.as_ref(),
|
||||||
@@ -255,7 +255,7 @@ impl VoxCPMGenerateRefact {
|
|||||||
return Err(anyhow!("need prompt_cache"));
|
return Err(anyhow!("need prompt_cache"));
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
Ok(audio)
|
Ok(audio)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -284,7 +284,7 @@ impl VoxCPMGenerateRefact {
|
|||||||
} else {
|
} else {
|
||||||
max_len
|
max_len
|
||||||
};
|
};
|
||||||
self.voxcpm.inference_stream(
|
self.model.inference_stream(
|
||||||
text_token,
|
text_token,
|
||||||
audio_feat,
|
audio_feat,
|
||||||
audio_mask,
|
audio_mask,
|
||||||
@@ -344,13 +344,13 @@ impl GenerateModel for VoxCPMGenerateRefact {
|
|||||||
retry_badcase_ratio_threshold,
|
retry_badcase_ratio_threshold,
|
||||||
)
|
)
|
||||||
.inspect_err(|_| {
|
.inspect_err(|_| {
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
})?;
|
})?;
|
||||||
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
|
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
|
||||||
// let wave_u8_str = String::from_utf8(wav_u8)?;
|
// let wave_u8_str = String::from_utf8(wav_u8)?;
|
||||||
let base64_audio = BASE64_STANDARD.encode(wav_u8);
|
let base64_audio = BASE64_STANDARD.encode(wav_u8);
|
||||||
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
||||||
self.voxcpm.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
Ok(response)
|
Ok(response)
|
||||||
}
|
}
|
||||||
#[allow(unused_variables)]
|
#[allow(unused_variables)]
|
||||||
|
|||||||
+41
-2
@@ -1,10 +1,11 @@
|
|||||||
use std::time::Instant;
|
use std::{pin::pin, time::Instant};
|
||||||
|
|
||||||
use aha::{
|
use aha::{
|
||||||
models::{GenerateModel, minicpm5::generate::MiniCPM5GenerateModel},
|
models::{GenerateModel, minicpm5::generate::MiniCPM5GenerateModel},
|
||||||
params::chat::ChatCompletionParameters,
|
params::chat::ChatCompletionParameters,
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
|
use rocket::futures::StreamExt;
|
||||||
#[test]
|
#[test]
|
||||||
fn minicpm5_generate() -> Result<()> {
|
fn minicpm5_generate() -> Result<()> {
|
||||||
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_minicpm5 minicpm5_generate -r -- --nocapture
|
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_minicpm5 minicpm5_generate -r -- --nocapture
|
||||||
@@ -16,7 +17,7 @@ fn minicpm5_generate() -> Result<()> {
|
|||||||
{
|
{
|
||||||
"temperature": 0.3,
|
"temperature": 0.3,
|
||||||
"top_p": 0.8,
|
"top_p": 0.8,
|
||||||
"model": "minicpm4",
|
"model": "minicpm5",
|
||||||
"messages": [
|
"messages": [
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
@@ -39,3 +40,41 @@ fn minicpm5_generate() -> Result<()> {
|
|||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn minicpm5_stream() -> Result<()> {
|
||||||
|
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_minicpm5 minicpm5_stream -r -- --nocapture
|
||||||
|
|
||||||
|
let save_dir =
|
||||||
|
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
||||||
|
let model_path = format!("{}/OpenBMB/MiniCPM5-1B/", save_dir);
|
||||||
|
|
||||||
|
let message = r#"
|
||||||
|
{
|
||||||
|
"model": "minicpm5",
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": "什么是AI"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"enable_thinking": true
|
||||||
|
}
|
||||||
|
"#;
|
||||||
|
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
|
let i_start = Instant::now();
|
||||||
|
let mut model = MiniCPM5GenerateModel::init(&model_path, None, None)?;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
|
let i_start = Instant::now();
|
||||||
|
let mut stream = pin!(model.generate_stream(mes)?);
|
||||||
|
while let Some(item) = stream.next().await {
|
||||||
|
println!("generate: \n {:?}", item);
|
||||||
|
}
|
||||||
|
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|||||||
+1
-1
@@ -42,7 +42,7 @@ fn qwen3_0_6b_generate() -> Result<()> {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn qwen3_0_6b_stream() -> Result<()> {
|
async fn qwen3_0_6b_stream() -> Result<()> {
|
||||||
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3_0_6b_stream -r -- --nocapture
|
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3 qwen3_0_6b_stream -r -- --nocapture
|
||||||
|
|
||||||
let save_dir =
|
let save_dir =
|
||||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
||||||
|
|||||||
Reference in New Issue
Block a user