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();
}
}