refactor deepseek_ocr/fun_asr_nano generate code
This commit is contained in:
@@ -46,6 +46,8 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
|||||||
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
|
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
|
||||||
|
|
||||||
## Changelog
|
## Changelog
|
||||||
|
### 2026-04-01
|
||||||
|
- refactor deepseek_ocr/fun_asr_nano generate code
|
||||||
|
|
||||||
### 2026-03-31
|
### 2026-03-31
|
||||||
- add server adn cli mod
|
- add server adn cli mod
|
||||||
|
|||||||
@@ -45,6 +45,8 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
|
|||||||
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
|
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
|
||||||
|
|
||||||
## 更新日志
|
## 更新日志
|
||||||
|
### 2026-04-01
|
||||||
|
- 重构 deepseek_ocr/fun_asr_nano 生成代码
|
||||||
|
|
||||||
### 2026-03-31
|
### 2026-03-31
|
||||||
- 新增 server 和 cli 模块
|
- 新增 server 和 cli 模块
|
||||||
|
|||||||
@@ -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-04-01
|
||||||
|
- refactor deepseek_ocr/fun_asr_nano generate code
|
||||||
|
|
||||||
### 2026-03-31
|
### 2026-03-31
|
||||||
- add server adn cli mod
|
- add server adn cli mod
|
||||||
- aha model name use modelscope id replace
|
- aha model name use modelscope id replace
|
||||||
|
|||||||
@@ -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-04-01
|
||||||
|
- 重构 deepseek_ocr/fun_asr_nano 生成代码
|
||||||
|
|
||||||
### 2026-03-31
|
### 2026-03-31
|
||||||
- 新增 server 和 cli 模块
|
- 新增 server 和 cli 模块
|
||||||
- aha模型名称使用 modelscope id 替换
|
- aha模型名称使用 modelscope id 替换
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ show_help() {
|
|||||||
echo " Tencent-Hunyuan/HunyuanOCR"
|
echo " Tencent-Hunyuan/HunyuanOCR"
|
||||||
echo " PaddlePaddle/PaddleOCR-VL"
|
echo " PaddlePaddle/PaddleOCR-VL"
|
||||||
echo " AI-ModelScope/RMBG-2.0"
|
echo " AI-ModelScope/RMBG-2.0"
|
||||||
echo " voxcpm"
|
echo " OpenBMB/VoxCPM-0.5B"
|
||||||
echo " OpenBMB/VoxCPM1.5"
|
echo " OpenBMB/VoxCPM1.5"
|
||||||
echo " ZhipuAI/GLM-ASR-Nano-2512"
|
echo " ZhipuAI/GLM-ASR-Nano-2512"
|
||||||
echo " FunAudioLLM/Fun-ASR-Nano-2512"
|
echo " FunAudioLLM/Fun-ASR-Nano-2512"
|
||||||
@@ -88,7 +88,7 @@ case $MODEL_ALIAS in
|
|||||||
"AI-ModelScope/RMBG-2.0")
|
"AI-ModelScope/RMBG-2.0")
|
||||||
MODEL_ID="briaai/RMBG-2.0"
|
MODEL_ID="briaai/RMBG-2.0"
|
||||||
;;
|
;;
|
||||||
"voxcpm")
|
"OpenBMB/VoxCPM-0.5B")
|
||||||
MODEL_ID="openbmb/VoxCPM-0.5B"
|
MODEL_ID="openbmb/VoxCPM-0.5B"
|
||||||
;;
|
;;
|
||||||
"OpenBMB/VoxCPM1.5")
|
"OpenBMB/VoxCPM1.5")
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ use candle_nn::{Init, VarBuilder};
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
bigvgan::config::BigVGANConfig,
|
bigvgan::config::BigVGANConfig,
|
||||||
common::{WNConv1d, WNConvTranspose1d},
|
common::modules::{WNConv1d, WNConvTranspose1d},
|
||||||
},
|
},
|
||||||
utils::tensor_utils::pad_replicate_last_dim,
|
utils::tensor_utils::pad_replicate_last_dim,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ use candle_core::{D, Tensor};
|
|||||||
use candle_nn::{BatchNorm, Conv1d, Conv2d, Module, ModuleT, VarBuilder, ops::sigmoid};
|
use candle_nn::{BatchNorm, Conv1d, Conv2d, Module, ModuleT, VarBuilder, ops::sigmoid};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::common::{get_batch_norm, get_conv1d, get_conv2d},
|
models::common::modules::{get_batch_norm, get_conv1d, get_conv2d},
|
||||||
utils::tensor_utils::{pool1d, statistics_pooling},
|
utils::tensor_utils::{pool1d, statistics_pooling},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,217 @@
|
|||||||
|
use anyhow::Result;
|
||||||
|
use candle_core::{DType, Device, Tensor};
|
||||||
|
use candle_transformers::generation::{LogitsProcessor, Sampling};
|
||||||
|
use rocket::async_stream::stream;
|
||||||
|
use rocket::futures::Stream;
|
||||||
|
use std::time::Instant;
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
models::common::{InferenceModel, MultiModalData},
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
|
||||||
|
tokenizer::TokenizerModel,
|
||||||
|
utils::{build_completion_chunk_response, build_completion_response_with_time},
|
||||||
|
};
|
||||||
|
pub fn get_logit_processor(
|
||||||
|
temperature: Option<f32>,
|
||||||
|
top_p: Option<f32>,
|
||||||
|
top_k: Option<usize>,
|
||||||
|
seed: u64,
|
||||||
|
) -> LogitsProcessor {
|
||||||
|
let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) });
|
||||||
|
match top_k {
|
||||||
|
None => LogitsProcessor::new(
|
||||||
|
seed,
|
||||||
|
temperature.map(|temp| temp as f64),
|
||||||
|
top_p.map(|tp| tp as f64),
|
||||||
|
),
|
||||||
|
Some(k) => {
|
||||||
|
let sampling = match temperature {
|
||||||
|
None => Sampling::ArgMax,
|
||||||
|
Some(temperature) => match top_p {
|
||||||
|
None => Sampling::TopK {
|
||||||
|
k,
|
||||||
|
temperature: temperature as f64,
|
||||||
|
},
|
||||||
|
Some(p) => Sampling::TopKThenTopP {
|
||||||
|
k,
|
||||||
|
p: p as f64,
|
||||||
|
temperature: temperature as f64,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
};
|
||||||
|
LogitsProcessor::from_sampling(seed, sampling)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct GenerationContext {
|
||||||
|
pub logit_processor: LogitsProcessor,
|
||||||
|
pub seqlen_offset: usize,
|
||||||
|
pub seq_len: usize,
|
||||||
|
pub sample_len: u32,
|
||||||
|
pub device: Device,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl GenerationContext {
|
||||||
|
pub fn new(
|
||||||
|
temperature: Option<f32>,
|
||||||
|
top_p: Option<f32>,
|
||||||
|
top_k: Option<usize>,
|
||||||
|
seed: u64,
|
||||||
|
initial_seq_len: usize,
|
||||||
|
max_tokens: u32,
|
||||||
|
device: Device,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
logit_processor: get_logit_processor(temperature, top_p, top_k, seed),
|
||||||
|
seqlen_offset: 0,
|
||||||
|
seq_len: initial_seq_len,
|
||||||
|
sample_len: max_tokens,
|
||||||
|
device,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn prepare_for_next_token(&mut self, token: u32) -> Result<Tensor> {
|
||||||
|
self.update_status();
|
||||||
|
self.create_input_ids(token)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn update_status(&mut self) {
|
||||||
|
self.seqlen_offset += self.seq_len;
|
||||||
|
self.seq_len = 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
fn create_input_ids(&self, token: u32) -> Result<Tensor> {
|
||||||
|
Ok(Tensor::from_vec(vec![token], (1, 1), &self.device)?)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 采样辅助函数
|
||||||
|
fn sample_and_push(
|
||||||
|
processor: &mut LogitsProcessor,
|
||||||
|
logits: &Tensor,
|
||||||
|
generated: &mut Vec<u32>,
|
||||||
|
) -> Result<u32> {
|
||||||
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
|
let token = processor.sample(&logits)?;
|
||||||
|
generated.push(token);
|
||||||
|
Ok(token)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn generate_generic<M: InferenceModel>(
|
||||||
|
model: &mut M,
|
||||||
|
tokenizer: &TokenizerModel,
|
||||||
|
input_ids: Tensor,
|
||||||
|
data: MultiModalData,
|
||||||
|
ctx: &mut GenerationContext,
|
||||||
|
model_name: &str,
|
||||||
|
) -> Result<ChatCompletionResponse> {
|
||||||
|
let prompt_tokens = ctx.seq_len as u32;
|
||||||
|
let mut generated = Vec::new();
|
||||||
|
let eos_ids = model.stop_token_ids();
|
||||||
|
let i_start = Instant::now();
|
||||||
|
let logits = model.forward_initial(&input_ids, ctx.seqlen_offset, data)?;
|
||||||
|
let next_token = sample_and_push(&mut ctx.logit_processor, &logits, &mut generated)?;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
let prompt_secs = i_duration.as_secs_f64();
|
||||||
|
let mut input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
|
||||||
|
// 自回归循环
|
||||||
|
let i_start = Instant::now();
|
||||||
|
for _ in 1..ctx.sample_len {
|
||||||
|
let logits = model.forward_step(&input_ids, ctx.seqlen_offset)?;
|
||||||
|
let next_token = sample_and_push(&mut ctx.logit_processor, &logits, &mut generated)?;
|
||||||
|
|
||||||
|
if eos_ids.contains(&next_token) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
}
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
let completion_secs = i_duration.as_secs_f64();
|
||||||
|
|
||||||
|
model.clear_cache();
|
||||||
|
|
||||||
|
let num_tokens = generated.len() as u32;
|
||||||
|
let text = tokenizer.token_decode(generated)?;
|
||||||
|
Ok(build_completion_response_with_time(
|
||||||
|
text,
|
||||||
|
model_name,
|
||||||
|
Some(num_tokens),
|
||||||
|
Some(completion_secs),
|
||||||
|
Some(prompt_tokens),
|
||||||
|
Some(prompt_secs),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn generate_stream_generic<M: InferenceModel>(
|
||||||
|
model: &mut M,
|
||||||
|
tokenizer: &TokenizerModel,
|
||||||
|
input_ids: Tensor,
|
||||||
|
data: MultiModalData,
|
||||||
|
temperature: Option<f32>,
|
||||||
|
top_p: Option<f32>,
|
||||||
|
top_k: Option<usize>,
|
||||||
|
seed: u64,
|
||||||
|
max_tokens: u32,
|
||||||
|
device: &Device,
|
||||||
|
model_name: &str,
|
||||||
|
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||||
|
let mut ctx = GenerationContext::new(
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
seed,
|
||||||
|
input_ids.dim(1)?,
|
||||||
|
max_tokens,
|
||||||
|
device.clone(),
|
||||||
|
);
|
||||||
|
let mut error_tokens = Vec::new();
|
||||||
|
let eos_ids = model.stop_token_ids();
|
||||||
|
let stream = stream! {
|
||||||
|
let mut input_ids = input_ids;
|
||||||
|
// 处理 unicode 错误累积
|
||||||
|
for _ in 0..ctx.sample_len {
|
||||||
|
let logits = if ctx.seqlen_offset == 0 {
|
||||||
|
model.forward_initial(&input_ids, ctx.seqlen_offset, data.clone())
|
||||||
|
} else {
|
||||||
|
model.forward_step(&input_ids, ctx.seqlen_offset)
|
||||||
|
}?;
|
||||||
|
|
||||||
|
let next_token = {
|
||||||
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
|
ctx.logit_processor.sample(&logits)?
|
||||||
|
};
|
||||||
|
|
||||||
|
// 解码(处理�的累积)
|
||||||
|
let decode_ids = if error_tokens.is_empty() {
|
||||||
|
vec![next_token]
|
||||||
|
} else {
|
||||||
|
let mut ids = error_tokens.clone();
|
||||||
|
ids.push(next_token);
|
||||||
|
ids
|
||||||
|
};
|
||||||
|
|
||||||
|
let decoded = tokenizer.token_decode(decode_ids)?;
|
||||||
|
|
||||||
|
if decoded.contains("�") {
|
||||||
|
error_tokens.push(next_token);
|
||||||
|
if error_tokens.len() > 3 {
|
||||||
|
error_tokens.clear();
|
||||||
|
}
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
error_tokens.clear();
|
||||||
|
yield Ok(build_completion_chunk_response(decoded, model_name, None, None));
|
||||||
|
|
||||||
|
if eos_ids.contains(&next_token) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
}
|
||||||
|
model.clear_cache();
|
||||||
|
};
|
||||||
|
Ok(stream)
|
||||||
|
}
|
||||||
+26
-1324
File diff suppressed because it is too large
Load Diff
@@ -95,6 +95,7 @@ impl WhichModel {
|
|||||||
| WhichModel::Qwen3_0_6B
|
| WhichModel::Qwen3_0_6B
|
||||||
| WhichModel::LFM2_1_2B
|
| WhichModel::LFM2_1_2B
|
||||||
| WhichModel::LFM2_5_1_2BInstruct => "llm",
|
| WhichModel::LFM2_5_1_2BInstruct => "llm",
|
||||||
|
// VLM models
|
||||||
WhichModel::Qwen2_5VL3B
|
WhichModel::Qwen2_5VL3B
|
||||||
| WhichModel::Qwen2_5VL7B
|
| WhichModel::Qwen2_5VL7B
|
||||||
| WhichModel::Qwen3VL2B
|
| WhichModel::Qwen3VL2B
|
||||||
@@ -122,6 +123,7 @@ impl WhichModel {
|
|||||||
| WhichModel::FunASRNano2512 => "asr",
|
| WhichModel::FunASRNano2512 => "asr",
|
||||||
// Image models
|
// Image models
|
||||||
WhichModel::RMBG2_0 => "image",
|
WhichModel::RMBG2_0 => "image",
|
||||||
|
// TTS models
|
||||||
WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "tts",
|
WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "tts",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -1,10 +1,13 @@
|
|||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -15,18 +18,15 @@ use crate::{
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{extract_metadata_value, find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, extract_metadata_value,
|
|
||||||
find_type_files, get_device, get_dtype, get_logit_processor,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct DeepseekOCRGenerateModel {
|
pub struct DeepseekOCRGenerateModel {
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
processor: DeepseekOCRProcessor,
|
processor: DeepseekOCRProcessor,
|
||||||
deepseekocr_model: DeepseekOCRModel,
|
deepseekocr_model: DeepseekOCRModel,
|
||||||
bos_token_id: u32,
|
// bos_token_id: u32,
|
||||||
eos_token_id: u32,
|
// eos_token_id: u32,
|
||||||
device: Device,
|
device: Device,
|
||||||
size: Vec<u32>,
|
size: Vec<u32>,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
@@ -52,8 +52,8 @@ 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 eos_token_id = cfg.eos_token_id;
|
||||||
let bos_token_id = cfg.bos_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 deepseekocr_model = DeepseekOCRModel::new(vb, cfg, version)?;
|
||||||
@@ -63,8 +63,8 @@ impl DeepseekOCRGenerateModel {
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
deepseekocr_model,
|
deepseekocr_model,
|
||||||
bos_token_id,
|
// bos_token_id,
|
||||||
eos_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(),
|
||||||
@@ -87,58 +87,37 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
|||||||
} else {
|
} else {
|
||||||
640
|
640
|
||||||
};
|
};
|
||||||
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let base_size = if self.version == 2 { 1024 } else { base_size };
|
let base_size = if self.version == 2 { 1024 } else { base_size };
|
||||||
let image_size = if self.version == 2 { 768 } else { image_size };
|
let image_size = if self.version == 2 { 768 } else { image_size };
|
||||||
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
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
|
.processor
|
||||||
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
||||||
let mut seqlen_offset = 0;
|
let max_tokens = mes.max_tokens.unwrap_or(1024);
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let mut ctx = GenerationContext::new(
|
||||||
let prompt_tokens = seq_len as u32;
|
mes.temperature,
|
||||||
let mut generate = Vec::new();
|
mes.top_p,
|
||||||
let logits = self.deepseekocr_model.forward(
|
|
||||||
&input_ids,
|
|
||||||
Some(&images_ori),
|
|
||||||
Some(&image_crop),
|
|
||||||
Some(&images_seq_mask),
|
|
||||||
Some(&images_spatial_crop_t),
|
|
||||||
seqlen_offset,
|
|
||||||
)?;
|
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
generate.push(next_token);
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
|
||||||
for _ in 1..sample_len {
|
|
||||||
let logits = self.deepseekocr_model.forward(
|
|
||||||
&input_ids,
|
|
||||||
None,
|
None,
|
||||||
None,
|
mes.seed.unwrap_or(34562) as u64,
|
||||||
None,
|
input_ids.dim(1)?,
|
||||||
None,
|
max_tokens,
|
||||||
seqlen_offset,
|
self.device.clone(),
|
||||||
)?;
|
);
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let data_vec = vec![
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
Some(images_ori),
|
||||||
generate.push(next_token);
|
Some(image_crop),
|
||||||
if next_token == self.bos_token_id || next_token == self.eos_token_id {
|
Some(images_seq_mask),
|
||||||
break;
|
Some(images_spatial_crop_t),
|
||||||
}
|
];
|
||||||
seqlen_offset += seq_len;
|
let data = MultiModalData::new(data_vec);
|
||||||
seq_len = 1;
|
generate_generic(
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
&mut self.deepseekocr_model,
|
||||||
}
|
&self.tokenizer,
|
||||||
let num_token = generate.len() as u32;
|
input_ids,
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
data,
|
||||||
self.deepseekocr_model.clear_kv_cache();
|
&mut ctx,
|
||||||
let response =
|
&self.model_name,
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
)
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -164,71 +143,37 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
|||||||
} else {
|
} else {
|
||||||
640
|
640
|
||||||
};
|
};
|
||||||
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
|
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let base_size = if self.version == 2 { 1024 } else { base_size };
|
let base_size = if self.version == 2 { 1024 } else { base_size };
|
||||||
let image_size = if self.version == 2 { 768 } else { image_size };
|
let image_size = if self.version == 2 { 768 } else { image_size };
|
||||||
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
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
|
.processor
|
||||||
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
.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 mut seqlen_offset = 0;
|
let temperature = mes.temperature;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let top_p = mes.top_p;
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let stream = stream! {
|
let max_tokens = mes.max_tokens.unwrap_or(1024);
|
||||||
let mut error_tokens = Vec::new();
|
let stream = generate_stream_generic(
|
||||||
let mut images_ori = Some(&images_ori);
|
&mut self.deepseekocr_model,
|
||||||
let mut image_crop = Some(&image_crop);
|
&self.tokenizer,
|
||||||
let mut images_seq_mask = Some(&images_seq_mask);
|
input_ids,
|
||||||
let mut images_spatial_crop_t = Some(&images_spatial_crop_t);
|
data,
|
||||||
for _ in 0..sample_len {
|
temperature,
|
||||||
let logits = self.deepseekocr_model.forward(
|
top_p,
|
||||||
&input_ids,
|
None,
|
||||||
images_ori,
|
seed,
|
||||||
image_crop,
|
max_tokens,
|
||||||
images_seq_mask,
|
&self.device,
|
||||||
images_spatial_crop_t,
|
&self.model_name,
|
||||||
seqlen_offset,
|
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
images_ori = None;
|
|
||||||
image_crop = None;
|
|
||||||
images_seq_mask = None;
|
|
||||||
images_spatial_crop_t = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.bos_token_id || next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
images_ori = None;
|
|
||||||
image_crop = None;
|
|
||||||
images_seq_mask = None;
|
|
||||||
images_spatial_crop_t = None;
|
|
||||||
}
|
|
||||||
self.deepseekocr_model.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -13,8 +13,11 @@ use candle_transformers::models::segment_anything::LayerNorm2d;
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::{
|
||||||
GateUpDownMLP, NaiveAttention, TwoLinearMLP, eager_attention_forward, get_conv2d,
|
InferenceModel,
|
||||||
get_layer_norm,
|
modules::{
|
||||||
|
GateUpDownMLP, NaiveAttention, QKVCatAttention, TwoLinearMLP,
|
||||||
|
eager_attention_forward, get_conv2d, get_layer_norm,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
deepseek_ocr::config::{DeepseekOCRConfig, DeepseekV2Config},
|
deepseek_ocr::config::{DeepseekOCRConfig, DeepseekV2Config},
|
||||||
qwen2::{Qwen2Config, Qwen2Decoder},
|
qwen2::{Qwen2Config, Qwen2Decoder},
|
||||||
@@ -608,45 +611,6 @@ impl CLIPVisionEmbeddings {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub struct NoTPAttention {
|
|
||||||
num_heads: usize,
|
|
||||||
head_dim: usize,
|
|
||||||
qkv_proj: Linear,
|
|
||||||
out_proj: Linear,
|
|
||||||
scaling: f64,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl NoTPAttention {
|
|
||||||
pub fn new(vb: VarBuilder, hidden_size: usize, num_heads: usize) -> Result<Self> {
|
|
||||||
let qkv_proj = linear(hidden_size, hidden_size * 3, vb.pp("qkv_proj"))?;
|
|
||||||
let out_proj = linear(hidden_size, hidden_size, vb.pp("out_proj"))?;
|
|
||||||
let head_dim = hidden_size / num_heads;
|
|
||||||
let scaling = 1.0 / (head_dim as f64).sqrt();
|
|
||||||
Ok(Self {
|
|
||||||
num_heads,
|
|
||||||
head_dim,
|
|
||||||
qkv_proj,
|
|
||||||
out_proj,
|
|
||||||
scaling,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
|
||||||
let (bs, seq_len, _) = xs.dims3()?;
|
|
||||||
let qkv = self.qkv_proj.forward(xs)?;
|
|
||||||
let qkv = qkv
|
|
||||||
.reshape((bs, seq_len, 3, self.num_heads, self.head_dim))?
|
|
||||||
.permute((2, 0, 3, 1, 4))?;
|
|
||||||
let q = qkv.i(0)?.contiguous()?;
|
|
||||||
let k = qkv.i(1)?.contiguous()?;
|
|
||||||
let v = qkv.i(2)?.contiguous()?;
|
|
||||||
let output = eager_attention_forward(&q, &k, &v, None, None, self.scaling)?;
|
|
||||||
let output = output.reshape((bs, seq_len, ()))?;
|
|
||||||
let output = self.out_proj.forward(&output)?;
|
|
||||||
Ok(output)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pub struct NoTPFeedForward {
|
pub struct NoTPFeedForward {
|
||||||
fc1: Linear,
|
fc1: Linear,
|
||||||
fc2: Linear,
|
fc2: Linear,
|
||||||
@@ -668,7 +632,7 @@ impl NoTPFeedForward {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct NoTPTransformerBlock {
|
pub struct NoTPTransformerBlock {
|
||||||
self_attn: NoTPAttention,
|
self_attn: QKVCatAttention,
|
||||||
mlp: NoTPFeedForward,
|
mlp: NoTPFeedForward,
|
||||||
layer_norm1: LayerNorm,
|
layer_norm1: LayerNorm,
|
||||||
layer_norm2: LayerNorm,
|
layer_norm2: LayerNorm,
|
||||||
@@ -681,7 +645,15 @@ impl NoTPTransformerBlock {
|
|||||||
ffn_hidden_size: usize,
|
ffn_hidden_size: usize,
|
||||||
eps: f64,
|
eps: f64,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
let self_attn = NoTPAttention::new(vb.pp("self_attn"), hidden_size, num_heads)?;
|
let self_attn = QKVCatAttention::new(
|
||||||
|
vb.pp("self_attn"),
|
||||||
|
hidden_size,
|
||||||
|
num_heads,
|
||||||
|
None,
|
||||||
|
true,
|
||||||
|
Some("qkv_proj"),
|
||||||
|
Some("out_proj"),
|
||||||
|
)?;
|
||||||
let mlp = NoTPFeedForward::new(vb.pp("mlp"), hidden_size, ffn_hidden_size)?;
|
let mlp = NoTPFeedForward::new(vb.pp("mlp"), hidden_size, ffn_hidden_size)?;
|
||||||
let layer_norm1 = get_layer_norm(vb.pp("layer_norm1"), eps, hidden_size, true)?;
|
let layer_norm1 = get_layer_norm(vb.pp("layer_norm1"), eps, hidden_size, true)?;
|
||||||
let layer_norm2 = get_layer_norm(vb.pp("layer_norm2"), eps, hidden_size, true)?;
|
let layer_norm2 = get_layer_norm(vb.pp("layer_norm2"), eps, hidden_size, true)?;
|
||||||
@@ -695,7 +667,7 @@ impl NoTPTransformerBlock {
|
|||||||
|
|
||||||
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
||||||
let x = self.layer_norm1.forward(xs)?;
|
let x = self.layer_norm1.forward(xs)?;
|
||||||
let x = self.self_attn.forward(&x)?;
|
let x = self.self_attn.forward(&x, None, None, None, false, false)?;
|
||||||
let res = x.add(xs)?;
|
let res = x.add(xs)?;
|
||||||
let x = self.layer_norm2.forward(&res)?;
|
let x = self.layer_norm2.forward(&res)?;
|
||||||
let x = self.mlp.forward(&x)?;
|
let x = self.mlp.forward(&x)?;
|
||||||
@@ -1204,6 +1176,7 @@ pub struct DeepseekOCRModel {
|
|||||||
image_newline: Option<Tensor>,
|
image_newline: Option<Tensor>,
|
||||||
view_seperator: Tensor,
|
view_seperator: Tensor,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl DeepseekOCRModel {
|
impl DeepseekOCRModel {
|
||||||
@@ -1262,6 +1235,7 @@ impl DeepseekOCRModel {
|
|||||||
let view_seperator = vb_m.get_with_hints(1280, "view_seperator", Init::Const(0.))?;
|
let view_seperator = vb_m.get_with_hints(1280, "view_seperator", Init::Const(0.))?;
|
||||||
let language_model = DeepseekV2Model::new(vb_m, config.language_config.clone())?;
|
let language_model = DeepseekV2Model::new(vb_m, config.language_config.clone())?;
|
||||||
let lm_head = linear_no_bias(config.hidden_size, config.vocab_size, vb.pp("lm_head"))?;
|
let lm_head = linear_no_bias(config.hidden_size, config.vocab_size, vb.pp("lm_head"))?;
|
||||||
|
let stop_token_ids = vec![config.eos_token_id, config.bos_token_id];
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
// config,
|
// config,
|
||||||
sam_model,
|
sam_model,
|
||||||
@@ -1271,6 +1245,7 @@ impl DeepseekOCRModel {
|
|||||||
image_newline,
|
image_newline,
|
||||||
view_seperator,
|
view_seperator,
|
||||||
lm_head,
|
lm_head,
|
||||||
|
stop_token_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1459,3 +1434,42 @@ impl DeepseekOCRModel {
|
|||||||
self.language_model.clear_kv_cache();
|
self.language_model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for DeepseekOCRModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 4 {
|
||||||
|
return Err(anyhow!(
|
||||||
|
"DeepseekOCR process data error, must have images_ori, image_crop, images_seq_mask, images_spatial_crop"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let images_ori = &data.data_vec[0];
|
||||||
|
let image_crop = &data.data_vec[1];
|
||||||
|
let images_seq_mask = &data.data_vec[2];
|
||||||
|
let images_spatial_crop = &data.data_vec[3];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
images_ori.as_ref(),
|
||||||
|
image_crop.as_ref(),
|
||||||
|
images_seq_mask.as_ref(),
|
||||||
|
images_spatial_crop.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,12 +1,15 @@
|
|||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
|
|
||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
|
use candle_core::{DType, Device, pickle::read_all_with_key};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -18,10 +21,7 @@ use crate::{
|
|||||||
qwen3::config::{Qwen3Config, Qwen3GenerationConfig},
|
qwen3::config::{Qwen3Config, Qwen3GenerationConfig},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype, get_logit_processor,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct FunAsrNanoGenerateModel {
|
pub struct FunAsrNanoGenerateModel {
|
||||||
@@ -30,8 +30,8 @@ pub struct FunAsrNanoGenerateModel {
|
|||||||
fun_asr_nano: FunAsrNanoModel,
|
fun_asr_nano: FunAsrNanoModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
eos_token_id1: u32,
|
// eos_token_id1: u32,
|
||||||
eos_token_id2: u32,
|
// eos_token_id2: u32,
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
@@ -77,7 +77,8 @@ 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 = FunAsrNanoModel::new(vb, &cfg, &llm_cfg)?;
|
let fun_asr_nano =
|
||||||
|
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()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -89,8 +90,8 @@ impl FunAsrNanoGenerateModel {
|
|||||||
fun_asr_nano,
|
fun_asr_nano,
|
||||||
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,
|
||||||
eos_token_id2: generation_config.eos_token_id[1] as u32,
|
// eos_token_id2: generation_config.eos_token_id[1] as u32,
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
@@ -105,42 +106,29 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
let max_tokens = mes.max_tokens.unwrap_or(1024);
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?;
|
||||||
let (speech, fbank_mask, mut input_ids) =
|
let speech = speech.to_dtype(self.dtype)?;
|
||||||
self.processor.process_info(&mes, &self.tokenizer)?;
|
let mut ctx = GenerationContext::new(
|
||||||
let mut speech = Some(speech.to_dtype(self.dtype)?);
|
temperature.into(),
|
||||||
let mut fbank_mask = Some(&fbank_mask);
|
top_p.into(),
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
top_k.into(),
|
||||||
let prompt_tokens = seq_len as u32;
|
seed,
|
||||||
let mut seqlen_offset = 0;
|
input_ids.dim(1)?,
|
||||||
let mut generate = Vec::new();
|
max_tokens,
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
self.device.clone(),
|
||||||
for _ in 0..sample_len {
|
);
|
||||||
let logits = self.fun_asr_nano.forward(
|
|
||||||
&input_ids,
|
let data_vec = vec![speech.into(), fbank_mask.into()];
|
||||||
speech.as_ref(),
|
let data = MultiModalData::new(data_vec);
|
||||||
fbank_mask,
|
generate_generic(
|
||||||
seqlen_offset,
|
&mut self.fun_asr_nano,
|
||||||
)?;
|
&self.tokenizer,
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
input_ids,
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
data,
|
||||||
generate.push(next_token);
|
&mut ctx,
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
&self.model_name,
|
||||||
break;
|
)
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
speech = None;
|
|
||||||
fbank_mask = None;
|
|
||||||
}
|
|
||||||
let num_token = generate.len() as u32;
|
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
|
||||||
self.fun_asr_nano.clear_kv_cache();
|
|
||||||
let response =
|
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -160,58 +148,75 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
let max_tokens = mes.max_tokens.unwrap_or(1024);
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
// let mut logit_processor =
|
||||||
|
// get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
||||||
let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?;
|
let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let speech = speech.to_dtype(self.dtype)?;
|
||||||
let mut seqlen_offset = 0;
|
let data_vec = vec![speech.into(), fbank_mask.into()];
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let data = MultiModalData::new(data_vec);
|
||||||
let stream = stream! {
|
let stream = generate_stream_generic(
|
||||||
let mut error_tokens = Vec::new();
|
&mut self.fun_asr_nano,
|
||||||
let mut speech = Some(speech.to_dtype(self.dtype)?);
|
&self.tokenizer,
|
||||||
let mut fbank_mask = Some(&fbank_mask);
|
input_ids,
|
||||||
let mut input_ids = input_ids;
|
data,
|
||||||
for _ in 0..sample_len {
|
temperature.into(),
|
||||||
let logits = self.fun_asr_nano.forward(
|
top_p.into(),
|
||||||
&input_ids,
|
top_k.into(),
|
||||||
speech.as_ref(),
|
seed,
|
||||||
fbank_mask,
|
max_tokens,
|
||||||
seqlen_offset,
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
// let mut seq_len = input_ids.dim(1)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
// let mut seqlen_offset = 0;
|
||||||
let mut decode_ids = Vec::new();
|
// let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
if !error_tokens.is_empty() {
|
// let stream = stream! {
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
// let mut error_tokens = Vec::new();
|
||||||
}
|
// let mut speech = Some(speech.to_dtype(self.dtype)?);
|
||||||
decode_ids.push(next_token);
|
// let mut fbank_mask = Some(&fbank_mask);
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
// let mut input_ids = input_ids;
|
||||||
if decoded_token.contains("�") {
|
// for _ in 0..sample_len {
|
||||||
error_tokens.push(next_token);
|
// let logits = self.fun_asr_nano.forward(
|
||||||
if error_tokens.len() > 3 {
|
// &input_ids,
|
||||||
error_tokens.clear();
|
// speech.as_ref(),
|
||||||
}
|
// fbank_mask,
|
||||||
seqlen_offset += seq_len;
|
// seqlen_offset,
|
||||||
seq_len = 1;
|
// )?;
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
// let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
speech = None;
|
// let next_token = logit_processor.sample(&logits)?;
|
||||||
fbank_mask = None;
|
// let mut decode_ids = Vec::new();
|
||||||
continue;
|
// if !error_tokens.is_empty() {
|
||||||
}
|
// decode_ids.extend_from_slice(&error_tokens);
|
||||||
error_tokens.clear();
|
// }
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
// decode_ids.push(next_token);
|
||||||
yield Ok(chunk);
|
// let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
// if decoded_token.contains("�") {
|
||||||
break;
|
// error_tokens.push(next_token);
|
||||||
}
|
// if error_tokens.len() > 3 {
|
||||||
seqlen_offset += seq_len;
|
// error_tokens.clear();
|
||||||
seq_len = 1;
|
// }
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
// seqlen_offset += seq_len;
|
||||||
speech = None;
|
// seq_len = 1;
|
||||||
fbank_mask = None;
|
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
}
|
// speech = None;
|
||||||
self.fun_asr_nano.clear_kv_cache();
|
// fbank_mask = None;
|
||||||
};
|
// continue;
|
||||||
|
// }
|
||||||
|
// error_tokens.clear();
|
||||||
|
// let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
||||||
|
// yield Ok(chunk);
|
||||||
|
// if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
||||||
|
// break;
|
||||||
|
// }
|
||||||
|
// seqlen_offset += seq_len;
|
||||||
|
// seq_len = 1;
|
||||||
|
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
|
// speech = None;
|
||||||
|
// fbank_mask = None;
|
||||||
|
// }
|
||||||
|
// self.fun_asr_nano.clear_kv_cache();
|
||||||
|
// };
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,12 +1,15 @@
|
|||||||
use anyhow::Result;
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{D, IndexOp, Tensor};
|
use candle_core::{D, IndexOp, Tensor};
|
||||||
use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, ops::softmax_last_dim};
|
use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, ops::softmax_last_dim};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::{
|
||||||
NaiveAttention, TwoLinearMLP, conv1d_depthwise, eager_attention_forward, get_conv1d,
|
InferenceModel,
|
||||||
get_layer_norm,
|
modules::{
|
||||||
|
NaiveAttention, TwoLinearMLP, conv1d_depthwise, eager_attention_forward,
|
||||||
|
get_conv1d, get_layer_norm,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
fun_asr_nano::config::FunASRNanoConfig,
|
fun_asr_nano::config::FunASRNanoConfig,
|
||||||
qwen3::{config::Qwen3Config, model::Qwen3Model},
|
qwen3::{config::Qwen3Config, model::Qwen3Model},
|
||||||
@@ -577,9 +580,15 @@ pub struct FunAsrNanoModel {
|
|||||||
audio_encoder: SenseVoiceEncoderSmall,
|
audio_encoder: SenseVoiceEncoderSmall,
|
||||||
audio_adaptor: AudioAdaptor,
|
audio_adaptor: AudioAdaptor,
|
||||||
llm: Qwen3Model,
|
llm: Qwen3Model,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
impl FunAsrNanoModel {
|
impl FunAsrNanoModel {
|
||||||
pub fn new(vb: VarBuilder, config: &FunASRNanoConfig, llm_cfg: &Qwen3Config) -> Result<Self> {
|
pub fn new(
|
||||||
|
vb: VarBuilder,
|
||||||
|
config: &FunASRNanoConfig,
|
||||||
|
llm_cfg: &Qwen3Config,
|
||||||
|
eos_ids: Vec<u32>,
|
||||||
|
) -> Result<Self> {
|
||||||
let input_size = config.frontend_conf.lfr_m * config.frontend_conf.n_mels;
|
let input_size = config.frontend_conf.lfr_m * config.frontend_conf.n_mels;
|
||||||
let audio_encoder = SenseVoiceEncoderSmall::new(
|
let audio_encoder = SenseVoiceEncoderSmall::new(
|
||||||
vb.pp("audio_encoder"),
|
vb.pp("audio_encoder"),
|
||||||
@@ -607,6 +616,7 @@ impl FunAsrNanoModel {
|
|||||||
audio_encoder,
|
audio_encoder,
|
||||||
audio_adaptor,
|
audio_adaptor,
|
||||||
llm,
|
llm,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -639,3 +649,38 @@ impl FunAsrNanoModel {
|
|||||||
self.llm.clear_kv_cache();
|
self.llm.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for FunAsrNanoModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 2 {
|
||||||
|
return Err(anyhow!(
|
||||||
|
"FunAsrNano process data error, must have speech, fbank_mask"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let speech = &data.data_vec[0];
|
||||||
|
let fbank_mask = &data.data_vec[1];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
speech.as_ref(),
|
||||||
|
fbank_mask.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
@@ -18,7 +19,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, linear_no
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::modules::{
|
||||||
LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
|
LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
|
||||||
},
|
},
|
||||||
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
|
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
//! GLM-OCR Inference and Generation
|
//! GLM-OCR Inference and Generation
|
||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, IndexOp, Tensor};
|
use candle_core::{DType, Device, IndexOp, Tensor};
|
||||||
@@ -21,7 +22,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, extract_user_text,
|
build_completion_chunk_response, build_completion_response, extract_user_text,
|
||||||
find_type_files, get_device, get_dtype, get_logit_processor, img_utils::extract_image_url,
|
find_type_files, get_device, get_dtype, img_utils::extract_image_url,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::GateUpDownMLP,
|
common::modules::GateUpDownMLP,
|
||||||
glm_ocr::config::{GlmOcrConfig, GlmOcrTextConfig, GlmOcrVisionConfig},
|
glm_ocr::config::{GlmOcrConfig, GlmOcrTextConfig, GlmOcrVisionConfig},
|
||||||
},
|
},
|
||||||
position_embed::rope::{apply_rotary_pos_emb_vision, glm_ocr_apply_rotary_pos_emb},
|
position_embed::rope::{apply_rotary_pos_emb_vision, glm_ocr_apply_rotary_pos_emb},
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
@@ -20,7 +21,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,9 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d},
|
common::modules::{
|
||||||
|
GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d,
|
||||||
|
},
|
||||||
hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig},
|
hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig},
|
||||||
},
|
},
|
||||||
position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin},
|
position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin},
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use crate::models::common::generate::get_logit_processor;
|
||||||
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
|
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
|
||||||
use crate::utils::build_completion_chunk_response;
|
use crate::utils::build_completion_chunk_response;
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -10,9 +11,7 @@ use crate::{
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{build_completion_response, find_type_files, get_device, get_dtype},
|
||||||
build_completion_response, find_type_files, get_device, get_dtype, get_logit_processor,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d},
|
common::modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d},
|
||||||
lfm2::config::Lfm2Config,
|
lfm2::config::Lfm2Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::RoPE,
|
position_embed::rope::RoPE,
|
||||||
|
|||||||
@@ -1,4 +1,7 @@
|
|||||||
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
|
use crate::{
|
||||||
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionParameters, ChatCompletionResponse},
|
||||||
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
@@ -13,7 +16,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
use rocket::async_stream::stream;
|
use rocket::async_stream::stream;
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{NaiveAttnTwoLinearMLPBlock, get_layer_norm},
|
common::modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm},
|
||||||
lfm2::model::Lfm2Decoder,
|
lfm2::model::Lfm2Decoder,
|
||||||
lfm2vl::config::{Lfm2VLConfig, Lfm2VLVisionConfig},
|
lfm2vl::config::{Lfm2VLConfig, Lfm2VLVisionConfig},
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{WNConv1d, conv1d_depthwise, get_conv1d, get_layer_norm},
|
common::modules::{WNConv1d, conv1d_depthwise, get_conv1d, get_layer_norm},
|
||||||
mask_gct::config::SemanticCodec,
|
mask_gct::config::SemanticCodec,
|
||||||
},
|
},
|
||||||
utils::{interpolate::interpolate_nearest_1d, tensor_utils::l2_normalize},
|
utils::{interpolate::interpolate_nearest_1d, tensor_utils::l2_normalize},
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use crate::models::common::generate::get_logit_processor;
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
@@ -12,7 +13,7 @@ use crate::models::minicpm4::model::MiniCPMModel;
|
|||||||
// use crate::models::GenerateStream;
|
// use crate::models::GenerateStream;
|
||||||
use crate::utils::{
|
use crate::utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
};
|
};
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_n
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{GateUpDownMLP, NaiveAttention},
|
common::modules::{GateUpDownMLP, NaiveAttention},
|
||||||
minicpm4::config::MiniCPM4Config,
|
minicpm4::config::MiniCPM4Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::compute_default_rope_parameters,
|
position_embed::rope::compute_default_rope_parameters,
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use crate::models::common::generate::get_logit_processor;
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
@@ -13,7 +14,7 @@ use crate::models::paddleocr_vl::processor::PaddleOCRVLProcessor;
|
|||||||
use crate::utils::tensor_utils::get_equal_mask;
|
use crate::utils::tensor_utils::get_equal_mask;
|
||||||
use crate::utils::{
|
use crate::utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
};
|
};
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ use num::integer::Roots;
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::modules::{
|
||||||
NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm,
|
NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm,
|
||||||
},
|
},
|
||||||
paddleocr_vl::config::{
|
paddleocr_vl::config::{
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ use candle_nn::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::common::{GateUpDownMLP, eager_attention_forward},
|
models::common::modules::{GateUpDownMLP, eager_attention_forward},
|
||||||
position_embed::rope::{RoPE, apply_rotary_pos_emb},
|
position_embed::rope::{RoPE, apply_rotary_pos_emb},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use crate::models::common::generate::get_logit_processor;
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
@@ -10,7 +11,7 @@ use rocket::futures::Stream;
|
|||||||
use crate::models::qwen2_5vl::config::Qwen2_5VLConfig;
|
use crate::models::qwen2_5vl::config::Qwen2_5VLConfig;
|
||||||
use crate::utils::{
|
use crate::utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
chat_template::ChatTemplate,
|
chat_template::ChatTemplate,
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ use candle_nn::{Init, Linear, Module, RmsNorm, VarBuilder, linear, linear_no_bia
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{GateUpDownMLP, eager_attention_forward},
|
common::modules::{GateUpDownMLP, eager_attention_forward},
|
||||||
qwen2::Qwen2DecoderLayer,
|
qwen2::Qwen2DecoderLayer,
|
||||||
qwen2_5vl::config::{Qwen2_5VLConfig, RopeScaling},
|
qwen2_5vl::config::{Qwen2_5VLConfig, RopeScaling},
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ pub struct Qwen3GenerationConfig {
|
|||||||
pub bos_token_id: usize,
|
pub bos_token_id: usize,
|
||||||
pub pad_token_id: usize,
|
pub pad_token_id: usize,
|
||||||
pub do_sample: bool,
|
pub do_sample: bool,
|
||||||
pub eos_token_id: Vec<usize>,
|
pub eos_token_id: Vec<u32>,
|
||||||
pub top_p: f32,
|
pub top_p: f32,
|
||||||
pub top_k: usize,
|
pub top_k: usize,
|
||||||
pub temperature: f32,
|
pub temperature: f32,
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use crate::models::common::generate::get_logit_processor;
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
@@ -12,7 +13,7 @@ use crate::models::qwen3::model::Qwen3Model;
|
|||||||
// use crate::models::GenerateStream;
|
// use crate::models::GenerateStream;
|
||||||
use crate::utils::{
|
use crate::utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
};
|
};
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{GateUpDownMLP, QKNormAttention},
|
common::modules::{GateUpDownMLP, QKNormAttention},
|
||||||
qwen3::config::Qwen3Config,
|
qwen3::config::Qwen3Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::RoPE,
|
position_embed::rope::RoPE,
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor, quantized::gguf_file};
|
use candle_core::{DType, Device, Tensor, quantized::gguf_file};
|
||||||
@@ -18,7 +19,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -10,9 +10,8 @@ use candle_nn::{
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::{
|
||||||
conv1d_depthwise, eager_attention_forward, get_conv1d,
|
|
||||||
gguf::{GateUpDownMLPGguf, Gguf, ProjKind, QuantizedLinear},
|
gguf::{GateUpDownMLPGguf, Gguf, ProjKind, QuantizedLinear},
|
||||||
softplus,
|
modules::{conv1d_depthwise, eager_attention_forward, get_conv1d, softplus},
|
||||||
},
|
},
|
||||||
qwen3_5::config::{Qwen3_5Config, Qwen3_5TextConfig},
|
qwen3_5::config::{Qwen3_5Config, Qwen3_5TextConfig},
|
||||||
qwen3vl::model::Qwen3VLVisionModel,
|
qwen3vl::model::Qwen3VLVisionModel,
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
@@ -21,7 +22,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{NaiveAttention, get_conv2d, get_layer_norm},
|
common::modules::{NaiveAttention, get_conv2d, get_layer_norm},
|
||||||
qwen3::model::Qwen3DecoderLayer,
|
qwen3::model::Qwen3DecoderLayer,
|
||||||
qwen3_asr::{
|
qwen3_asr::{
|
||||||
config::{
|
config::{
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
use crate::params::chat::{
|
use crate::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
models::common::generate::get_logit_processor,
|
||||||
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
@@ -17,7 +18,7 @@ use crate::{
|
|||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
get_dtype, get_logit_processor,
|
get_dtype,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -10,8 +10,8 @@ use candle_nn::{
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::{
|
||||||
eager_attention_forward, get_layer_norm,
|
|
||||||
gguf::{Gguf, ProjKind, TwoLinearMLPGguf},
|
gguf::{Gguf, ProjKind, TwoLinearMLPGguf},
|
||||||
|
modules::{eager_attention_forward, get_layer_norm},
|
||||||
},
|
},
|
||||||
qwen3::model::Qwen3DecoderLayer,
|
qwen3::model::Qwen3DecoderLayer,
|
||||||
qwen3vl::config::{
|
qwen3vl::config::{
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ use candle_nn::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::common::{
|
models::common::modules::{
|
||||||
Conv2dWithBN, TwoLinearMLP, deform_conv2d_kernel, get_batch_norm, get_conv2d,
|
Conv2dWithBN, TwoLinearMLP, deform_conv2d_kernel, get_batch_norm, get_conv2d,
|
||||||
get_layer_norm,
|
get_layer_norm,
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ use candle_nn::{Embedding, Module, RmsNorm, VarBuilder, embedding, rms_norm};
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{GateUpDownMLP, NaiveAttention},
|
common::modules::{GateUpDownMLP, NaiveAttention},
|
||||||
voxcpm::config::VoxMiniCPM4Config,
|
voxcpm::config::VoxMiniCPM4Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::compute_default_rope_parameters,
|
position_embed::rope::compute_default_rope_parameters,
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::modules::{
|
||||||
GLU, TwoLinearMLP, conv1d_depthwise, eager_attention_forward, get_conv1d,
|
GLU, TwoLinearMLP, conv1d_depthwise, eager_attention_forward, get_conv1d,
|
||||||
get_layer_norm,
|
get_layer_norm,
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -138,6 +138,8 @@ pub struct ChatCompletionParameters {
|
|||||||
/// So 0.1 means only the tokens comprising the top 10% probability mass are considered.
|
/// So 0.1 means only the tokens comprising the top 10% probability mass are considered.
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub top_p: Option<f32>,
|
pub top_p: Option<f32>,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
pub top_k: Option<usize>,
|
||||||
/// A list of tools the model may call. Currently, only functions are supported as a tool.
|
/// A list of tools the model may call. Currently, only functions are supported as a tool.
|
||||||
/// Use this to provide a list of functions the model may generate JSON inputs for.
|
/// Use this to provide a list of functions the model may generate JSON inputs for.
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
|||||||
@@ -7,14 +7,14 @@ pub struct Usage {
|
|||||||
pub prompt_tokens: Option<u32>,
|
pub prompt_tokens: Option<u32>,
|
||||||
/// Number of tokens in the prompt.
|
/// Number of tokens in the prompt.
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub prompt_ms: Option<f64>,
|
pub prompt_secs: Option<f64>,
|
||||||
/// Number of tokens in the completion.
|
/// Number of tokens in the completion.
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub completion_tokens: Option<u32>,
|
pub completion_tokens: Option<u32>,
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub completion_ms: Option<f64>,
|
pub completion_secs: Option<f64>,
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub completion_per_token_ms: Option<f64>,
|
pub completion_per_token_secs: Option<f64>,
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub completion_tps: Option<f64>,
|
pub completion_tps: Option<f64>,
|
||||||
/// Number of tokens in the entire response.
|
/// Number of tokens in the entire response.
|
||||||
|
|||||||
+14
-47
@@ -26,7 +26,6 @@ use candle_core::{
|
|||||||
pickle::{Object, Stack, TensorInfo, read_all_with_key},
|
pickle::{Object, Stack, TensorInfo, read_all_with_key},
|
||||||
};
|
};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use candle_transformers::generation::{LogitsProcessor, Sampling};
|
|
||||||
use dirs::home_dir;
|
use dirs::home_dir;
|
||||||
use half::{bf16, f16, slice::HalfFloatSliceExt};
|
use half::{bf16, f16, slice::HalfFloatSliceExt};
|
||||||
use modelscope::ModelScope;
|
use modelscope::ModelScope;
|
||||||
@@ -583,10 +582,10 @@ pub fn build_completion_response(
|
|||||||
} else {
|
} else {
|
||||||
Some(Usage {
|
Some(Usage {
|
||||||
prompt_tokens,
|
prompt_tokens,
|
||||||
prompt_ms: None,
|
prompt_secs: None,
|
||||||
completion_tokens,
|
completion_tokens,
|
||||||
completion_ms: None,
|
completion_secs: None,
|
||||||
completion_per_token_ms: None,
|
completion_per_token_secs: None,
|
||||||
completion_tps: None,
|
completion_tps: None,
|
||||||
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
||||||
prompt_tokens_details: None,
|
prompt_tokens_details: None,
|
||||||
@@ -601,28 +600,29 @@ pub fn build_completion_response_with_time(
|
|||||||
res: String,
|
res: String,
|
||||||
model_name: &str,
|
model_name: &str,
|
||||||
completion_tokens: Option<u32>,
|
completion_tokens: Option<u32>,
|
||||||
completion_ms: Option<f64>,
|
completion_secs: Option<f64>,
|
||||||
prompt_tokens: Option<u32>,
|
prompt_tokens: Option<u32>,
|
||||||
prompt_ms: Option<f64>,
|
prompt_secs: Option<f64>,
|
||||||
) -> ChatCompletionResponse {
|
) -> ChatCompletionResponse {
|
||||||
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
||||||
None
|
None
|
||||||
} else {
|
} else {
|
||||||
let (completion_per_token_ms, completion_tps) = if let Some(prompt_tokens) = prompt_tokens
|
let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) =
|
||||||
&& let Some(prompt_ms) = prompt_ms
|
completion_tokens
|
||||||
|
&& let Some(completion_secs) = completion_secs
|
||||||
{
|
{
|
||||||
let per_token_ms = prompt_ms / prompt_tokens as f64;
|
let per_token_secs = completion_secs / completion_tokens as f64;
|
||||||
let tps = prompt_tokens as f64 / (prompt_ms / 1000.0);
|
let tps = completion_tokens as f64 / completion_secs;
|
||||||
(Some(per_token_ms), Some(tps))
|
(Some(per_token_secs), Some(tps))
|
||||||
} else {
|
} else {
|
||||||
(None, None)
|
(None, None)
|
||||||
};
|
};
|
||||||
Some(Usage {
|
Some(Usage {
|
||||||
prompt_tokens,
|
prompt_tokens,
|
||||||
prompt_ms,
|
prompt_secs,
|
||||||
completion_tokens,
|
completion_tokens,
|
||||||
completion_ms,
|
completion_secs,
|
||||||
completion_per_token_ms,
|
completion_per_token_secs,
|
||||||
completion_tps,
|
completion_tps,
|
||||||
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
||||||
prompt_tokens_details: None,
|
prompt_tokens_details: None,
|
||||||
@@ -708,39 +708,6 @@ pub fn build_completion_chunk_response(
|
|||||||
response
|
response
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get_logit_processor(
|
|
||||||
temperature: Option<f32>,
|
|
||||||
top_p: Option<f32>,
|
|
||||||
top_k: Option<usize>,
|
|
||||||
seed: u64,
|
|
||||||
) -> LogitsProcessor {
|
|
||||||
let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) });
|
|
||||||
match top_k {
|
|
||||||
None => LogitsProcessor::new(
|
|
||||||
seed,
|
|
||||||
temperature.map(|temp| temp as f64),
|
|
||||||
top_p.map(|tp| tp as f64),
|
|
||||||
),
|
|
||||||
Some(k) => {
|
|
||||||
let sampling = match temperature {
|
|
||||||
None => Sampling::ArgMax,
|
|
||||||
Some(temperature) => match top_p {
|
|
||||||
None => Sampling::TopK {
|
|
||||||
k,
|
|
||||||
temperature: temperature as f64,
|
|
||||||
},
|
|
||||||
Some(p) => Sampling::TopKThenTopP {
|
|
||||||
k,
|
|
||||||
p: p as f64,
|
|
||||||
temperature: temperature as f64,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
};
|
|
||||||
LogitsProcessor::from_sampling(seed, sampling)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn extract_mes(mes: &ChatCompletionParameters) -> Result<Vec<(String, String)>> {
|
pub fn extract_mes(mes: &ChatCompletionParameters) -> Result<Vec<(String, String)>> {
|
||||||
let mut mes_vec = Vec::new();
|
let mut mes_vec = Vec::new();
|
||||||
for chat_mes in mes.messages.clone() {
|
for chat_mes in mes.messages.clone() {
|
||||||
|
|||||||
@@ -127,10 +127,6 @@ async fn deepseek_ocr_stream() -> Result<()> {
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
|
||||||
"role": "assistant",
|
|
||||||
"content": ""
|
|
||||||
}
|
|
||||||
],
|
],
|
||||||
"metadata": {"base_size": "640", "image_size": "640", "crop_mode": "false"}
|
"metadata": {"base_size": "640", "image_size": "640", "crop_mode": "false"}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -54,7 +54,7 @@ fn fun_asr_nano_generate() -> Result<()> {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn fun_asr_nano_stream() -> Result<()> {
|
async fn fun_asr_nano_stream() -> Result<()> {
|
||||||
// RUST_BACKTRACE=1 cargo test -F cuda fun_asr_nano_stream -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda --test test_fun_asr_nano fun_asr_nano_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"))?;
|
||||||
let model_path = format!("{}/FunAudioLLM/Fun-ASR-Nano-2512/", save_dir);
|
let model_path = format!("{}/FunAudioLLM/Fun-ASR-Nano-2512/", save_dir);
|
||||||
|
|||||||
Reference in New Issue
Block a user