add qwen3 and fun-asr-nano
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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)))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user