add qwen3 and fun-asr-nano

This commit is contained in:
jhqxxx
2026-01-15 21:57:12 +08:00
parent 2c1d5e3a14
commit d9b803d27e
39 changed files with 2577 additions and 307 deletions
+87 -7
View File
@@ -1,4 +1,4 @@
use anyhow::Result;
use anyhow::{Result, anyhow};
use candle_core::{D, Tensor};
use candle_nn::{
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, Embedding,
@@ -6,6 +6,7 @@ use candle_nn::{
conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear, linear_no_bias,
rms_norm,
};
use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator};
use crate::{
position_embed::rope::{RoPE, apply_rotary_pos_emb},
@@ -126,6 +127,9 @@ impl NaiveAttention {
num_key_value_heads: usize,
head_dim: Option<usize>,
bias: bool,
q_proj_pp_name: Option<&str>,
k_proj_pp_name: Option<&str>,
v_proj_pp_name: Option<&str>,
o_proj_pp_name: Option<&str>,
) -> Result<Self> {
let num_kv_groups = num_attention_heads / num_key_value_heads;
@@ -133,12 +137,27 @@ impl NaiveAttention {
None => hidden_size / num_attention_heads,
Some(dim) => dim,
};
let q_proj_pp_name = q_proj_pp_name.unwrap_or("q_proj");
let k_proj_pp_name = k_proj_pp_name.unwrap_or("k_proj");
let v_proj_pp_name = v_proj_pp_name.unwrap_or("v_proj");
let o_proj_pp_name = o_proj_pp_name.unwrap_or("o_proj");
let (q_proj, k_proj, v_proj, o_proj) = if bias {
(
linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?,
linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?,
linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?,
linear(
hidden_size,
num_attention_heads * head_dim,
vb.pp(q_proj_pp_name),
)?,
linear(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(k_proj_pp_name),
)?,
linear(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(v_proj_pp_name),
)?,
linear(
num_attention_heads * head_dim,
hidden_size,
@@ -147,9 +166,21 @@ impl NaiveAttention {
)
} else {
(
linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?,
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?,
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?,
linear_no_bias(
hidden_size,
num_attention_heads * head_dim,
vb.pp(q_proj_pp_name),
)?,
linear_no_bias(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(k_proj_pp_name),
)?,
linear_no_bias(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(v_proj_pp_name),
)?,
linear_no_bias(
num_attention_heads * head_dim,
hidden_size,
@@ -305,6 +336,9 @@ impl NaiveAttnTwoLinearMLPBlock {
num_key_value_heads,
head_dim,
attn_bias,
None,
None,
None,
o_proj_pp_name,
)?;
let mlp = TwoLinearMLP::new(
@@ -386,6 +420,9 @@ impl NaiveAttnGateUpDownMLPBlock {
num_key_value_heads,
head_dim,
attn_bias,
None,
None,
None,
o_proj_pp_name,
)?;
let mlp = GateUpDownMLP::new(
@@ -811,3 +848,46 @@ impl LlamaForCausalLM {
self.model.clear_kv_cache();
}
}
pub fn conv1d_group_parallel(xs: &Tensor, conv1d: &Conv1d) -> Result<Tensor> {
let groups = conv1d.config().groups;
let xs = if groups == 1 {
xs.conv1d_with_algo(
conv1d.weight(),
conv1d.config().padding,
conv1d.config().stride,
conv1d.config().dilation,
groups,
conv1d.config().cudnn_fwd_algo,
)?
} else {
let blocks = xs.chunk(groups, 1)?;
let kernel = conv1d.weight().chunk(groups, 0)?;
let blocks = blocks
// .iter()
.par_iter()
.zip(&kernel)
.map(|(block, kernel)| {
block
.conv1d_with_algo(
kernel,
conv1d.config().padding,
conv1d.config().stride,
conv1d.config().dilation,
1,
conv1d.config().cudnn_fwd_algo,
)
.map_err(|e| anyhow!(format!("tensor conv1d_with_algo error:{}", e)))
})
.collect::<Result<Vec<Tensor>>>()?;
Tensor::cat(&blocks, 1)?
};
match conv1d.bias() {
None => Ok(xs),
Some(bias) => {
let b = bias.dims1()?;
let bias = bias.reshape((1, b, 1))?;
Ok(xs.broadcast_add(&bias)?)
}
}
}
+3
View File
@@ -986,6 +986,9 @@ impl DeepseekV2DecoderLayer {
None,
false,
None,
None,
None,
None,
)?;
let mlp = if layer_id >= config.first_k_dense_replace
&& layer_id.is_multiple_of(config.moe_layer_freq)
+84
View File
@@ -0,0 +1,84 @@
use serde::Deserialize;
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct FunASRNanoConfig {
pub audio_encoder_conf: AudioEncoderConf,
pub llm_conf: LlmConf,
pub audio_adaptor_conf: AudioAdaptorConf,
pub detach_ctc_decoder: bool,
pub ctc_decoder_conf: CtcDecoderConf,
pub ctc_weight: f64,
pub ctc_conf: CtcConf,
pub frontend_conf: FrontendConf,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct AudioEncoderConf {
pub output_size: usize,
pub attention_heads: usize,
pub linear_units: usize,
pub num_blocks: usize,
pub tp_blocks: usize,
pub dropout_rate: f64,
pub positional_dropout_rate: f64,
pub attention_dropout_rate: f64,
pub input_layer: String,
pub pos_enc_class: String,
pub normalize_before: bool,
pub kernel_size: usize,
pub sanm_shfit: usize,
pub selfattention_layer_type: String,
pub freeze: bool,
pub freeze_layer_num: i32,
pub feat_permute: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct LlmConf {
pub hub: String,
pub freeze: bool,
pub llm_dtype: String,
pub init_param_path: String,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct AudioAdaptorConf {
pub downsample_rate: usize,
pub use_low_frame_rate: bool,
pub ffn_dim: usize,
pub llm_dim: usize,
pub encoder_dim: usize,
pub n_layer: usize,
pub freeze: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct CtcDecoderConf {
pub downsample_rate: u32,
pub ffn_dim: u32,
pub llm_dim: u32,
pub encoder_dim: u32,
pub n_layer: u32,
pub freeze: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct CtcConf {
pub dropout_rate: f64,
pub ctc_type: String,
pub reduce: bool,
pub ignore_nan_grad: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct FrontendConf {
pub fs: usize,
pub window: String,
pub n_mels: usize,
pub frame_length: f32,
pub frame_shift: f32,
pub lfr_m: usize,
pub lfr_n: usize,
#[serde(skip_serializing_if = "Option::is_none")]
pub cmvn_file: Option<serde_yaml::Value>,
}
+207
View File
@@ -0,0 +1,207 @@
use std::collections::HashMap;
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::{Result, anyhow};
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
models::{
GenerateModel,
fun_asr_nano::{
config::FunASRNanoConfig, model::FunAsrNanoModel, processor::FunAsrNanoProcessor,
},
qwen3::config::{Qwen3Config, Qwen3GenerationConfig},
},
tokenizer::TokenizerModel,
utils::{
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
get_dtype, get_logit_processor,
},
};
pub struct FunAsrNanoGenerateModel {
tokenizer: TokenizerModel,
processor: FunAsrNanoProcessor,
fun_asr_nano: FunAsrNanoModel,
device: Device,
dtype: DType,
eos_token_id1: u32,
eos_token_id2: u32,
generation_config: Qwen3GenerationConfig,
model_name: String,
}
impl FunAsrNanoGenerateModel {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let llm_config_path = path.to_string() + "/Qwen3-0.6B";
let tokenizer = TokenizerModel::init(&llm_config_path)?;
let generation_config_path = llm_config_path.clone() + "/generation_config.json";
let generation_config: Qwen3GenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
let config_path = llm_config_path + "/config.json";
let llm_cfg: Qwen3Config = serde_json::from_slice(&std::fs::read(config_path)?)?;
let device = get_device(device);
let config_path = path.to_string() + "/config.yaml";
let cfg: FunASRNanoConfig = serde_yaml::from_slice(&std::fs::read(config_path)?)?;
let cfg_dtype = cfg.llm_conf.llm_dtype.as_str();
let dtype = get_dtype(dtype, cfg_dtype);
let processor = FunAsrNanoProcessor::new(&cfg.frontend_conf, &device)?;
let model_list = find_type_files(path, "pt")?;
let mut dict_to_hashmap = HashMap::new();
for m in model_list {
let dict = read_all_with_key(m, Some("state_dict"))?;
for (k, v) in dict {
dict_to_hashmap.insert(k, v);
}
}
let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device);
let fun_asr_nano = FunAsrNanoModel::new(vb, &cfg, &llm_cfg)?;
Ok(Self {
tokenizer,
processor,
fun_asr_nano,
device,
dtype,
eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config,
model_name: "fun-asr-nano".to_string(),
})
}
}
impl GenerateModel for FunAsrNanoGenerateModel {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let top_k = self.generation_config.top_k;
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let (speech, fbank_mask, mut input_ids) =
self.processor.process_info(&mes, &self.tokenizer)?;
let mut speech = Some(speech.to_dtype(self.dtype)?);
let mut fbank_mask = Some(&fbank_mask);
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024);
for _ in 0..sample_len {
let logits = self.fun_asr_nano.forward(
&input_ids,
speech.as_ref(),
fbank_mask,
seqlen_offset,
)?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
let next_token = logit_processor.sample(&logits)?;
generate.push(next_token);
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;
}
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));
Ok(response)
}
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ Send
+ Unpin
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let top_k = self.generation_config.top_k;
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let sample_len = mes.max_tokens.unwrap_or(1024);
let stream = stream! {
let mut error_tokens = Vec::new();
let mut speech = Some(speech.to_dtype(self.dtype)?);
let mut fbank_mask = Some(&fbank_mask);
let mut input_ids = input_ids;
for _ in 0..sample_len {
let logits = self.fun_asr_nano.forward(
&input_ids,
speech.as_ref(),
fbank_mask,
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)?;
speech = None;
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)))
}
}
+4
View File
@@ -0,0 +1,4 @@
pub mod config;
pub mod generate;
pub mod model;
pub mod processor;
+646
View File
@@ -0,0 +1,646 @@
use anyhow::Result;
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, ops::softmax_last_dim};
use crate::{
models::{
common::{
NaiveAttention, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
},
fun_asr_nano::config::FunASRNanoConfig,
qwen3::{config::Qwen3Config, model::Qwen3Model},
},
position_embed::sinusoidal_pe::SinusoidalPositionEncoderCat,
utils::tensor_utils::{get_equal_mask, mask_filled, masked_scatter_dim0},
};
pub struct MultiHeadedAttentionSANM {
head_dim: usize,
n_head: usize,
linear_out: Linear,
linear_q_k_v: Linear,
fsmn_block: Conv1d,
left_padding: usize,
right_padding: usize,
scaling: f64,
}
impl MultiHeadedAttentionSANM {
pub fn new(
vb: VarBuilder,
n_head: usize,
in_dim: usize,
hidden_dim: usize,
kernel_size: usize,
sanm_shfit: usize,
) -> Result<Self> {
let head_dim = hidden_dim / n_head;
let linear_out = linear(hidden_dim, hidden_dim, vb.pp("linear_out"))?;
let linear_q_k_v = linear(in_dim, hidden_dim * 3, vb.pp("linear_q_k_v"))?;
let fsmn_block = get_conv1d(
vb.pp("fsmn_block"),
hidden_dim,
hidden_dim,
kernel_size,
0,
1,
1,
hidden_dim,
false,
)?;
let mut left_padding = (kernel_size - 1) / 2;
if sanm_shfit > 0 {
left_padding += sanm_shfit;
}
let right_padding = kernel_size - 1 - left_padding;
let scaling = (head_dim as f64).powf(-0.5);
Ok(Self {
head_dim,
n_head,
linear_out,
linear_q_k_v,
fsmn_block,
left_padding,
right_padding,
scaling,
})
}
pub fn forward_fsmn(
&self,
inputs: &Tensor,
mask: Option<&Tensor>,
mask_shfit_chunk: Option<&Tensor>,
) -> Result<Tensor> {
let mut inputs = inputs.clone();
let mask = if let Some(mask) = mask {
let mut mask = mask.unsqueeze(D::Minus1)?.unsqueeze(0)?;
if let Some(mask_shfit_chunk) = mask_shfit_chunk {
mask = mask.broadcast_mul(mask_shfit_chunk)?;
}
inputs = inputs.broadcast_mul(&mask)?;
Some(mask)
} else {
None
};
let xs = inputs.transpose(1, 2)?;
let xs = xs.pad_with_zeros(D::Minus1, self.left_padding, self.right_padding)?;
let xs = self.fsmn_block.forward(&xs)?;
let xs = xs.transpose(1, 2)?;
let mut xs = xs.add(&inputs)?;
if let Some(mask) = mask {
xs = xs.broadcast_mul(&mask)?;
}
Ok(xs)
}
pub fn forward_qkv(&self, xs: &Tensor) -> Result<(Tensor, Tensor, Tensor, Tensor)> {
let (b, t, _) = xs.dims3()?;
let q_k_v = self
.linear_q_k_v
.forward(xs)?
.reshape((b, t, 3, self.n_head, ()))?
.permute((2, 0, 3, 1, 4))?
.contiguous()?;
let q_h = q_k_v.i(0)?.contiguous()?;
let k_h = q_k_v.i(1)?.contiguous()?;
let v_h = q_k_v.i(2)?.contiguous()?;
let v = v_h.transpose(1, 2)?.reshape((b, t, ()))?;
Ok((q_h, k_h, v_h, v))
}
pub fn forward_attention(
&self,
values: &Tensor,
scores: &Tensor,
mask: Option<&Tensor>,
mask_att_chunk_encoder: Option<&Tensor>,
) -> Result<Tensor> {
let bs = scores.dim(0)?;
let attn = if let Some(mask) = mask {
let mask = if let Some(mask_att_chunk_encoder) = mask_att_chunk_encoder {
mask.mul(mask_att_chunk_encoder)?
} else {
mask.clone()
};
// mask: rank = 2
let mask = get_equal_mask(&mask, 0)?;
let scores = mask_filled(scores, &mask, f32::NEG_INFINITY)?;
let attn = softmax_last_dim(&scores)?;
mask_filled(&attn, &mask, 0.0)?
} else {
softmax_last_dim(scores)?
};
let xs = attn.matmul(values)?;
let xs =
xs.transpose(1, 2)?
.contiguous()?
.reshape((bs, (), self.n_head * self.head_dim))?;
let xs = self.linear_out.forward(&xs)?;
Ok(xs)
}
pub fn forward_simple(&self, xs: &Tensor) -> Result<Tensor> {
let (b, t, _) = xs.dims3()?;
let q_k_v = self.linear_q_k_v.forward(xs)?;
let dim = self.head_dim * self.n_head;
let q_h = q_k_v
.narrow(D::Minus1, 0, dim)?
.reshape((b, t, self.n_head, ()))?
.permute((0, 2, 1, 3))?;
let k_h = q_k_v
.narrow(D::Minus1, dim, dim)?
.reshape((b, t, self.n_head, ()))?
.permute((0, 2, 1, 3))?;
let v = q_k_v.narrow(D::Minus1, dim * 2, dim)?;
let v_h = v.reshape((b, t, self.n_head, ()))?.permute((0, 2, 1, 3))?;
let fsmn_memory = v.transpose(1, 2)?;
let fsmn_memory = fsmn_memory
.pad_with_zeros(D::Minus1, self.left_padding, self.right_padding)?
.contiguous()?;
let fsmn_memory = self.fsmn_block.forward(&fsmn_memory)?;
// let fsmn_memory = conv1d_group_parallel(&fsmn_memory, &self.fsmn_block)?;
let fsmn_memory = fsmn_memory.transpose(1, 2)?;
let fsmn_memory = fsmn_memory.add(&v)?;
let att_outs = eager_attention_forward(&q_h, &k_h, &v_h, None, None, self.scaling)?;
let att_outs = att_outs.reshape((b, t, ()))?;
let att_outs = self.linear_out.forward(&att_outs)?;
let att_outs = att_outs.add(&fsmn_memory)?;
Ok(att_outs)
}
pub fn forward(
&self,
xs: &Tensor,
mask: Option<&Tensor>,
mask_shfit_chunk: Option<&Tensor>,
mask_att_chunk_encoder: Option<&Tensor>,
) -> Result<Tensor> {
let (q_h, k_h, v_h, v) = self.forward_qkv(xs)?;
let fsmn_memory = self.forward_fsmn(&v, mask, mask_shfit_chunk)?;
let q_h = q_h.affine(self.scaling, 0.0)?;
let scores = q_h.matmul(&k_h.transpose(D::Minus2, D::Minus1)?)?;
let attn_outs = self.forward_attention(&v_h, &scores, mask, mask_att_chunk_encoder)?;
let att_outs = attn_outs.add(&fsmn_memory)?;
Ok(att_outs)
}
}
pub struct EncoderLayerSANM {
self_attn: MultiHeadedAttentionSANM,
feed_forward: TwoLinearMLP,
norm1: LayerNorm,
norm2: LayerNorm,
concat_linear: Option<Linear>,
normalize_before: bool,
in_dim: usize,
hidden_dim: usize,
}
impl EncoderLayerSANM {
pub fn new(
vb: VarBuilder,
in_dim: usize,
hidden_dim: usize,
n_head: usize,
kernel_size: usize,
sanm_shfit: usize,
hidden_units: usize,
normalize_before: bool,
concat_after: bool,
) -> Result<Self> {
let self_attn = MultiHeadedAttentionSANM::new(
vb.pp("self_attn"),
n_head,
in_dim,
hidden_dim,
kernel_size,
sanm_shfit,
)?;
let feed_forward = TwoLinearMLP::new(
vb.pp("feed_forward"),
hidden_dim,
hidden_units,
hidden_dim,
candle_nn::Activation::Relu,
true,
"w_1",
"w_2",
)?;
let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, in_dim)?;
let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, hidden_dim)?;
let concat_linear = if concat_after {
let lin = linear(hidden_dim * 2, hidden_dim, vb.pp("concat_linear"))?;
Some(lin)
} else {
None
};
Ok(Self {
self_attn,
feed_forward,
norm1,
norm2,
concat_linear,
normalize_before,
in_dim,
hidden_dim,
})
}
pub fn forward(
&self,
xs: &Tensor,
mask: Option<&Tensor>,
mask_shfit_chunk: Option<&Tensor>,
mask_att_chunk_encoder: Option<&Tensor>,
) -> Result<Tensor> {
let stoch_layer_coeff = 1.0f64;
let residual = xs.clone();
let mut xs = if self.normalize_before {
self.norm1.forward(xs)?
} else {
xs.clone()
};
if self.concat_linear.is_some() {
let attn =
self.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?;
let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?;
if self.in_dim == self.hidden_dim {
let x_concat = self
.concat_linear
.as_ref()
.unwrap()
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&x_concat)?;
} else {
xs = self
.concat_linear
.as_ref()
.unwrap()
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
}
} else if self.in_dim == self.hidden_dim {
let attn = self
.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&attn)?;
} else {
xs = self
.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?
.affine(stoch_layer_coeff, 0.0)?;
}
if !self.normalize_before {
xs = self.norm1.forward(&xs)?;
}
let residual = xs.clone();
if self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
xs = self
.feed_forward
.forward(&xs)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&xs)?;
if !self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
Ok(xs)
}
pub fn forward_simple(&self, xs: &Tensor) -> Result<Tensor> {
let residual = xs.clone();
let mut xs = self.norm1.forward(xs)?;
if self.in_dim == self.hidden_dim {
let attn = self.self_attn.forward_simple(&xs)?;
xs = residual.add(&attn)?;
} else {
xs = self.self_attn.forward_simple(&xs)?;
}
let residual = xs.clone();
let xs = self.norm2.forward(&xs)?;
let xs = self.feed_forward.forward(&xs)?;
let xs = residual.add(&xs)?;
Ok(xs)
}
}
pub struct SenseVoiceEncoderSmall {
embed: SinusoidalPositionEncoderCat,
encoders0: EncoderLayerSANM,
encoders: Vec<EncoderLayerSANM>,
tp_encoders: Vec<EncoderLayerSANM>,
after_norm: LayerNorm,
tp_norm: LayerNorm,
scaling: f64,
}
impl SenseVoiceEncoderSmall {
pub fn new(
vb: VarBuilder,
input_size: usize,
output_size: usize,
attention_heads: usize,
linear_units: usize,
num_blocks: usize,
tp_blocks: usize,
normalize_before: bool,
kernel_size: usize,
sanm_shfit: usize,
) -> Result<Self> {
let embed = SinusoidalPositionEncoderCat::new(Some(input_size), true, vb.device())?;
let encoders0 = EncoderLayerSANM::new(
vb.pp("encoders0.0"),
input_size,
output_size,
attention_heads,
kernel_size,
sanm_shfit,
linear_units,
normalize_before,
false,
)?;
let mut encoders = vec![];
let vb_encoders = vb.pp("encoders");
for i in 0..(num_blocks - 1) {
let encoder_i = EncoderLayerSANM::new(
vb_encoders.pp(i),
output_size,
output_size,
attention_heads,
kernel_size,
sanm_shfit,
linear_units,
normalize_before,
false,
)?;
encoders.push(encoder_i);
}
let vb_tp_encoders = vb.pp("tp_encoders");
let mut tp_encoders = vec![];
for i in 0..tp_blocks {
let tp_blocks_i = EncoderLayerSANM::new(
vb_tp_encoders.pp(i),
output_size,
output_size,
attention_heads,
kernel_size,
sanm_shfit,
linear_units,
normalize_before,
false,
)?;
tp_encoders.push(tp_blocks_i);
}
let after_norm = get_layer_norm(vb.pp("after_norm"), 1e-5, output_size)?;
let tp_norm = get_layer_norm(vb.pp("tp_norm"), 1e-5, output_size)?;
let scaling = (output_size as f64).powf(0.5);
Ok(Self {
embed,
encoders0,
encoders,
tp_encoders,
after_norm,
tp_norm,
scaling,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let xs = xs.affine(self.scaling, 0.0)?;
let xs = self.embed.forward(&xs, 0)?;
let mut xs = self.encoders0.forward_simple(&xs)?;
for encoder_layer in &self.encoders {
xs = encoder_layer.forward_simple(&xs)?;
}
xs = self.after_norm.forward(&xs)?;
for tp_layer in &self.tp_encoders {
xs = tp_layer.forward_simple(&xs)?;
}
xs = self.tp_norm.forward(&xs)?;
Ok(xs)
}
}
pub struct AdaptorEncoderLayer {
self_attn: NaiveAttention,
feed_forward: TwoLinearMLP,
norm1: LayerNorm,
norm2: LayerNorm,
concat_linear: Option<Linear>,
normalize_before: bool,
}
impl AdaptorEncoderLayer {
pub fn new(
vb: VarBuilder,
llm_dim: usize,
n_head: usize,
normalize_before: bool,
concat_after: bool,
) -> Result<Self> {
let self_attn = NaiveAttention::new(
vb.pp("self_attn"),
llm_dim,
n_head,
n_head,
None,
true,
Some("linear_q"),
Some("linear_k"),
Some("linear_v"),
Some("linear_out"),
)?;
let feed_forward = TwoLinearMLP::new(
vb.pp("feed_forward"),
llm_dim,
llm_dim / 4,
llm_dim,
candle_nn::Activation::Relu,
true,
"w_1",
"w_2",
)?;
let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, llm_dim)?;
let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, llm_dim)?;
let concat_linear = if concat_after {
let lin = linear(llm_dim * 2, llm_dim, vb.pp("concat_linear"))?;
Some(lin)
} else {
None
};
Ok(Self {
self_attn,
feed_forward,
norm1,
norm2,
concat_linear,
normalize_before,
})
}
pub fn forward(&self, xs: &Tensor, mask: Option<&Tensor>) -> Result<Tensor> {
let stoch_layer_coeff = 1.0f64;
let residual = xs.clone();
let mut xs = if self.normalize_before {
self.norm1.forward(xs)?
} else {
xs.clone()
};
if self.concat_linear.is_some() {
let attn = self.self_attn.forward(&xs, None, None, mask, false)?;
let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?;
let x_concat = self
.concat_linear
.as_ref()
.unwrap()
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&x_concat)?;
} else {
let attn = self
.self_attn
.forward(&xs, None, None, mask, false)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&attn)?;
}
if !self.normalize_before {
xs = self.norm1.forward(&xs)?;
}
let residual = xs.clone();
if self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
xs = self
.feed_forward
.forward(&xs)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&xs)?;
if !self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
Ok(xs)
}
}
pub struct AudioAdaptor {
k: usize,
linear1: Linear,
linear2: Linear,
blocks: Vec<AdaptorEncoderLayer>,
}
impl AudioAdaptor {
pub fn new(
vb: VarBuilder,
downsample_rate: usize,
encoder_dim: usize,
llm_dim: usize,
ffn_dim: usize,
n_layer: usize,
attention_heads: usize,
) -> Result<Self> {
let linear1 = linear(encoder_dim * downsample_rate, ffn_dim, vb.pp("linear1"))?;
let linear2 = linear(ffn_dim, llm_dim, vb.pp("linear2"))?;
let mut blocks = vec![];
let vb_blocks = vb.pp("blocks");
for i in 0..n_layer {
let layer =
AdaptorEncoderLayer::new(vb_blocks.pp(i), llm_dim, attention_heads, true, false)?;
blocks.push(layer);
}
Ok(Self {
k: downsample_rate,
linear1,
linear2,
blocks,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let (bs, seq_len, dim) = xs.dims3()?;
let chunk_num = (seq_len - 1) / self.k + 1;
let pad_num = chunk_num * self.k - seq_len;
let xs = xs.pad_with_zeros(1, 0, pad_num)?;
let xs = xs.contiguous()?.reshape((bs, chunk_num, dim * self.k))?;
let xs = self.linear1.forward(&xs)?.relu()?;
let mut xs = self.linear2.forward(&xs)?;
for block in &self.blocks {
xs = block.forward(&xs, None)?;
}
Ok(xs)
}
}
pub struct FunAsrNanoModel {
audio_encoder: SenseVoiceEncoderSmall,
audio_adaptor: AudioAdaptor,
llm: Qwen3Model,
}
impl FunAsrNanoModel {
pub fn new(vb: VarBuilder, config: &FunASRNanoConfig, llm_cfg: &Qwen3Config) -> Result<Self> {
let input_size = config.frontend_conf.lfr_m * config.frontend_conf.n_mels;
let audio_encoder = SenseVoiceEncoderSmall::new(
vb.pp("audio_encoder"),
input_size,
config.audio_encoder_conf.output_size,
config.audio_encoder_conf.attention_heads,
config.audio_encoder_conf.linear_units,
config.audio_encoder_conf.num_blocks,
config.audio_encoder_conf.tp_blocks,
config.audio_encoder_conf.normalize_before,
config.audio_encoder_conf.kernel_size,
config.audio_encoder_conf.sanm_shfit,
)?;
let audio_adaptor = AudioAdaptor::new(
vb.pp("audio_adaptor"),
config.audio_adaptor_conf.downsample_rate,
config.audio_adaptor_conf.encoder_dim,
config.audio_adaptor_conf.llm_dim,
config.audio_adaptor_conf.ffn_dim,
config.audio_adaptor_conf.n_layer,
8,
)?;
let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"))?;
Ok(Self {
audio_encoder,
audio_adaptor,
llm,
})
}
pub fn forward(
&mut self,
input_ids: &Tensor,
speech: Option<&Tensor>,
fbank_mask: Option<&Tensor>,
seqlen_offset: usize,
) -> Result<Tensor> {
let mut inputs_embeds = self.llm.embedding_token_id(input_ids)?;
if let Some(speech) = speech
&& let Some(fbank_mask) = fbank_mask
{
let speech = self.audio_encoder.forward(speech)?;
let encoder_out = self.audio_adaptor.forward(&speech)?;
let speech_token_len = fbank_mask.sum_all()?.to_scalar::<u32>()?;
let audio_embed = encoder_out
.squeeze(0)?
.narrow(0, 0, speech_token_len as usize)?;
inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embed, fbank_mask)?;
}
let logits = self
.llm
.forward(None, Some(&inputs_embeds), seqlen_offset)?;
Ok(logits)
}
pub fn clear_kv_cache(&mut self) {
self.llm.clear_kv_cache();
}
}
+110
View File
@@ -0,0 +1,110 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{D, Device, Tensor};
use crate::{
models::fun_asr_nano::config::FrontendConf,
tokenizer::TokenizerModel,
utils::{
audio_utils::{
apply_lfr, extract_audios, get_waveform_and_window_properties, kaldi_fbank,
kaldi_get_mel_banks,
},
extract_user_text,
},
};
pub struct FunAsrNanoProcessor {
fronted_conf: FrontendConf,
device: Device,
prompt_prefix: String,
prompt_suffix: String,
window_shift: usize,
window_size: usize,
padded_window_size: usize,
mel_energies: Tensor,
}
impl FunAsrNanoProcessor {
pub fn new(fronted_conf: &FrontendConf, device: &Device) -> Result<Self> {
let (window_shift, window_size, padded_window_size) = get_waveform_and_window_properties(
fronted_conf.fs,
fronted_conf.frame_shift,
fronted_conf.frame_length,
true,
)?;
let (mel_energies, _) = kaldi_get_mel_banks(
fronted_conf.n_mels,
padded_window_size,
fronted_conf.fs as f32,
20.0,
0.0,
device,
)?;
let mel_energies = mel_energies.pad_with_zeros(D::Minus1, 0, 1)?.t()?;
Ok(Self {
fronted_conf: fronted_conf.clone(),
device: device.clone(),
prompt_prefix:
"<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n"
.to_string(),
prompt_suffix: "<|im_end|>\n<|im_start|>assistant\n".to_string(),
window_shift,
window_size,
padded_window_size,
mel_energies,
})
}
pub fn extract_fbank(&self, audio: &Tensor) -> Result<(Tensor, usize)> {
let waveform = audio.affine(32768.0, 0.0)?;
let mut mat = kaldi_fbank(
&waveform,
&self.mel_energies,
self.window_shift,
self.window_size,
self.padded_window_size,
1.0,
// 0.0,
// "hamming",
// self.fronted_conf.fs,
// true,
)?;
mat = mat.squeeze(0)?;
if self.fronted_conf.lfr_m != 1 || self.fronted_conf.lfr_n != 1 {
mat = apply_lfr(&mat, self.fronted_conf.lfr_m, self.fronted_conf.lfr_n)?;
}
let feat_length = mat.dim(0)?;
let mat = mat.unsqueeze(0)?;
Ok((mat, feat_length))
}
pub fn process_info(
&self,
mes: &ChatCompletionParameters,
tokenizer: &TokenizerModel,
) -> Result<(Tensor, Tensor, Tensor)> {
let user_text = extract_user_text(mes)?;
let sub_prompt = self.prompt_prefix.clone() + &user_text;
let mut source_ids = vec![];
let mut fbank_mask = vec![];
let sub_token = tokenizer.text_encode_vec(sub_prompt, true)?;
source_ids.extend_from_slice(&sub_token);
fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]);
let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?;
let audio = &audio_tensors[0];
let (speech, speech_lengths) = self.extract_fbank(audio)?;
let olens = 1 + (speech_lengths - 3 + 2) / 2;
let olens = 1 + (olens - 3 + 2) / 2;
let fake_token_len = (olens - 1) / 2 + 1;
source_ids.extend_from_slice(&vec![0u32; fake_token_len]);
fbank_mask.extend_from_slice(&vec![1u32; fake_token_len]);
let sub_token = tokenizer.text_encode_vec(self.prompt_suffix.clone(), true)?;
source_ids.extend_from_slice(&sub_token);
fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]);
let input_ids = Tensor::from_slice(&source_ids, (1, source_ids.len()), &self.device)?;
let fbank_mask = Tensor::from_slice(&fbank_mask, (1, fbank_mask.len()), &self.device)?;
Ok((speech, fbank_mask, input_ids))
}
}
+1 -1
View File
@@ -70,7 +70,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
Some(s) => s as u64,
};
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let render_text = self.chat_template.apply_chat_template(&mes)?;
let render_text: String = self.chat_template.apply_chat_template(&mes)?;
let (input_features, audio_token_lengths, replace_text) =
self.processor.process_info(&mes, &render_text)?;
let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
+14 -36
View File
@@ -3,12 +3,13 @@ use std::f32;
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor};
use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
use crate::{
models::glm_asr_nano::config::GlmAsrNanoProcessorConfig,
utils::{
audio_utils::{create_hann_window, extract_audios, mel_filter_bank, stft_audio},
audio_utils::{
apply_stft, create_hann_window, extract_audios, extract_frames, mel_filter_bank,
},
tensor_utils::{pad_reflect_last_dim, split_tensor},
},
};
@@ -81,48 +82,19 @@ impl GlmAsrNanoProcessor {
})
}
/// 提取音频帧
pub fn extract_frames(&self, waveform: &Tensor, n_frames: usize) -> Result<Tensor> {
let mut frames = Vec::with_capacity(n_frames);
for i in 0..n_frames {
let start = i * self.hop_length;
let frame = waveform.narrow(D::Minus1, start, self.n_fft)?;
frames.push(frame);
}
let result = Tensor::cat(&frames, D::Minus1)?;
let bs = result.dim(0)?;
let reshaped = result.reshape((bs, n_frames, self.n_fft))?;
Ok(reshaped)
}
pub fn extract_fbank_features(&self, waveform: &Tensor) -> Result<Tensor> {
let pad = self.n_fft / 2;
let waveform = pad_reflect_last_dim(waveform, (pad, pad))?;
let (batch_size, samples) = waveform.dims2()?;
let (_, samples) = waveform.dims2()?;
// 计算输出维度
let n_frames = (samples - self.n_fft) / self.hop_length + 1;
// (bs, n_frames, n_fft)
let frames = self.extract_frames(&waveform, n_frames)?;
let frames = extract_frames(&waveform, self.n_fft, self.hop_length)?;
// 应用汉明窗口
let result = frames.broadcast_mul(&self.window)?;
// 傅立叶变换
let mut wave_fft = vec![];
for bs in 0..batch_size {
let wave_i = result.i(bs)?;
let wave_i_vec = wave_i.to_vec2::<f32>()?;
let wave_i_fft_vec: Result<Vec<Vec<f32>>> = wave_i_vec
.par_iter()
.map(|frame_wave| stft_audio(self.n_fft, frame_wave))
.collect();
let wave_i_fft_vec = wave_i_fft_vec?;
let wave_i_fft = Tensor::new(wave_i_fft_vec, &self.device)?.unsqueeze(0)?;
wave_fft.push(wave_i_fft);
}
let magnitudes = Tensor::cat(&wave_fft, 0)?.transpose(D::Minus1, D::Minus2)?;
let magnitudes = apply_stft(&result)?.transpose(D::Minus1, D::Minus2)?;
let magnitudes = magnitudes.narrow(D::Minus1, 0, n_frames - 1)?;
let mel_spec = self.mel_filters.broadcast_matmul(&magnitudes)?;
let mel_spec = mel_spec.clamp(1e-10f32, f32::INFINITY)?;
@@ -141,12 +113,18 @@ impl GlmAsrNanoProcessor {
let audio_len = audio.dim(0)?;
let pad_num = self.n_samples - audio_len;
let audio_pad = audio.pad_with_zeros(0, 0, pad_num)?;
let audio_pad = if pad_num > 0 {
audio.pad_with_zeros(0, 0, pad_num)?
} else {
audio
};
// (n_samples) -> (1, n_samples)
let audio_pad = audio_pad.unsqueeze(0)?;
pad_audio.push(audio_pad);
let mut mask = vec![1u32; audio_len];
mask.extend_from_slice(&vec![0u32; pad_num]);
if pad_num > 0 {
mask.extend_from_slice(&vec![0u32; pad_num]);
}
input_features_mask.push(mask);
}
let input_features = Tensor::cat(&pad_audio, 0)?;
+3
View File
@@ -111,6 +111,9 @@ impl MiniCPMDecoderLayer {
None,
false,
None,
None,
None,
None,
)?;
let mlp = GateUpDownMLP::new(
vb.pp("mlp"),
+23 -2
View File
@@ -1,10 +1,12 @@
pub mod common;
pub mod deepseek_ocr;
pub mod fun_asr_nano;
pub mod glm_asr_nano;
pub mod hunyuan_ocr;
pub mod minicpm4;
pub mod paddleocr_vl;
pub mod qwen2_5vl;
pub mod qwen3;
pub mod qwen3vl;
pub mod rmbg2_0;
pub mod voxcpm;
@@ -17,11 +19,12 @@ use rocket::futures::Stream;
use crate::models::{
deepseek_ocr::generate::DeepseekOCRGenerateModel,
fun_asr_nano::generate::FunAsrNanoGenerateModel,
glm_asr_nano::generate::GlmAsrNanoGenerateModel,
hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel,
paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel,
qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model,
voxcpm::generate::VoxCPMGenerate,
qwen3::generate::Qwen3GenerateModel, qwen3vl::generate::Qwen3VLGenerateModel,
rmbg2_0::generate::RMBG2_0Model, voxcpm::generate::VoxCPMGenerate,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
@@ -32,6 +35,8 @@ pub enum WhichModel {
Qwen2_5vl3B,
#[value(name = "qwen2.5vl-7b")]
Qwen2_5vl7B,
#[value(name = "qwen3-0.6b")]
Qwen3_0_6B,
#[value(name = "qwen3vl-2b")]
Qwen3vl2B,
#[value(name = "qwen3vl-4b")]
@@ -54,6 +59,8 @@ pub enum WhichModel {
VoxCPM1_5,
#[value(name = "glm-asr-nano-2512")]
GlmASRNano2512,
#[value(name = "fun-asr-nano-2512")]
FunASRNano2512,
}
pub trait GenerateModel {
@@ -74,6 +81,7 @@ pub trait GenerateModel {
pub enum ModelInstance<'a> {
MiniCPM4(MiniCPMGenerateModel<'a>),
Qwen2_5VL(Qwen2_5VLGenerateModel<'a>),
Qwen3(Qwen3GenerateModel<'a>),
Qwen3VL(Qwen3VLGenerateModel<'a>),
DeepSeekOCR(DeepseekOCRGenerateModel),
HunyuanOCR(HunyuanOCRGenerateModel<'a>),
@@ -81,6 +89,7 @@ pub enum ModelInstance<'a> {
RMBG2_0(Box<RMBG2_0Model>),
VoxCPM(Box<VoxCPMGenerate>),
GlmASRNano(GlmAsrNanoGenerateModel<'a>),
FunASRNano(FunAsrNanoGenerateModel),
}
impl<'a> GenerateModel for ModelInstance<'a> {
@@ -88,6 +97,7 @@ impl<'a> GenerateModel for ModelInstance<'a> {
match self {
ModelInstance::MiniCPM4(model) => model.generate(mes),
ModelInstance::Qwen2_5VL(model) => model.generate(mes),
ModelInstance::Qwen3(model) => model.generate(mes),
ModelInstance::Qwen3VL(model) => model.generate(mes),
ModelInstance::DeepSeekOCR(model) => model.generate(mes),
ModelInstance::HunyuanOCR(model) => model.generate(mes),
@@ -95,6 +105,7 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::RMBG2_0(model) => model.generate(mes),
ModelInstance::VoxCPM(model) => model.generate(mes),
ModelInstance::GlmASRNano(model) => model.generate(mes),
ModelInstance::FunASRNano(model) => model.generate(mes),
}
}
@@ -112,6 +123,7 @@ impl<'a> GenerateModel for ModelInstance<'a> {
match self {
ModelInstance::MiniCPM4(model) => model.generate_stream(mes),
ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes),
ModelInstance::Qwen3(model) => model.generate_stream(mes),
ModelInstance::Qwen3VL(model) => model.generate_stream(mes),
ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes),
ModelInstance::HunyuanOCR(model) => model.generate_stream(mes),
@@ -119,6 +131,7 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::RMBG2_0(model) => model.generate_stream(mes),
ModelInstance::VoxCPM(model) => model.generate_stream(mes),
ModelInstance::GlmASRNano(model) => model.generate_stream(mes),
ModelInstance::FunASRNano(model) => model.generate_stream(mes),
}
}
}
@@ -137,6 +150,10 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_
let model = Qwen2_5VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen2_5VL(model)
}
WhichModel::Qwen3_0_6B => {
let model = Qwen3GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3(model)
}
WhichModel::Qwen3vl2B => {
let model = Qwen3VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen3VL(model)
@@ -181,6 +198,10 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_
let model = GlmAsrNanoGenerateModel::init(path, None, None)?;
ModelInstance::GlmASRNano(model)
}
WhichModel::FunASRNano2512 => {
let model = FunAsrNanoGenerateModel::init(path, None, None)?;
ModelInstance::FunASRNano(model)
}
};
Ok(model)
}
+44
View File
@@ -0,0 +1,44 @@
use candle_nn::Activation;
use serde::Deserialize;
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct Qwen3Config {
pub attention_bias: bool,
pub attention_dropout: f64,
pub bos_token_id: u32,
pub eos_token_id: u32,
pub head_dim: usize,
pub hidden_act: Activation,
pub hidden_size: usize,
pub initializer_range: f64,
pub intermediate_size: usize,
pub max_position_embeddings: usize,
pub max_window_layers: usize,
pub num_attention_heads: usize,
pub num_hidden_layers: usize,
pub num_key_value_heads: usize,
pub rms_norm_eps: f64,
pub rope_theta: f32,
pub tie_word_embeddings: bool,
pub torch_dtype: String,
pub use_cache: bool,
pub use_sliding_window: bool,
pub vocab_size: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct Qwen3GenerationConfig {
pub bos_token_id: usize,
pub pad_token_id: usize,
pub do_sample: bool,
pub eos_token_id: Vec<usize>,
pub top_p: f32,
pub top_k: usize,
pub temperature: f32,
#[serde(default = "default_repetition_penalty")]
pub repetition_penalty: f32,
}
fn default_repetition_penalty() -> f32 {
1.0
}
+172
View File
@@ -0,0 +1,172 @@
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::{Result, anyhow};
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
use crate::models::qwen3::model::Qwen3Model;
// use crate::models::GenerateStream;
use crate::utils::{
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
get_dtype, get_logit_processor,
};
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
pub struct Qwen3GenerateModel<'a> {
chat_template: ChatTemplate<'a>,
tokenizer: TokenizerModel,
qwen3: Qwen3Model,
device: Device,
eos_token_id1: u32,
eos_token_id2: u32,
generation_config: Qwen3GenerationConfig,
model_name: String,
}
impl<'a> Qwen3GenerateModel<'a> {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let chat_template = ChatTemplate::init(path)?;
let tokenizer = TokenizerModel::init(path)?;
let config_path = path.to_string() + "/config.json";
let cfg: Qwen3Config = serde_json::from_slice(&std::fs::read(config_path)?)?;
let device = &get_device(device);
let cfg_dtype = cfg.torch_dtype.as_str();
let dtype = get_dtype(dtype, cfg_dtype);
let model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
let qwen3 = Qwen3Model::new(&cfg, vb)?;
let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: Qwen3GenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
Ok(Qwen3GenerateModel {
chat_template,
tokenizer,
qwen3,
device: device.clone(),
eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config,
model_name: "qwen3".to_string(),
})
}
}
impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let top_k = self.generation_config.top_k;
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(2048);
for _ in 0..sample_len {
let logits = self.qwen3.forward(Some(&input_ids), None, seqlen_offset)?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
let next_token = logit_processor.sample(&logits)?;
generate.push(next_token);
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)?;
}
let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?;
self.qwen3.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ Send
+ Unpin
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let top_k = self.generation_config.top_k;
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let sample_len = mes.max_tokens.unwrap_or(512);
let stream = stream! {
let mut error_tokens = Vec::new();
for _ in 0..sample_len {
let logits = self.qwen3.forward(
Some(&input_ids),
None,
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)?;
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)?;
}
self.qwen3.clear_kv_cache();
};
Ok(Box::new(Box::pin(stream)))
}
}
+3
View File
@@ -0,0 +1,3 @@
pub mod config;
pub mod generate;
pub mod model;
+277
View File
@@ -0,0 +1,277 @@
use anyhow::Result;
use candle_core::Tensor;
use candle_nn::{
Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear, linear_no_bias, rms_norm,
};
use crate::{
models::{
common::{GateUpDownMLP, eager_attention_forward},
qwen3::config::Qwen3Config,
},
position_embed::rope::{RoPE, apply_rotary_pos_emb},
utils::tensor_utils::prepare_causal_attention_mask,
};
pub struct Qwen3Attention {
q_proj: Linear,
k_proj: Linear,
v_proj: Linear,
o_proj: Linear,
q_norm: RmsNorm,
k_norm: RmsNorm,
num_attention_heads: usize,
num_key_value_heads: usize,
num_kv_groups: usize,
head_dim: usize,
scaling: f64,
kv_cache: Option<(Tensor, Tensor)>,
}
impl Qwen3Attention {
pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result<Self> {
let hidden_size = config.hidden_size;
let num_attention_heads = config.num_attention_heads;
let head_dim = config.head_dim;
let num_key_value_heads = config.num_key_value_heads;
let num_kv_groups = num_attention_heads / num_key_value_heads;
let scaling = 1f64 / f64::sqrt(head_dim as f64);
let (q_proj, k_proj, v_proj, o_proj) = if config.attention_bias {
let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
} else {
let q_proj =
linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj =
linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
};
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
let k_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?;
Ok(Self {
q_proj,
k_proj,
v_proj,
o_proj,
q_norm,
k_norm,
num_attention_heads,
num_key_value_heads,
num_kv_groups,
head_dim,
scaling,
kv_cache: None,
})
}
pub fn forward(
&mut self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
let (b_sz, q_len, _) = xs.dims3()?;
let query_states = self.q_proj.forward(xs)?.reshape((
b_sz,
q_len,
self.num_attention_heads,
self.head_dim,
))?;
let query_states = self.q_norm.forward(&query_states)?.transpose(1, 2)?;
let key_states = self.k_proj.forward(xs)?.reshape((
b_sz,
q_len,
self.num_key_value_heads,
self.head_dim,
))?;
let key_states = self.k_norm.forward(&key_states)?.transpose(1, 2)?;
let value_states = self.v_proj.forward(xs)?;
let value_states = value_states
.reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))?
.transpose(1, 2)?;
let (query_states, key_states) =
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?;
let (key_states, value_states) = match &self.kv_cache {
None => (key_states, value_states),
Some((prev_k, prev_v)) => {
let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
(key_states, value_states)
}
};
self.kv_cache = Some((key_states.clone(), value_states.clone()));
let attn_output = eager_attention_forward(
&query_states,
&key_states,
&value_states,
Some(self.num_kv_groups),
attention_mask,
self.scaling,
)?;
let attn_output =
attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?;
let attn_output = attn_output.apply(&self.o_proj)?;
Ok(attn_output)
}
pub fn clear_kv_cache(&mut self) {
self.kv_cache = None
}
}
pub struct Qwen3DecoderLayer {
self_attn: Qwen3Attention,
mlp: GateUpDownMLP,
input_layernorm: RmsNorm,
post_attention_layernorm: RmsNorm,
}
impl Qwen3DecoderLayer {
pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result<Self> {
let self_attn = Qwen3Attention::new(config, vb.pp("self_attn"))?;
let mlp = GateUpDownMLP::new(
vb.pp("mlp"),
config.hidden_size,
config.intermediate_size,
config.hidden_act,
false,
)?;
let input_layernorm = rms_norm(
config.hidden_size,
config.rms_norm_eps,
vb.pp("input_layernorm"),
)?;
let post_attention_layernorm = rms_norm(
config.hidden_size,
config.rms_norm_eps,
vb.pp("post_attention_layernorm"),
)?;
Ok(Self {
self_attn,
mlp,
input_layernorm,
post_attention_layernorm,
})
}
pub fn forward(
&mut self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?;
let xs = residual.add(&xs)?;
let residual = xs.clone();
let xs = self.post_attention_layernorm.forward(&xs)?;
let xs = self.mlp.forward(&xs)?;
let xs = residual.add(&xs)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
self.self_attn.clear_kv_cache();
}
}
pub struct Qwen3Model {
embed_tokens: Embedding,
layers: Vec<Qwen3DecoderLayer>,
norm: RmsNorm,
rotary_emb: RoPE,
lm_head: Linear,
}
impl Qwen3Model {
pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result<Self> {
let vb = vb.pp("model");
let vocab_size = config.vocab_size;
let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![];
let vb_l = vb.pp("layers");
for layer_idx in 0..config.num_hidden_layers {
let layer = Qwen3DecoderLayer::new(config, vb_l.pp(layer_idx))?;
layers.push(layer)
}
let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?;
let head_dim = config.head_dim;
let rotary_emb = RoPE::new(head_dim, config.rope_theta, vb.device())?;
let lm_head = if config.tie_word_embeddings {
Linear::new(embed_tokens.embeddings().clone(), None)
} else {
linear_no_bias(config.hidden_size, config.vocab_size, vb.pp("lm_head"))?
};
Ok(Self {
embed_tokens,
layers,
norm,
rotary_emb,
lm_head,
})
}
pub fn forward(
&mut self,
input_ids: Option<&Tensor>,
inputs_embeds: Option<&Tensor>,
seqlen_offset: usize,
) -> Result<Tensor> {
if input_ids.is_none() && inputs_embeds.is_none() {
return Err(anyhow::anyhow!(
"You must specify exactly one of input_ids or inputs_embeds"
));
}
let inputs_embeds = if let Some(inputs_embeds) = inputs_embeds {
inputs_embeds.clone()
} else {
let input_ids = input_ids.unwrap();
self.embedding_token_id(input_ids)?
};
let (bs, seq_len, _) = inputs_embeds.dims3()?;
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
bs,
seq_len,
0,
inputs_embeds.device(),
)?)
}
};
let (cos, sin) = self
.rotary_emb
.forward(seqlen_offset, seq_len, inputs_embeds.device())?;
let mut hidden_states = inputs_embeds;
for decode_layer in &mut self.layers {
hidden_states =
decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?;
}
hidden_states = self.norm.forward(&hidden_states)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
let logits = self.lm_head.forward(&hidden_state)?;
Ok(logits)
}
pub fn embedding_token_id(&self, input_ids: &Tensor) -> Result<Tensor> {
Ok(self.embed_tokens.forward(input_ids)?)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.layers.iter_mut() {
layer.clear_kv_cache()
}
}
}
+28 -12
View File
@@ -1,5 +1,7 @@
use candle_nn::Activation;
use crate::models::qwen3::config::Qwen3Config;
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct Size {
pub longest_edge: usize,
@@ -46,6 +48,32 @@ pub struct Qwen3VLTextConfig {
pub vocab_size: usize,
}
pub fn qwen3vl_text_config2qwen3_config(cfg: &Qwen3VLTextConfig) -> Qwen3Config {
Qwen3Config {
attention_bias: cfg.attention_bias,
attention_dropout: cfg.attention_dropout as f64,
bos_token_id: cfg.bos_token_id as u32,
eos_token_id: cfg.eos_token_id as u32,
head_dim: cfg.head_dim,
hidden_act: cfg.hidden_act,
hidden_size: cfg.hidden_size,
initializer_range: cfg.initializer_range as f64,
intermediate_size: cfg.intermediate_size,
max_position_embeddings: cfg.max_position_embeddings,
max_window_layers: 0,
num_attention_heads: cfg.num_attention_heads,
num_hidden_layers: cfg.num_hidden_layers,
num_key_value_heads: cfg.num_key_value_heads,
rms_norm_eps: cfg.rms_norm_eps,
rope_theta: cfg.rope_theta,
tie_word_embeddings: true,
torch_dtype: cfg.dtype.clone(),
use_cache: cfg.use_cache,
use_sliding_window: false,
vocab_size: cfg.vocab_size,
}
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct Qwen3VLVisionConfig {
pub deepstack_visual_indexes: Vec<usize>,
@@ -73,15 +101,3 @@ pub struct Qwen3VLConfig {
pub vision_end_token_id: usize,
pub vision_start_token_id: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct Qwen3VLGenerationConfig {
pub bos_token_id: usize,
pub pad_token_id: usize,
pub do_sample: bool,
pub eos_token_id: Vec<usize>,
pub top_p: f32,
pub top_k: usize,
pub temperature: f32,
pub repetition_penalty: f32,
}
+4 -7
View File
@@ -11,11 +11,8 @@ use crate::{
chat_template::ChatTemplate,
models::{
GenerateModel,
qwen3vl::{
config::{Qwen3VLConfig, Qwen3VLGenerationConfig},
model::Qwen3VLModel,
processor::Qwen3VLProcessor,
},
qwen3::config::Qwen3GenerationConfig,
qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor},
},
tokenizer::TokenizerModel,
utils::{
@@ -32,7 +29,7 @@ pub struct Qwen3VLGenerateModel<'a> {
device: Device,
eos_token_id1: u32,
eos_token_id2: u32,
generation_config: Qwen3VLGenerationConfig,
generation_config: Qwen3GenerationConfig,
model_name: String,
}
@@ -50,7 +47,7 @@ impl<'a> Qwen3VLGenerateModel<'a> {
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
let qwen3_vl = Qwen3VLModel::new(cfg, vb)?;
let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: Qwen3VLGenerationConfig =
let generation_config: Qwen3GenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
Ok(Self {
chat_template,
+9 -178
View File
@@ -7,12 +7,14 @@ use candle_nn::{
use crate::{
models::{
common::{GateUpDownMLP, TwoLinearMLP, eager_attention_forward, get_layer_norm},
qwen3vl::config::{Qwen3VLConfig, Qwen3VLTextConfig, Qwen3VLVisionConfig},
common::{TwoLinearMLP, eager_attention_forward, get_layer_norm},
qwen3::model::Qwen3DecoderLayer,
qwen3vl::config::{
Qwen3VLConfig, Qwen3VLTextConfig, Qwen3VLVisionConfig, qwen3vl_text_config2qwen3_config,
},
},
position_embed::rope::{
Qwen2_5VisionRotaryEmbedding, Qwen3VLTextRotaryEmbedding, apply_rotary_pos_emb,
apply_rotary_pos_emb_vision,
Qwen2_5VisionRotaryEmbedding, Qwen3VLTextRotaryEmbedding, apply_rotary_pos_emb_vision,
},
utils::tensor_utils::{
bitor_tensor, get_vision_next_indices, linspace, mask_index_add, masked_scatter_dim0,
@@ -519,181 +521,9 @@ impl Qwen3VLVisionModel {
}
}
pub struct Qwen3VLTextAttention {
q_proj: Linear,
k_proj: Linear,
v_proj: Linear,
o_proj: Linear,
q_norm: RmsNorm,
k_norm: RmsNorm,
num_attention_heads: usize,
num_key_value_heads: usize,
num_kv_groups: usize,
head_dim: usize,
scaling: f64,
kv_cache: Option<(Tensor, Tensor)>,
}
impl Qwen3VLTextAttention {
pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result<Self> {
let hidden_size = config.hidden_size;
let num_attention_heads = config.num_attention_heads;
let head_dim = config.head_dim;
let num_key_value_heads = config.num_key_value_heads;
let num_kv_groups = num_attention_heads / num_key_value_heads;
let scaling = 1f64 / f64::sqrt(head_dim as f64);
let (q_proj, k_proj, v_proj, o_proj) = if config.attention_bias {
let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
} else {
let q_proj =
linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj =
linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
};
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
let k_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?;
Ok(Self {
q_proj,
k_proj,
v_proj,
o_proj,
q_norm,
k_norm,
num_attention_heads,
num_key_value_heads,
num_kv_groups,
head_dim,
scaling,
kv_cache: None,
})
}
pub fn forward(
&mut self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
let (b_sz, q_len, _) = xs.dims3()?;
let query_states = self.q_proj.forward(xs)?.reshape((
b_sz,
q_len,
self.num_attention_heads,
self.head_dim,
))?;
let query_states = self.q_norm.forward(&query_states)?.transpose(1, 2)?;
let key_states = self.k_proj.forward(xs)?.reshape((
b_sz,
q_len,
self.num_key_value_heads,
self.head_dim,
))?;
let key_states = self.k_norm.forward(&key_states)?.transpose(1, 2)?;
let value_states = self.v_proj.forward(xs)?;
let value_states = value_states
.reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))?
.transpose(1, 2)?;
let (query_states, key_states) =
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?;
let (key_states, value_states) = match &self.kv_cache {
None => (key_states, value_states),
Some((prev_k, prev_v)) => {
let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
(key_states, value_states)
}
};
self.kv_cache = Some((key_states.clone(), value_states.clone()));
let attn_output = eager_attention_forward(
&query_states,
&key_states,
&value_states,
Some(self.num_kv_groups),
attention_mask,
self.scaling,
)?;
let attn_output =
attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?;
let attn_output = attn_output.apply(&self.o_proj)?;
Ok(attn_output)
}
pub fn clear_kv_cache(&mut self) {
self.kv_cache = None
}
}
pub struct Qwen3VLTextDecoderLayer {
self_attn: Qwen3VLTextAttention,
mlp: GateUpDownMLP,
input_layernorm: RmsNorm,
post_attention_layernorm: RmsNorm,
}
impl Qwen3VLTextDecoderLayer {
pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result<Self> {
let self_attn = Qwen3VLTextAttention::new(config.clone(), vb.pp("self_attn"))?;
let mlp = GateUpDownMLP::new(
vb.pp("mlp"),
config.hidden_size,
config.intermediate_size,
config.hidden_act,
false,
)?;
let input_layernorm = rms_norm(
config.hidden_size,
config.rms_norm_eps,
vb.pp("input_layernorm"),
)?;
let post_attention_layernorm = rms_norm(
config.hidden_size,
config.rms_norm_eps,
vb.pp("post_attention_layernorm"),
)?;
Ok(Self {
self_attn,
mlp,
input_layernorm,
post_attention_layernorm,
})
}
pub fn forward(
&mut self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?;
let xs = residual.add(&xs)?;
let residual = xs.clone();
let xs = self.post_attention_layernorm.forward(&xs)?;
let xs = self.mlp.forward(&xs)?;
let xs = residual.add(&xs)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
self.self_attn.clear_kv_cache();
}
}
pub struct Qwen3VLTextModel {
embed_tokens: Embedding,
layers: Vec<Qwen3VLTextDecoderLayer>,
layers: Vec<Qwen3DecoderLayer>,
norm: RmsNorm,
rotary_emb: Qwen3VLTextRotaryEmbedding,
mrope_section: Vec<usize>,
@@ -706,7 +536,8 @@ impl Qwen3VLTextModel {
let mut layers = vec![];
let vb_l = vb.pp("layers");
for layer_idx in 0..config.num_hidden_layers {
let layer = Qwen3VLTextDecoderLayer::new(config.clone(), vb_l.pp(layer_idx))?;
let qwen3_cfg = qwen3vl_text_config2qwen3_config(&config);
let layer = Qwen3DecoderLayer::new(&qwen3_cfg, vb_l.pp(layer_idx))?;
layers.push(layer)
}
let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?;
+18 -10
View File
@@ -147,6 +147,7 @@ impl VoxCPMGenerate {
}
None => self.generate_simple(target_text)?,
};
self.voxcpm.clear_kv_cache();
Ok(audio)
}
@@ -197,6 +198,7 @@ impl VoxCPMGenerate {
// retry_badcase,
retry_badcase_ratio_threshold,
)?;
self.voxcpm.clear_kv_cache();
Ok(audio)
}
@@ -223,19 +225,25 @@ impl GenerateModel for VoxCPMGenerate {
} else {
None
};
let audio = self.voxcpm.generate(
target_text,
prompt_text,
prompt_wav_path,
min_len,
max_len,
inference_timesteps,
cfg_value,
retry_badcase_ratio_threshold,
)?;
let audio = self
.voxcpm
.generate(
target_text,
prompt_text,
prompt_wav_path,
min_len,
max_len,
inference_timesteps,
cfg_value,
retry_badcase_ratio_threshold,
)
.inspect_err(|_| {
self.voxcpm.clear_kv_cache();
})?;
let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate as u32)?;
let base64_audio = BASE64_STANDARD.encode(wav_u8);
let response = build_audio_completion_response(&base64_audio, &self.model_name);
self.voxcpm.clear_kv_cache();
Ok(response)
}
#[allow(unused_variables)]
+3
View File
@@ -120,6 +120,9 @@ impl MiniCPMDecoderLayer {
None,
false,
None,
None,
None,
None,
)?;
let mlp = GateUpDownMLP::new(
vb.pp("mlp"),
+7 -1
View File
@@ -716,9 +716,15 @@ impl VoxCPMModel {
.permute((0, 3, 1, 2))?
.reshape((b, d, ()))?
.contiguous()?;
// self.base_lm.clear_kv_cache();
// self.residual_lm.clear_kv_cache();
self.clear_kv_cache();
Ok(feat_pred)
}
pub fn clear_kv_cache(&mut self) {
self.base_lm.clear_kv_cache();
self.residual_lm.clear_kv_cache();
Ok(feat_pred)
}
pub fn build_prompt_cache(