add glm-asr-nano
This commit is contained in:
+201
-10
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,4 +1,4 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
pub mod processor;
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -205,6 +205,7 @@ impl Qwen3VLVisionBlock {
|
||||
vb.pp("mlp"),
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
config.hidden_size,
|
||||
config.hidden_act,
|
||||
true,
|
||||
"linear_fc1",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
Reference in New Issue
Block a user