add glm-asr-nano

This commit is contained in:
jhqxxx
2026-01-07 21:46:01 +08:00
parent 981457347a
commit 6ef8facfde
22 changed files with 1918 additions and 117 deletions
+201 -10
View File
@@ -1,12 +1,16 @@
use anyhow::Result;
use candle_core::{D, Tensor};
use candle_nn::{
Activation, BatchNorm, BatchNormConfig, Conv2d, Conv2dConfig, LayerNorm, LayerNormConfig,
Linear, Module, RmsNorm, VarBuilder, batch_norm, conv2d, conv2d_no_bias, layer_norm, linear,
linear_no_bias, rms_norm,
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, Embedding,
LayerNorm, LayerNormConfig, Linear, Module, RmsNorm, VarBuilder, batch_norm, conv1d,
conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear, linear_no_bias,
rms_norm,
};
use crate::{position_embed::rope::apply_rotary_pos_emb, utils::tensor_utils::repeat_kv};
use crate::{
position_embed::rope::{RoPE, apply_rotary_pos_emb},
utils::tensor_utils::{prepare_causal_attention_mask, repeat_kv},
};
#[derive(Debug, Clone)]
pub struct GateUpDownMLP {
@@ -63,8 +67,11 @@ pub struct TwoLinearMLP {
impl TwoLinearMLP {
pub fn new(
vb: VarBuilder,
embedding_dim: usize,
mlp_dim: usize,
// embedding_dim: usize,
// mlp_dim: usize,
in_dim: usize,
middle_dim: usize,
out_dim: usize,
act: Activation,
bias: bool,
linear1_pp_name: &str,
@@ -72,13 +79,13 @@ impl TwoLinearMLP {
) -> Result<Self> {
let (linear1, linear2) = if bias {
(
linear(embedding_dim, mlp_dim, vb.pp(linear1_pp_name))?,
linear(mlp_dim, embedding_dim, vb.pp(linear2_pp_name))?,
linear(in_dim, middle_dim, vb.pp(linear1_pp_name))?,
linear(middle_dim, out_dim, vb.pp(linear2_pp_name))?,
)
} else {
(
linear_no_bias(embedding_dim, mlp_dim, vb.pp(linear1_pp_name))?,
linear_no_bias(mlp_dim, embedding_dim, vb.pp(linear2_pp_name))?,
linear_no_bias(in_dim, middle_dim, vb.pp(linear1_pp_name))?,
linear_no_bias(middle_dim, out_dim, vb.pp(linear2_pp_name))?,
)
};
Ok(Self {
@@ -304,6 +311,7 @@ impl NaiveAttnTwoLinearMLPBlock {
vb.pp(mlp_pp_name),
hidden_size,
intermediate_size,
hidden_size,
hidden_act,
mlp_bias,
linear1_pp_name,
@@ -503,6 +511,32 @@ pub fn get_conv2d(
Ok(conv2d)
}
pub fn get_conv1d(
vb: VarBuilder,
in_c: usize,
out_c: usize,
kernel_size: usize,
padding: usize,
stride: usize,
dilation: usize,
groups: usize,
bias: bool,
) -> Result<Conv1d> {
let cfg = Conv1dConfig {
padding,
stride,
dilation,
groups,
cudnn_fwd_algo: None,
};
let conv1d = if bias {
conv1d(in_c, out_c, kernel_size, cfg, vb)?
} else {
conv1d_no_bias(in_c, out_c, kernel_size, cfg, vb)?
};
Ok(conv1d)
}
pub fn get_layer_norm(vb: VarBuilder, eps: f64, dim: usize) -> Result<LayerNorm> {
let ln_config = LayerNormConfig {
eps,
@@ -620,3 +654,160 @@ pub fn deform_conv2d_kernel(
}
Ok(out)
}
pub struct LlamaModel {
pub embed_tokens: Embedding,
layers: Vec<NaiveAttnGateUpDownMLPBlock>,
norm: RmsNorm,
rotary_emb: RoPE,
}
impl LlamaModel {
pub fn new(
vb: VarBuilder,
vocab_size: usize,
hidden_size: usize,
num_hidden_layers: usize,
num_attention_heads: usize,
num_key_value_heads: Option<usize>,
head_dim: Option<usize>,
attn_bias: bool,
attn_pp_name: &str,
o_proj_pp_name: Option<&str>,
intermediate_size: usize,
hidden_act: Activation,
mlp_bias: bool,
mlp_pp_name: &str,
norm_eps: f64,
input_norm_pp_name: &str,
post_norm_pp_name: &str,
rope_theta_base: f32,
) -> Result<Self> {
let embed_tokens = embedding(vocab_size, hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![];
let vb_layers = vb.pp("layers");
for i in 0..num_hidden_layers {
let layers_i = NaiveAttnGateUpDownMLPBlock::new(
vb_layers.pp(i),
hidden_size,
num_attention_heads,
num_key_value_heads,
head_dim,
attn_bias,
attn_pp_name,
o_proj_pp_name,
intermediate_size,
hidden_act,
mlp_bias,
mlp_pp_name,
norm_eps,
input_norm_pp_name,
post_norm_pp_name,
)?;
layers.push(layers_i);
}
let norm = rms_norm(hidden_size, norm_eps, vb.pp("norm"))?;
let head_dim = head_dim.unwrap_or(hidden_size / num_attention_heads);
let rotary_emb = RoPE::new(head_dim, rope_theta_base, vb.device())?;
Ok(Self {
embed_tokens,
layers,
norm,
rotary_emb,
})
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let (cos, sin) = self
.rotary_emb
.forward(seqlen_offset, seq_len, inputs_embeds.device())?;
let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
b_size,
seq_len,
0,
xs.device(),
)?)
}
};
for layer in self.layers.iter_mut() {
xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?;
}
let xs = xs.apply(&self.norm)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.layers.iter_mut() {
layer.clear_kv_cache()
}
}
}
pub struct LlamaForCausalLM {
pub model: LlamaModel,
lm_head: Linear,
}
impl LlamaForCausalLM {
pub fn new(
vb: VarBuilder,
vocab_size: usize,
hidden_size: usize,
num_hidden_layers: usize,
num_attention_heads: usize,
num_key_value_heads: Option<usize>,
head_dim: Option<usize>,
attn_bias: bool,
attn_pp_name: &str,
o_proj_pp_name: Option<&str>,
intermediate_size: usize,
hidden_act: Activation,
mlp_bias: bool,
mlp_pp_name: &str,
norm_eps: f64,
input_norm_pp_name: &str,
post_norm_pp_name: &str,
rope_theta_base: f32,
) -> Result<Self> {
let model = LlamaModel::new(
vb.pp("model"),
vocab_size,
hidden_size,
num_hidden_layers,
num_attention_heads,
num_key_value_heads,
head_dim,
attn_bias,
attn_pp_name,
o_proj_pp_name,
intermediate_size,
hidden_act,
mlp_bias,
mlp_pp_name,
norm_eps,
input_norm_pp_name,
post_norm_pp_name,
rope_theta_base,
)?;
let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?;
Ok(Self { model, lm_head })
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let outputs = self.model.forward(inputs_embeds, seqlen_offset)?;
let seq_len = outputs.dim(1)?;
let hidden_state = outputs.narrow(1, seq_len - 1, 1)?;
let logits = self.lm_head.forward(&hidden_state)?;
Ok(logits)
}
pub fn clear_kv_cache(&mut self) {
self.model.clear_kv_cache();
}
}
+1 -1
View File
@@ -271,7 +271,7 @@ impl Block {
)?;
let norm2 = get_layer_norm(vb.pp("norm2"), eps, dim)?;
let mlp_dim = (dim as f32 * mlp_ratio) as usize;
let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, act, true, "lin1", "lin2")?;
let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, dim, act, true, "lin1", "lin2")?;
Ok(Self {
norm1,
attn,
+66 -2
View File
@@ -1,4 +1,5 @@
use serde::{Deserialize};
use candle_nn::Activation;
use serde::Deserialize;
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrNanoProcessorConfig {
@@ -21,4 +22,67 @@ pub struct FeatureExtractor {
pub padding_value: f32,
pub return_attention_mask: bool,
pub sampling_rate: usize,
}
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrNanoConfig {
pub audio_config: GlmAsrAudioConfig,
pub audio_token_id: u32,
pub dtype: String,
pub hidden_size: usize,
pub projector_hidden_act: Activation,
pub text_config: GlmAsrTextConfig,
pub vocab_size: usize,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrAudioConfig {
pub attention_dropout: f64,
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 num_attention_heads: usize,
pub num_hidden_layers: usize,
pub num_key_value_heads: usize,
pub num_mel_bins: usize,
pub partial_rotary_factor: f64,
pub rope_parameters: GlmAsrRopeParameters,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrRopeParameters {
pub partial_rotary_factor: f64,
pub rope_theta: f32,
pub rope_type: String,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrTextConfig {
pub attention_bias: bool,
pub attention_dropout: f64,
pub eos_token_id: Vec<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 mlp_bias: bool,
pub num_attention_heads: usize,
pub num_hidden_layers: usize,
pub num_key_value_heads: usize,
pub pretraining_tp: usize,
pub rms_norm_eps: f64,
pub rope_parameters: GlmAsrTextRopeParameters,
pub use_cache: bool,
pub vocab_size: usize,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrTextRopeParameters {
pub rope_theta: f32,
pub rope_type: String,
}
+149 -18
View File
@@ -1,24 +1,37 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{DType, Device};
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::{
chat_template::ChatTemplate,
models::glm_asr_nano::{config::GlmAsrNanoProcessorConfig, processor::GlmAsrNanoProcessor},
models::{
GenerateModel,
glm_asr_nano::{
config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor,
},
},
tokenizer::TokenizerModel,
utils::{get_device, get_dtype},
utils::{
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
get_dtype, get_logit_processor,
},
};
pub struct GlmAsrNanoGenerateModel<'a> {
chat_template: ChatTemplate<'a>,
tokenizer: TokenizerModel,
processor: GlmAsrNanoProcessor,
// glm_asr_nano: GlmAsrNanoModel,
glm_asr_nano: GlmAsrNanoModel,
device: Device,
// eos_token_id1: u32,
// eos_token_id2: u32,
// eos_token_id3: u32,
// generation_config: GlmAsrNanoGenerationConfig,
dtype: DType,
eos_token_id1: u32,
eos_token_id2: u32,
eos_token_id3: u32,
model_name: String,
}
@@ -28,23 +41,141 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
let tokenizer = TokenizerModel::init(path)?;
let device = get_device(device);
let processor = GlmAsrNanoProcessor::new(path, &device, DType::F32)?;
// let cfg_dtype = cfg.dtype.as_str();
// let dtype = get_dtype(dtype, cfg_dtype);
let config_path = path.to_string() + "/config.json";
let cfg: GlmAsrNanoConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
let cfg_dtype = cfg.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 glm_asr_nano = GlmAsrNanoModel::new(vb, cfg)?;
Ok(Self {
chat_template,
tokenizer,
processor,
glm_asr_nano,
device,
// eos_token_id1,
// eos_token_id2,
// eos_token_id3,
dtype,
eos_token_id1: 59246,
eos_token_id2: 59253,
eos_token_id3: 59255,
model_name: "glm-asr-nano".to_string(),
})
}
}
pub fn generate(&self, mes: ChatCompletionParameters) -> Result<()> {
impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
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 audio = self.processor.process_info(&mes)?;
Ok(())
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)?;
let mut input_features = Some(input_features.to_dtype(self.dtype)?);
let mut audio_token_lengths = Some(audio_token_lengths);
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let mut generate: Vec<u32> = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024);
for _ in 0..sample_len {
let logits = self.glm_asr_nano.forward(
input_features.as_ref(),
audio_token_lengths.as_ref(),
&input_ids,
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
|| next_token == self.eos_token_id3
{
break;
}
seqlen_offset += seq_len;
seq_len = 1;
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
input_features = None;
audio_token_lengths = None;
}
let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?;
self.glm_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 seed = match mes.seed {
None => 34562u64,
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 (input_features, audio_token_lengths, replace_text) =
self.processor.process_info(&mes, &render_text)?;
let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
let sample_len = mes.max_tokens.unwrap_or(1024);
let stream = stream! {
let mut error_tokens = Vec::new();
let mut input_features = Some(input_features.to_dtype(self.dtype)?);
let mut audio_token_lengths = Some(audio_token_lengths);
let mut input_ids = input_ids;
for _ in 0..sample_len {
let logits =
self.glm_asr_nano
.forward(input_features.as_ref(), audio_token_lengths.as_ref(), &input_ids, 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)?;
input_features = None;
audio_token_lengths = 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 || next_token == self.eos_token_id3{
break;
}
seqlen_offset += seq_len;
seq_len = 1;
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
input_features = None;
audio_token_lengths = None;
}
self.glm_asr_nano.clear_kv_cache();
};
Ok(Box::new(Box::pin(stream)))
}
}
+1 -1
View File
@@ -1,4 +1,4 @@
pub mod config;
pub mod generate;
pub mod model;
pub mod processor;
pub mod processor;
+320
View File
@@ -0,0 +1,320 @@
use anyhow::Result;
use candle_core::{IndexOp, Tensor};
use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, linear_no_bias};
use crate::{
models::{
common::{
LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
},
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
},
position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb},
utils::tensor_utils::{get_equal_mask, masked_scatter_dim0},
};
#[derive(Debug, Clone)]
// pub struct AttentionNobias {
pub struct GlmAsrAttention {
q_proj: Linear,
k_proj: Linear,
v_proj: Linear,
o_proj: Linear,
num_heads: usize,
num_kv_heads: usize,
num_kv_groups: usize,
head_dim: usize,
middle_size: usize,
}
impl GlmAsrAttention {
pub fn new(
vb: VarBuilder,
hidden_size: usize,
num_attention_heads: usize,
num_key_value_heads: usize,
head_dim: Option<usize>,
) -> Result<Self> {
let num_kv_groups = num_attention_heads / num_key_value_heads;
let head_dim = match head_dim {
None => hidden_size / num_attention_heads,
Some(dim) => dim,
};
let q_proj = linear(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(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"))?;
Ok(Self {
q_proj,
k_proj,
v_proj,
o_proj,
num_heads: num_attention_heads,
num_kv_heads: num_key_value_heads,
num_kv_groups,
head_dim,
middle_size: num_attention_heads * head_dim,
})
}
pub fn forward(
&self,
xs: &Tensor,
cos: Option<&Tensor>,
sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
tof32: bool,
) -> Result<Tensor> {
let (b_sz, q_len, _) = xs.dims3()?;
let query_states = self.q_proj.forward(xs)?;
let key_states = self.k_proj.forward(xs)?;
let value_states = self.v_proj.forward(xs)?;
let query_states = query_states
.reshape((b_sz, q_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let key_states = key_states
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
.transpose(1, 2)?;
let value_states = value_states
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
.transpose(1, 2)?;
let (query_states, key_states) = if let Some(cos) = cos
&& let Some(sin) = sin
{
glm_asr_apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?
} else {
(query_states, key_states)
};
let scale = 1f64 / f64::sqrt(self.head_dim as f64);
let attn_output = eager_attention_forward(
&query_states,
&key_states,
&value_states,
Some(self.num_kv_groups),
attention_mask,
scale,
)?;
let attn_output = attn_output.reshape((b_sz, q_len, self.middle_size))?;
let attn_output = attn_output.apply(&self.o_proj)?;
Ok(attn_output)
}
}
pub struct GlmAsrEncoderLayer {
self_attn: GlmAsrAttention,
mlp: TwoLinearMLP,
input_layernorm: LayerNorm,
post_attention_layernorm: LayerNorm,
}
impl GlmAsrEncoderLayer {
pub fn new(vb: VarBuilder, audio_cfg: &GlmAsrAudioConfig) -> Result<Self> {
let self_attn = GlmAsrAttention::new(
vb.pp("self_attn"),
audio_cfg.hidden_size,
audio_cfg.num_attention_heads,
audio_cfg.num_key_value_heads,
Some(audio_cfg.head_dim),
)?;
let mlp = TwoLinearMLP::new(
vb.pp("mlp"),
audio_cfg.hidden_size,
audio_cfg.intermediate_size,
audio_cfg.hidden_size,
audio_cfg.hidden_act,
true,
"fc1",
"fc2",
)?;
let input_layernorm =
get_layer_norm(vb.pp("input_layernorm"), 1e-5, audio_cfg.hidden_size)?;
let post_attention_layernorm = get_layer_norm(
vb.pp("post_attention_layernorm"),
1e-5,
audio_cfg.hidden_size,
)?;
Ok(Self {
self_attn,
mlp,
input_layernorm,
post_attention_layernorm,
})
}
pub fn forward(
&self,
xs: &Tensor,
cos: Option<&Tensor>,
sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
tof32: bool,
) -> Result<Tensor> {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
let xs = self
.self_attn
.forward(&xs, cos, sin, attention_mask, tof32)?;
let residual = residual.add(&xs)?;
let xs = self.post_attention_layernorm.forward(&residual)?;
let xs = self.mlp.forward(&xs)?;
let xs = residual.add(&xs)?;
Ok(xs)
}
}
pub struct GlmAsrEncoder {
conv1: Conv1d,
conv2: Conv1d,
layers: Vec<GlmAsrEncoderLayer>,
norm: LayerNorm,
rotary_emb: RoPE,
}
impl GlmAsrEncoder {
pub fn new(vb: VarBuilder, audio_cfg: &GlmAsrAudioConfig) -> Result<Self> {
let conv1 = get_conv1d(
vb.pp("conv1"),
audio_cfg.num_mel_bins,
audio_cfg.hidden_size,
3,
1,
1,
1,
1,
true,
)?;
let conv2 = get_conv1d(
vb.pp("conv2"),
audio_cfg.hidden_size,
audio_cfg.hidden_size,
3,
1,
2,
1,
1,
true,
)?;
let mut layers = vec![];
let vb_layers = vb.pp("layers");
for i in 0..audio_cfg.num_hidden_layers {
let layer_i = GlmAsrEncoderLayer::new(vb_layers.pp(i), audio_cfg)?;
layers.push(layer_i);
}
let norm = get_layer_norm(vb.pp("norm"), 1e-5, audio_cfg.hidden_size)?;
let dim = (audio_cfg.head_dim as f64 * audio_cfg.partial_rotary_factor) as usize;
let rotary_emb = RoPE::new(dim, audio_cfg.rope_parameters.rope_theta, vb.device())?;
Ok(Self {
conv1,
conv2,
layers,
norm,
rotary_emb,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let xs = self.conv1.forward(xs)?.gelu()?;
let xs = self.conv2.forward(&xs)?.gelu()?;
let mut xs = xs.transpose(1, 2)?;
let (_, seq_len, _) = xs.dims3()?;
let (cos, sin) = self.rotary_emb.forward(0, seq_len, xs.device())?;
for encoder_layer in &self.layers {
xs = encoder_layer.forward(&xs, Some(&cos), Some(&sin), None, false)?;
}
let xs = self.norm.forward(&xs)?;
Ok(xs)
}
}
pub struct GlmAsrNanoModel {
config: GlmAsrNanoConfig,
audio_tower: GlmAsrEncoder,
multi_modal_projector: TwoLinearMLP,
language_model: LlamaForCausalLM,
}
impl GlmAsrNanoModel {
pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig) -> Result<Self> {
let audio_tower = GlmAsrEncoder::new(vb.pp("audio_tower"), &config.audio_config)?;
let multi_modal_projector = TwoLinearMLP::new(
vb.pp("multi_modal_projector"),
config.audio_config.intermediate_size,
config.text_config.hidden_size * 2,
config.text_config.hidden_size,
config.projector_hidden_act,
true,
"linear_1",
"linear_2",
)?;
let language_model = LlamaForCausalLM::new(
vb.pp("language_model"),
config.text_config.vocab_size,
config.text_config.hidden_size,
config.text_config.num_hidden_layers,
config.text_config.num_attention_heads,
Some(config.text_config.num_key_value_heads),
Some(config.text_config.head_dim),
config.text_config.attention_bias,
"self_attn",
Some("o_proj"),
config.text_config.intermediate_size,
config.text_config.hidden_act,
config.text_config.mlp_bias,
"mlp",
config.text_config.rms_norm_eps,
"input_layernorm",
"post_attention_layernorm",
config.text_config.rope_parameters.rope_theta,
)?;
Ok(Self {
config,
audio_tower,
multi_modal_projector,
language_model,
})
}
pub fn get_audio_features(
&self,
input_features: &Tensor,
audio_token_lengths: &[u32],
) -> Result<Tensor> {
let audio_hidden_states = self.audio_tower.forward(input_features)?;
let bs = audio_hidden_states.dim(0)?;
let audio_hidden_states =
audio_hidden_states.reshape((bs, (), self.config.audio_config.intermediate_size))?;
let audio_embeds = self.multi_modal_projector.forward(&audio_hidden_states)?;
let mut valid_audios = vec![];
for (i, &len) in audio_token_lengths.iter().enumerate() {
let len = len as usize;
let audio_i = audio_embeds.i((i, 0..len, ..))?;
valid_audios.push(audio_i);
}
let audio_embeds = Tensor::cat(&valid_audios, 0)?;
Ok(audio_embeds)
}
pub fn forward(
&mut self,
input_features: Option<&Tensor>,
audio_token_lengths: Option<&Vec<u32>>,
input_ids: &Tensor,
seqlen_offset: usize,
) -> Result<Tensor> {
let mut inputs_embeds = self.language_model.model.embed_tokens.forward(input_ids)?;
if let Some(input_features) = input_features
&& let Some(audio_token_len) = audio_token_lengths
{
let audio_token_mask = get_equal_mask(input_ids, self.config.audio_token_id)?;
let audio_embeds = self.get_audio_features(input_features, audio_token_len)?;
inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embeds, &audio_token_mask)?;
}
let logits = self.language_model.forward(&inputs_embeds, seqlen_offset)?;
Ok(logits)
}
pub fn clear_kv_cache(&mut self) {
self.language_model.clear_kv_cache();
}
}
+187 -46
View File
@@ -1,30 +1,30 @@
use std::f32;
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{DType, Device, IndexOp, Tensor};
use candle_core::{D, DType, Device, IndexOp, Tensor};
use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
use crate::{
models::glm_asr_nano::config::GlmAsrNanoProcessorConfig,
tokenizer::TokenizerModel,
utils::{audio_utils::extract_audios, extract_user_text},
utils::{
audio_utils::{create_hann_window, extract_audios, mel_filter_bank, stft_audio},
tensor_utils::{pad_reflect_last_dim, split_tensor},
},
};
pub struct WhisperFeatureExtractor {
feature_size: usize,
sampling_rate: usize,
padding_value: f32,
hop_length: usize,
chunk_length: usize,
n_fft: usize,
dither: f32,
}
pub struct GlmAsrNanoProcessor {
sampling_rate: usize,
chunk_length: usize,
n_samples: usize,
n_fft: usize,
window: Tensor,
mel_filters: Tensor,
hop_length: usize,
audio_token: String,
audio_token_id: u32,
// audio_token_id: u32,
max_audio_len: usize,
default_transcription_prompt: String,
// default_transcription_prompt: String,
device: Device,
}
@@ -43,52 +43,193 @@ impl GlmAsrNanoProcessor {
let processor_cfg: GlmAsrNanoProcessorConfig =
serde_json::from_slice(&std::fs::read(processor_config_path)?)?;
let audio_token = processor_cfg.audio_token.clone();
let audio_token_id = 59260u32;
// let audio_token_id = 59260u32;
let max_audio_len = processor_cfg.max_audio_len;
let default_transcription_prompt = processor_cfg.default_transcription_prompt.clone();
// let default_transcription_prompt = processor_cfg.default_transcription_prompt.clone();
let sampling_rate = processor_cfg.feature_extractor.sampling_rate;
let chunk_length = processor_cfg.feature_extractor.chunk_length;
let n_samples = processor_cfg.feature_extractor.n_samples;
let n_fft = processor_cfg.feature_extractor.n_fft;
let hop_length = processor_cfg.feature_extractor.hop_length;
let window = create_hann_window(n_fft, dtype, device)?;
let window = window.unsqueeze(0)?.unsqueeze(0)?;
let mel_filters = mel_filter_bank(
1 + n_fft / 2,
processor_cfg.feature_extractor.feature_size,
0.0,
8000.0,
sampling_rate as f32,
Some("slaney"),
crate::utils::audio_utils::MelScale::Slaney,
false,
device,
)?
.t()?;
Ok(Self {
sampling_rate,
chunk_length,
n_samples,
n_fft,
window,
mel_filters,
hop_length,
audio_token,
audio_token_id,
// audio_token_id,
max_audio_len,
default_transcription_prompt,
// default_transcription_prompt,
device: device.clone(),
})
}
// pub fn process_audio(&self, audios: Vec<Tensor>) -> Result<Tensor> {
// let window_size = self.sampling_rate * self.chunk_length;
// let max_windows = self.max_audio_len / self.chunk_length;
// let mut per_sample_windows = vec![];
// let mut flat_chunks = vec![];
// for audio_el in audios {
// let n_samples = audio_el.dim(0)?;
// let n_win = ((n_samples + window_size - 1) / window_size).max(1);
// let n_win = if n_win > max_windows {
// max_windows
// } else {
// n_win
// };
// per_sample_windows.push(n_win);
// let time_cap = (n_win * window_size).min(n_samples);
// for i in 0..n_win {
// let start = i * window_size;
// let end = ((i + 1) * window_size).min(time_cap);
// flat_chunks.push(audio_el.i(start..end)?);
// }
// }
// }
/// 提取音频帧
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 n_frames = (samples - self.n_fft) / self.hop_length + 1;
// (bs, n_frames, n_fft)
let frames = self.extract_frames(&waveform, n_frames)?;
// 应用汉明窗口
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 = 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)?;
let ln_spec = mel_spec.log()?;
let log10_spec = ln_spec.broadcast_div(&Tensor::new(f32::ln(10.0), mel_spec.device())?)?;
let max_val = log10_spec.max_all()?.affine(1.0, -8.0)?;
let log10_spec = log10_spec.broadcast_maximum(&max_val)?;
let log_spec = log10_spec.affine(1.0, 4.0)?.affine(1.0 / 4.0, 0.0)?;
Ok(log_spec)
}
pub fn feature_extractor(&self, raw_speech: Vec<Tensor>) -> Result<(Tensor, Tensor)> {
let mut pad_audio = vec![];
let mut input_features_mask = vec![];
for audio in raw_speech {
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)?;
// (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]);
input_features_mask.push(mask);
}
let input_features = Tensor::cat(&pad_audio, 0)?;
let input_features_mask = Tensor::new(input_features_mask, input_features.device())?;
let input_features = self.extract_fbank_features(&input_features)?;
let (_, audio_len) = input_features_mask.dims2()?;
let mask_idx: Vec<u32> = (0..audio_len)
.step_by(self.hop_length)
.map(|i| i as u32)
.collect();
let mask_idx = Tensor::new(mask_idx, &self.device)?;
let input_features_mask = input_features_mask.index_select(&mask_idx, D::Minus1)?;
Ok((input_features, input_features_mask))
}
pub fn process_audio(&self, audios: Vec<Tensor>) -> Result<(Tensor, Tensor, Vec<usize>)> {
let window_size = self.sampling_rate * self.chunk_length;
let max_windows = self.max_audio_len / self.chunk_length;
let mut per_sample_windows = vec![];
let mut flat_chunks = vec![];
for audio_el in audios {
let audio_el = if audio_el.rank() == 2 {
audio_el.squeeze(0)?
} else {
audio_el
};
let n_samples = audio_el.dim(0)?;
let n_win = ((n_samples + window_size - 1) / window_size).max(1);
let n_win = if n_win > max_windows {
max_windows
} else {
n_win
};
per_sample_windows.push(n_win);
let time_cap = (n_win * window_size).min(n_samples);
for i in 0..n_win {
let start = i * window_size;
let end = ((i + 1) * window_size).min(time_cap);
flat_chunks.push(audio_el.i(start..end)?);
}
}
let (input_features, input_features_mask) = self.feature_extractor(flat_chunks)?;
Ok((input_features, input_features_mask, per_sample_windows))
}
pub fn get_audio_token_length(&self, audio_lens: Vec<u32>) -> Result<Vec<u32>> {
let merge_factor = 4;
let audio_lens = audio_lens
.iter()
.map(|i| (i + 2 - 3) + 1) // (pad=1, ks=3, stride=1)
.collect::<Vec<u32>>()
.iter()
.map(|i| (i + 2 - 3) / 2 + 1) // (pad=1, ks=3, stride=2)
.collect::<Vec<u32>>();
let num_tokens = audio_lens
.iter()
.map(|i| (i - merge_factor) / merge_factor + 1)
.collect::<Vec<u32>>();
Ok(num_tokens)
}
pub fn process_info(
&self,
mes: &ChatCompletionParameters
) -> Result<Tensor> {
mes: &ChatCompletionParameters,
render_text: &str,
) -> Result<(Tensor, Vec<u32>, String)> {
let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?;
println!("audio: {}", audio_tensors[0]);
// let audio = self.process_audio(audio_tensors)?;
Ok(audio_tensors[0].clone())
let (input_features, input_features_mask, per_sample_windows) =
self.process_audio(audio_tensors)?;
let audio_lengths = input_features_mask.sum(D::Minus1)?;
let audio_vec = split_tensor(&audio_lengths, &per_sample_windows, 0)?;
let audio_vec: Vec<u32> = audio_vec
.iter()
.map(|t| t.sum_all().unwrap().to_scalar::<u32>().unwrap())
.collect();
let audio_token_lengths = self.get_audio_token_length(audio_vec)?;
let mut text = render_text.to_string();
for audio_len in audio_token_lengths.clone() {
let replace = "<|placeholder|>".repeat(audio_len as usize);
text = text.replacen(&self.audio_token, &replace, 1);
}
text = text.replace("<|placeholder|>", &self.audio_token);
Ok((input_features, audio_token_lengths, text))
}
}
+11 -1
View File
@@ -1,5 +1,6 @@
pub mod common;
pub mod deepseek_ocr;
pub mod glm_asr_nano;
pub mod hunyuan_ocr;
pub mod minicpm4;
pub mod paddleocr_vl;
@@ -7,7 +8,6 @@ pub mod qwen2_5vl;
pub mod qwen3vl;
pub mod rmbg2_0;
pub mod voxcpm;
pub mod glm_asr_nano;
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
@@ -17,6 +17,7 @@ use rocket::futures::Stream;
use crate::models::{
deepseek_ocr::generate::DeepseekOCRGenerateModel,
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,
@@ -51,6 +52,8 @@ pub enum WhichModel {
VoxCPM,
#[value(name = "voxcpm1.5")]
VoxCPM1_5,
#[value(name = "glm-asr-nano-2512")]
GlmASRNano2512,
}
pub trait GenerateModel {
@@ -77,6 +80,7 @@ pub enum ModelInstance<'a> {
PaddleOCRVL(Box<PaddleOCRVLGenerateModel<'a>>),
RMBG2_0(Box<RMBG2_0Model>),
VoxCPM(Box<VoxCPMGenerate>),
GlmASRNano(GlmAsrNanoGenerateModel<'a>),
}
impl<'a> GenerateModel for ModelInstance<'a> {
@@ -90,6 +94,7 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::PaddleOCRVL(model) => model.generate(mes),
ModelInstance::RMBG2_0(model) => model.generate(mes),
ModelInstance::VoxCPM(model) => model.generate(mes),
ModelInstance::GlmASRNano(model) => model.generate(mes),
}
}
@@ -113,6 +118,7 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::PaddleOCRVL(model) => model.generate_stream(mes),
ModelInstance::RMBG2_0(model) => model.generate_stream(mes),
ModelInstance::VoxCPM(model) => model.generate_stream(mes),
ModelInstance::GlmASRNano(model) => model.generate_stream(mes),
}
}
}
@@ -171,6 +177,10 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_
let model = VoxCPMGenerate::init(path, None, None)?;
ModelInstance::VoxCPM(Box::new(model))
}
WhichModel::GlmASRNano2512 => {
let model = GlmAsrNanoGenerateModel::init(path, None, None)?;
ModelInstance::GlmASRNano(model)
}
};
Ok(model)
}
+1
View File
@@ -205,6 +205,7 @@ impl Qwen3VLVisionBlock {
vb.pp("mlp"),
config.hidden_size,
config.intermediate_size,
config.hidden_size,
config.hidden_act,
true,
"linear_fc1",
+1 -1
View File
@@ -255,7 +255,7 @@ impl SwinTransformerBlock {
)?;
let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, dim)?;
let mlp_dim = (dim as f32 * mlp_ratio) as usize;
let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, act, true, "fc1", "fc2")?;
let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, dim, act, true, "fc1", "fc2")?;
Ok(Self {
norm1,
attn,
+2 -5
View File
@@ -728,11 +728,8 @@ impl VoxCPMModel {
) -> Result<HashMap<String, Tensor>> {
let text_token = self.tokenizer.encode(prompt_text)?;
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
let mut audio = load_audio_with_resample(
&prompt_wav_path,
&self.device,
Some(self.sample_rate),
)?;
let mut audio =
load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;