diff --git a/Cargo.lock b/Cargo.lock index 65de170..87ba2e6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -29,12 +29,14 @@ dependencies = [ "candle-transformers", "chrono", "ffmpeg-next", + "hound", "image", "minijinja", "num", "openai_dive", "reqwest", "rocket", + "rubato", "serde", "serde_json", "tokenizers", @@ -1489,6 +1491,12 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" +[[package]] +name = "hound" +version = "3.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "62adaabb884c94955b19907d60019f4e145d091c75345379e70d1ee696f7854f" + [[package]] name = "http" version = "0.2.12" @@ -2618,6 +2626,15 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "primal-check" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc0d895b311e3af9902528fbb8f928688abbd95872819320517cc24ca6b2bd08" +dependencies = [ + "num-integer", +] + [[package]] name = "proc-macro-crate" version = "3.4.0" @@ -2901,6 +2918,15 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "realfft" +version = "3.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f821338fddb99d089116342c46e9f1fbf3828dba077674613e734e01d6ea8677" +dependencies = [ + "rustfft", +] + [[package]] name = "reborrow" version = "0.5.5" @@ -3127,6 +3153,18 @@ dependencies = [ "uncased", ] +[[package]] +name = "rubato" +version = "0.16.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5258099699851cfd0082aeb645feb9c084d9a5e1f1b8d5372086b989fc5e56a1" +dependencies = [ + "num-complex", + "num-integer", + "num-traits", + "realfft", +] + [[package]] name = "rustc-demangle" version = "0.1.26" @@ -3139,6 +3177,20 @@ version = "2.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "357703d41365b4b27c590e3ed91eabb1b663f07c4c084095e60cbed4362dff0d" +[[package]] +name = "rustfft" +version = "6.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "21db5f9893e91f41798c88680037dba611ca6674703c1a18601b01a72c8adb89" +dependencies = [ + "num-complex", + "num-integer", + "num-traits", + "primal-check", + "strength_reduce", + "transpose", +] + [[package]] name = "rustix" version = "1.1.2" @@ -3471,6 +3523,12 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f" +[[package]] +name = "strength_reduce" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe895eb47f22e2ddd4dabc02bce419d2e643c8e3b585c78158b349195bc24d82" + [[package]] name = "strsim" version = "0.11.1" @@ -3984,6 +4042,16 @@ dependencies = [ "tracing-log", ] +[[package]] +name = "transpose" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ad61aed86bc3faea4300c7aee358b4c6d0c8d6ccc36524c96e4c92ccf26e77e" +dependencies = [ + "num-integer", + "strength_reduce", +] + [[package]] name = "try-lock" version = "0.2.5" diff --git a/Cargo.toml b/Cargo.toml index c4c6dd4..05727ab 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -26,6 +26,8 @@ uuid = { version = "1.18.1", features = ["v4"]} chrono = "0.4.42" rocket = "0.5.1" tokio = "1.47.1" +hound = "3.5.1" +rubato = "0.16.2" [features] flash-attn=["candle-flash-attn"] diff --git a/assets/audio/example.wav b/assets/audio/example.wav new file mode 100644 index 0000000..dcc2b46 Binary files /dev/null and b/assets/audio/example.wav differ diff --git a/assets/audio/voice_06.wav b/assets/audio/voice_06.wav new file mode 100644 index 0000000..7444036 Binary files /dev/null and b/assets/audio/voice_06.wav differ diff --git a/src/lib.rs b/src/lib.rs index 3f7d85d..e86488f 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -7,27 +7,27 @@ pub mod position_embed; pub mod tokenizer; pub mod utils; -pub enum ModelType { - Qwen2_5VL, - MiniCPM4, -} +// pub enum ModelType { +// Qwen2_5VL, +// MiniCPM4, +// } -impl ModelType { - pub fn init( - model_type: ModelType, - model_path: &str, - device: Option<&Device>, - dtype: Option, - ) -> Result> { - match model_type { - ModelType::Qwen2_5VL => { - let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?; - Ok(Box::new(model) as Box) - }, - ModelType::MiniCPM4 => { - let model = MiniCPMGenerateModel::init(model_path, device, dtype)?; - Ok(Box::new(model)as Box) - } - } - } -} +// impl ModelType { +// pub fn init( +// model_type: ModelType, +// model_path: &str, +// device: Option<&Device>, +// dtype: Option, +// ) -> Result> { +// match model_type { +// ModelType::Qwen2_5VL => { +// let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?; +// Ok(Box::new(model) as Box) +// }, +// ModelType::MiniCPM4 => { +// let model = MiniCPMGenerateModel::init(model_path, device, dtype)?; +// Ok(Box::new(model)as Box) +// } +// } +// } +// } diff --git a/src/models/base_modules/mod.rs b/src/models/base_modules/mod.rs index be470d1..774719e 100644 --- a/src/models/base_modules/mod.rs +++ b/src/models/base_modules/mod.rs @@ -116,6 +116,7 @@ impl AttentionNobias { cos: &Tensor, sin: &Tensor, attention_mask: Option<&Tensor>, + tof32: bool, ) -> Result { let (b_sz, q_len, _) = xs.dims3()?; let query_states = self.q_proj.forward(xs)?; @@ -131,7 +132,7 @@ impl AttentionNobias { .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? .transpose(1, 2)?; let (query_states, key_states) = - apply_rotary_pos_emb(&query_states, &key_states, cos, sin)?; + apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?; let key_states = repeat_kv(key_states, self.num_kv_groups)?.contiguous()?; let value_states = repeat_kv(value_states, self.num_kv_groups)?.contiguous()?; @@ -184,6 +185,7 @@ impl AttentionNobias { cos: &Tensor, sin: &Tensor, attention_mask: Option<&Tensor>, + tof32: bool, ) -> Result { let (b_sz, q_len, _) = xs.dims3()?; let query_states = self.q_proj.forward(xs)?; @@ -199,7 +201,7 @@ impl AttentionNobias { .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? .transpose(1, 2)?; let (query_states, key_states) = - apply_rotary_pos_emb(&query_states, &key_states, cos, sin)?; + apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?; let (key_states, value_states) = match &self.kv_cache { None => (key_states, value_states), Some((prev_k, prev_v)) => { diff --git a/src/models/minicpm4/config.rs b/src/models/minicpm4/config.rs index af4d7d4..0b9d0f2 100644 --- a/src/models/minicpm4/config.rs +++ b/src/models/minicpm4/config.rs @@ -23,10 +23,7 @@ pub struct MiniCPM4Config { pub rope_scaling: RopeScalingConfig, pub torch_dtype: String, pub vocab_size: usize, - // pub use_mup: bool, - pub scale_emb:f32, + pub scale_emb: f64, pub dim_model_base: usize, pub scale_depth: f32, - // pub rope_theta: f32, - // pub kv_channels: i32, } \ No newline at end of file diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 7230e43..6c5f644 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -2,15 +2,14 @@ use crate::models::minicpm4::config::MiniCPM4Config; use crate::models::minicpm4::model::MiniCPMModel; // use crate::models::GenerateStream; use crate::utils::utils::{ - build_completion_chunk_response, build_completion_response, find_safetensors_files, get_device, - get_dtype, get_logit_processor, + build_completion_chunk_response, build_completion_response, find_type_files, get_device, get_dtype, get_logit_processor }; use crate::{ chat_template::chat_template::ChatTemplate, models::GenerateModel, tokenizer::tokenizer::TokenizerModel, }; use anyhow::{Result, anyhow}; -use candle_core::{D, DType, Device, IndexOp, Tensor}; +use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use openai_dive::v1::resources::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, @@ -27,7 +26,7 @@ pub struct MiniCPMGenerateModel<'a> { im_end_id: u32, } -impl<'a> MiniCPMGenerateModel<'a> { +impl <'a> MiniCPMGenerateModel<'a> { pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { let chat_template = ChatTemplate::init(path)?; let tokenizer = TokenizerModel::init(path)?; @@ -38,7 +37,7 @@ impl<'a> MiniCPMGenerateModel<'a> { let dtype = get_dtype(dtype, cfg_dtype); let endoftext_id = cfg.eos_token_id[0]; let im_end_id = cfg.eos_token_id[1]; - let model_list = find_safetensors_files(&path)?; + let model_list = find_type_files(&path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; let minicpm = MiniCPMModel::new(vb, cfg)?; @@ -64,7 +63,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { let mut generate = Vec::new(); let sample_len = match mes.max_tokens { Some(max) => max, - None => 512, + None => 2048, }; for _ in 0..sample_len { let logits = self.minicpm.forward_step(&input_ids, seqlen_offset)?; diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index 9e17c91..7ffee13 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -7,38 +7,41 @@ use crate::{ utils::tensor_utils::prepare_causal_attention_mask, }; use anyhow::{Ok, Result}; -use candle_core::{D, DType, Device, Tensor, Var}; -use candle_nn::{embedding, rms_norm, Embedding, Linear, Module, RmsNorm, VarBuilder}; +use candle_core::{D, Device, Tensor}; +use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_norm}; pub struct MiniCPMLongRoPE { - head_dim: usize, - rope_theta: f32, - max_position_embeddings: usize, short_factor: Vec, long_factor: Vec, original_max_position_embeddings: usize, + max_seq_len_cached: usize, + scaling_factor: f64, inv_freq: Tensor, cos_cached: Tensor, sin_cached: Tensor, + device: Device, } impl MiniCPMLongRoPE { pub fn new(cfg: &MiniCPM4Config, device: &Device) -> Result { let head_dim = cfg.hidden_size / cfg.num_attention_heads; let rope_theta = 10000.0; - let max_position_embeddings = cfg.max_position_embeddings; let short_factor = cfg.rope_scaling.short_factor.clone(); let long_factor = cfg.rope_scaling.short_factor.clone(); let original_max_position_embeddings = cfg.rope_scaling.original_max_position_embeddings; - let scale = max_position_embeddings / original_max_position_embeddings; + let max_position_embeddings = cfg.max_position_embeddings; + let scale = max_position_embeddings as f64 / original_max_position_embeddings as f64; let scaling_factor = - (1.0 + (scale as f64).ln() + (original_max_position_embeddings as f64).ln()).sqrt(); + (1.0 + scale.ln() / (original_max_position_embeddings as f64).ln()).sqrt(); let inv_freq = compute_default_rope_parameters(head_dim, rope_theta); - let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?; + let inv_freq = + Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?; + let max_seq_len_cached = max_position_embeddings; let t = Tensor::arange(0.0_f32, max_position_embeddings as f32, device)? .reshape((max_position_embeddings, 1))?; // short_factor.len() = 32 // head_dim = 1024 / 16 = 64, inv_freq.len() = 32 - let ext_factors = Tensor::from_slice(&short_factor, (1, short_factor.len()), device)?; + let ext_factors = + Tensor::from_slice(&short_factor, (1, short_factor.len()), device)?; let ext_factors = Tensor::ones_like(&ext_factors)?.div(&ext_factors)?; // (seq_len, 1) matmul (1, 32) -> (seq_len, 32) * (1, 32)-> (seq_len, 32) let freqs = t.matmul(&ext_factors)?.broadcast_mul(&inv_freq)?; @@ -47,41 +50,46 @@ impl MiniCPMLongRoPE { let cos_cached = emb.cos()?.affine(scaling_factor, 0.0)?; let sin_cached = emb.sin()?.affine(scaling_factor, 0.0)?; Ok(Self { - head_dim, - rope_theta, - max_position_embeddings, short_factor, long_factor, original_max_position_embeddings, + max_seq_len_cached, + scaling_factor, inv_freq, cos_cached, sin_cached, + device: device.clone(), }) } - pub fn update_cos_sin_cache(&mut self, seqlen: usize, device: &Device) -> Result<()> { - let t = Tensor::arange(0.0_f32, seqlen as f32, device)?.reshape((seqlen, 1))?; - let mut ext_factors = - Tensor::from_slice(&self.short_factor, (1, self.short_factor.len()), device)?; + pub fn update_cos_sin_cache(&mut self, seqlen: usize) -> Result<()> { + self.max_seq_len_cached = seqlen; + let t = Tensor::arange(0.0_f32, seqlen as f32, &self.device)? + .reshape((seqlen, 1))?; + let mut ext_factors = Tensor::from_slice( + &self.short_factor, + (1, self.short_factor.len()), + &self.device, + )?; if seqlen > self.original_max_position_embeddings { ext_factors = - Tensor::from_slice(&self.long_factor, (1, self.long_factor.len()), device)?; + Tensor::from_slice(&self.long_factor, (1, self.long_factor.len()), &self.device)?; } let ext_factors = Tensor::ones_like(&ext_factors)?.div(&ext_factors)?; let freqs = t.matmul(&ext_factors)?.broadcast_mul(&self.inv_freq)?; let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?; - let scale = seqlen / self.original_max_position_embeddings; - let scaling_factor = - (1.0 + (scale as f64).ln() + (self.original_max_position_embeddings as f64).ln()) - .sqrt(); - let cos_cached = emb.cos()?.affine(scaling_factor, 0.0)?; - let sin_cached = emb.sin()?.affine(scaling_factor, 0.0)?; + let cos_cached = emb.cos()?.affine(self.scaling_factor, 0.0)?; + let sin_cached = emb.sin()?.affine(self.scaling_factor, 0.0)?; self.cos_cached = cos_cached; self.sin_cached = sin_cached; Ok(()) } - pub fn forward(&self, pos_offset: usize, seqlen: usize) -> Result<(Tensor, Tensor)> { + pub fn forward(&mut self, pos_offset: usize, seqlen: usize) -> Result<(Tensor, Tensor)> { + if pos_offset + seqlen > self.max_seq_len_cached { + let _ = self.update_cos_sin_cache(pos_offset + seqlen)?; + } let cos = self.cos_cached.narrow(0, pos_offset, seqlen)?; let sin = self.sin_cached.narrow(0, pos_offset, seqlen)?; + Ok((cos, sin)) } } @@ -133,13 +141,21 @@ impl MiniCPMDecoderLayer { sin: &Tensor, attention_mask: Option<&Tensor>, ) -> Result { - let residual = xs; + let residual = xs.clone(); let xs = self.input_layernorm.forward(xs)?; - let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?; - let xs = (xs + residual)?; + let xs = self.self_attn.forward(&xs, cos, sin, attention_mask, true)?; + let xs = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; let residual = &xs; let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?; - let xs = (residual + xs)?; + let xs = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; Ok(xs) } @@ -150,13 +166,21 @@ impl MiniCPMDecoderLayer { sin: &Tensor, attention_mask: Option<&Tensor>, ) -> Result { - let residual = xs; + let residual = xs.clone(); let xs = self.input_layernorm.forward(xs)?; - let xs = self.self_attn.forward_step(&xs, cos, sin, attention_mask)?; - let xs = (xs + residual)?; + let xs = self.self_attn.forward_step(&xs, cos, sin, attention_mask, true)?; + let xs = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; let residual = &xs; let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?; - let xs = (residual + xs)?; + let xs = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; Ok(xs) } pub fn clear_kv_cache(&mut self) { @@ -175,6 +199,7 @@ pub struct MiniCPMModel { impl MiniCPMModel { pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> Result { + let vb = vb.pp("model"); let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?; let mut layers = Vec::with_capacity(cfg.num_hidden_layers); let vb_layers = vb.pp("layers"); @@ -191,13 +216,16 @@ impl MiniCPMModel { layers, norm, rope_emb, - lm_head + lm_head, }) } - pub fn forward(&self, input_ids: &Tensor, position_id: usize) -> Result { + pub fn forward(&mut self, input_ids: &Tensor, position_id: usize) -> Result { let (bs, seq_len) = input_ids.dims2()?; - let input_embeds = self.embed_tokens.forward(&input_ids)?; + let input_embeds = self + .embed_tokens + .forward(&input_ids)? + .affine(self.cfg.scale_emb, 0.0)?; let attention_mask: Option<&Tensor> = { if seq_len <= 1 { None @@ -210,7 +238,7 @@ impl MiniCPMModel { )?) } }; - + let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let mut hidden_states = input_embeds; for decode_layer in &self.layers { @@ -218,13 +246,20 @@ impl MiniCPMModel { } hidden_states = self.norm.forward(&hidden_states)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; + let hidden_state = hidden_state.affine( + 1.0 / (self.cfg.hidden_size / self.cfg.dim_model_base) as f64, + 0.0, + )?; let logits = self.lm_head.forward(&hidden_state)?; Ok(logits) } pub fn forward_step(&mut self, input_ids: &Tensor, position_id: usize) -> Result { let (bs, seq_len) = input_ids.dims2()?; - let input_embeds = self.embed_tokens.forward(&input_ids)?; + let input_embeds = self + .embed_tokens + .forward(&input_ids)? + .affine(self.cfg.scale_emb, 0.0)?; let attention_mask: Option<&Tensor> = { if seq_len <= 1 { None @@ -237,14 +272,18 @@ impl MiniCPMModel { )?) } }; - let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let mut hidden_states = input_embeds; for decode_layer in &mut self.layers { - hidden_states = decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?; + hidden_states = + decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?; } hidden_states = self.norm.forward(&hidden_states)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; + let hidden_state = hidden_state.affine( + 1.0 / (self.cfg.hidden_size / self.cfg.dim_model_base) as f64, + 0.0, + )?; let logits = self.lm_head.forward(&hidden_state)?; Ok(logits) } diff --git a/src/models/mod.rs b/src/models/mod.rs index 94a8ece..1f16059 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -1,18 +1,15 @@ pub mod base_modules; pub mod minicpm4; pub mod qwen2_5vl; +pub mod voxcpm; use anyhow::Result; -use candle_core::{DType, Device}; use openai_dive::v1::resources::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; use rocket::futures::Stream; pub trait GenerateModel { - // fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result - // where - // Self: Sized; fn generate(&mut self, mes: ChatCompletionParameters) -> Result; fn generate_stream( &mut self, diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 4536d56..67d25c1 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -1,8 +1,7 @@ // use crate::models::GenerateStream; use crate::models::qwen2_5vl::config::Qwen2_5VLConfig; use crate::utils::utils::{ - build_completion_chunk_response, build_completion_response, find_safetensors_files, get_device, - get_dtype, get_logit_processor, + build_completion_chunk_response, build_completion_response, find_type_files, get_device, get_dtype, get_logit_processor }; use crate::{ chat_template::chat_template::ChatTemplate, @@ -43,7 +42,8 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { let pre_processor = Qwen2_5VLProcessor::new(device, dtype)?; let endoftext_id = cfg.bos_token_id; let im_end_id = cfg.eos_token_id; - let model_list = find_safetensors_files(&path)?; + // let model_list = find_safetensors_files(&path)?; + let model_list = find_type_files(&path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; let qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?; @@ -86,7 +86,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let mut generate = Vec::new(); let sample_len = match mes.max_tokens { Some(max) => max, - None => 512, + None => 1024, }; for _ in 0..sample_len { let logits = self.qwen2_5_vl.forward( diff --git a/src/models/qwen2_5vl/model.rs b/src/models/qwen2_5vl/model.rs index 1b4c1f7..e6d61d3 100644 --- a/src/models/qwen2_5vl/model.rs +++ b/src/models/qwen2_5vl/model.rs @@ -631,7 +631,7 @@ impl Qwen2_5VLTextAttention { .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? .transpose(1, 2)?; let (query_states, key_states) = - apply_rotary_pos_emb(&query_states, &key_states, cos, sin)?; + 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)) => { diff --git a/src/models/voxcpm/audio_vae.rs b/src/models/voxcpm/audio_vae.rs new file mode 100644 index 0000000..1c6c918 --- /dev/null +++ b/src/models/voxcpm/audio_vae.rs @@ -0,0 +1,586 @@ +use anyhow::{Error, Ok, Result}; +use candle_core::{D, IndexOp, Tensor}; +use candle_nn::{Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Module, VarBuilder}; +use std::result::Result::Ok as StdOk; + +pub struct CausalConv1d { + conv1d: Conv1d, + padding: usize, +} + +impl CausalConv1d { + // CausalConv1d::new(scaled_weight, bias, padding, dilation, stride)?; + pub fn new( + weight: Tensor, + bias: Option, + // in_c: usize, + // out_c: usize, + // kernel_size: usize, + padding: usize, + dilation: usize, + groups: usize, + stride: usize, + ) -> Result { + let config = Conv1dConfig { + padding: 0, + stride, + dilation, + groups, + cudnn_fwd_algo: None, + }; + + let conv1d = Conv1d::new(weight, bias, config); + Ok(Self { conv1d, padding }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let x_pad = x.pad_with_zeros(D::Minus1, self.padding * 2, 0)?; + let x = self.conv1d.forward(&x_pad)?; + Ok(x) + } +} + +pub struct CausalConvTranspose1d { + conv_transpose1d: ConvTranspose1d, + padding: usize, + output_padding: usize, + config: ConvTranspose1dConfig, +} + +impl CausalConvTranspose1d { + pub fn new( + weight: Tensor, + bias: Option, + padding: usize, + dilation: usize, + output_padding: usize, + groups: usize, + stride: usize, + ) -> Result { + let config = ConvTranspose1dConfig { + padding: 0, + output_padding, + stride, + dilation, + groups, + }; + + let conv_transpose1d = ConvTranspose1d::new(weight, bias, config.clone()); + Ok(Self { + conv_transpose1d, + padding, + output_padding, + config + }) + } + pub fn forward(&self, x: &Tensor) -> Result { + println!("transpose conv input x: {:?}", x); + println!("transpose conv config stride: {:?}", self.config.stride); + println!("transpose conv config padding: {:?}", self.config.padding); + println!("transpose conv config output_padding: {:?}", self.config.output_padding); + println!("transpose conv config groups: {:?}", self.config.groups); + println!("transpose conv config dilation: {:?}", self.config.dilation); + println!("transpose conv config weight: {:?}", self.conv_transpose1d.weight()); + + let x = self.conv_transpose1d.forward(x)?; + println!("transpose conv after x: {:?}", x); + println!("transpose conv after self.padding: {:?}", self.padding); + println!("transpose conv after self.output_padding: {:?}", self.output_padding); + let last_dim = x.dim(D::Minus1)?; + let select_num = last_dim - (self.padding * 2 - self.output_padding); + println!("transpose conv after select_num: {:?}", select_num); + let x = x.narrow(D::Minus1, 0, select_num)?; + println!("transpose conv after x: {:?}", x); + Ok(x) + } +} + +pub struct WNCausalConv1d { + conv: CausalConv1d, +} +impl WNCausalConv1d { + pub fn new( + vb: VarBuilder, + in_c: usize, + out_c: usize, + kernel_size: usize, + dilation: usize, + padding: usize, + groups: usize, + stride: usize, + ) -> Result { + let in_c = in_c / groups; + let weight_g = vb.get((out_c, 1, 1), "weight_g")?; + let weight_v = vb.get((out_c, in_c, kernel_size), "weight_v")?; + let bias = match vb.get(out_c, "bias") { + StdOk(b) => Some(b), + Err(_) => None, + }; + let weight_norm = weight_v.sqr()?.sum_keepdim(D::Minus1)?.sqrt()?; + let normalized_weight = weight_v.broadcast_div(&weight_norm)?; + let scaled_weight = normalized_weight.broadcast_mul(&weight_g)?; + let conv = CausalConv1d::new(scaled_weight, bias, padding, dilation, groups, stride)?; + Ok(Self { conv }) + } + pub fn forward(&self, x: &Tensor) -> Result { + println!("conv1d: x: {:?}", x); + println!("conv weight: : {:?}", self.conv.conv1d.weight()); + let x = self.conv.forward(x)?; + println!("conv1d: WN causal x: {:?}", x); + Ok(x) + } +} + +pub struct WNCausalConvTranspose1d { + conv_transpose: CausalConvTranspose1d, +} + +impl WNCausalConvTranspose1d { + pub fn new( + vb: VarBuilder, + in_c: usize, + out_c: usize, + dilation: usize, + kernel_size: usize, + padding: usize, + output_padding: usize, + groups: usize, + stride: usize, + ) -> Result { + let in_c = in_c / groups; + let weight_g = vb.get((in_c, 1, 1), "weight_g")?; + let weight_v = vb.get((in_c, out_c, kernel_size), "weight_v")?; + let bias = match vb.get(out_c, "bias") { + StdOk(b) => Some(b), + Err(_) => None, + }; + let weight_norm = weight_v.sqr()?.sum_keepdim(D::Minus1)?.sqrt()?; + let normalized_weight = weight_v.broadcast_div(&weight_norm)?; + let scaled_weight = normalized_weight.broadcast_mul(&weight_g)?; + let conv_transpose = CausalConvTranspose1d::new( + scaled_weight, + bias, + padding, + dilation, + output_padding, + groups, + stride, + )?; + Ok(Self { conv_transpose }) + } + pub fn forward(&self, x: &Tensor) -> Result { + let x = self.conv_transpose.forward(x)?; + Ok(x) + } +} + +pub struct Snake1d { + alpha: Tensor, +} +impl Snake1d { + pub fn new(vb: VarBuilder, channels: usize) -> Result { + let alpha = vb.get((1, channels, 1), "alpha")?; + Ok(Self { alpha }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let dims = x.dims(); + let x = x.reshape((dims[0], dims[1], ()))?; + let alpha_ = self.alpha.affine(1.0, 1e-9)?.recip()?; + let alpha_ = x + .broadcast_mul(&self.alpha)? + .sin()? + .powf(2.0)? + .broadcast_mul(&alpha_)?; + let x = x.add(&alpha_)?; + let x = x.reshape(dims)?; + Ok(x) + } +} + +pub struct CausalResidualUnit { + // pad: usize, + block0: Snake1d, + block1: WNCausalConv1d, + block2: Snake1d, + block3: WNCausalConv1d, +} + +impl CausalResidualUnit { + pub fn new( + vb: VarBuilder, + dim: usize, + dilation: usize, + kernel: usize, + groups: usize, + ) -> Result { + let pad = ((7 - 1) * dilation) / 2; + let block0 = Snake1d::new(vb.pp("block.0"), dim)?; + let block1 = + WNCausalConv1d::new(vb.pp("block.1"), dim, dim, kernel, dilation, pad, groups, 1)?; + let block2 = Snake1d::new(vb.pp("block.2"), dim)?; + let block3 = WNCausalConv1d::new(vb.pp("block.3"), dim, dim, 1, 1, 0, 1, 1)?; + Ok(Self { + // pad, + block0, + block1, + block2, + block3, + }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + println!("causal residual unit x: {:?}", x); + // let orig_dim = x.dims(); + let last_dim_x = x.dim(D::Minus1)?; + let mut res_x = x.clone(); + let y = self.block0.forward(x)?; + let y = self.block1.forward(&y)?; + let y = self.block2.forward(&y)?; + let y = self.block3.forward(&y)?; + println!("causal residual unit y: {:?}", y); + // let dim = y.dims(); + let last_dim_y = y.dim(D::Minus1)?; + println!("last_dim_x: {:?}", last_dim_x); + println!("last_dim_y: {:?}", last_dim_y); + let pad = (last_dim_x - last_dim_y) / 2; + if pad > 0 { + res_x = res_x.narrow(D::Minus1, pad, last_dim_y)?; + } + let x = y.add(&res_x)?; + Ok(x) + } +} + +pub struct CausalEncoderBlock { + block0: CausalResidualUnit, + block1: CausalResidualUnit, + block2: CausalResidualUnit, + block3: Snake1d, + block4: WNCausalConv1d, +} + +impl CausalEncoderBlock { + pub fn new( + vb: VarBuilder, + in_dim: Option, + out_dim: usize, + stride: usize, + groups: usize, + ) -> Result { + let in_dim = match in_dim { + Some(d) => d, + None => out_dim / 2, + }; + let block0 = CausalResidualUnit::new(vb.pp("block.0"), in_dim, 1, 7, groups)?; + let block1 = CausalResidualUnit::new(vb.pp("block.1"), in_dim, 3, 7, groups)?; + let block2 = CausalResidualUnit::new(vb.pp("block.2"), in_dim, 9, 7, groups)?; + let block3 = Snake1d::new(vb.pp("block.3"), in_dim)?; + let padding = (stride as f32 / 2.0).ceil() as usize; + let block4 = WNCausalConv1d::new( + vb.pp("block.4"), + in_dim, + out_dim, + 2 * stride, + 1, + padding, + 1, + stride, + )?; + Ok(Self { + block0, + block1, + block2, + block3, + block4, + }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let x = self.block0.forward(x)?; + let x = self.block1.forward(&x)?; + let x = self.block2.forward(&x)?; + let x = self.block3.forward(&x)?; + let x = self.block4.forward(&x)?; + Ok(x) + } +} + +pub struct CausalEncoder { + block0: WNCausalConv1d, + block1_4: Vec, + fc_mu: WNCausalConv1d, + fc_logvar: WNCausalConv1d, +} + +impl CausalEncoder { + pub fn new( + vb: VarBuilder, + d_model: usize, + laten_dim: usize, + strides: Vec, + depthwise: bool, + ) -> Result { + let mut d_model = d_model; + let mut groups = 1; + let block0 = WNCausalConv1d::new(vb.pp("block.0"), 1, d_model, 7, 1, 3, 1, 1)?; + let vb_block = vb.pp("block"); + let mut block1_4 = Vec::new(); + for (i, stride) in strides.iter().enumerate() { + d_model *= 2; + groups = if depthwise { d_model / 2 } else { 1 }; + let block_i = CausalEncoderBlock::new(vb_block.pp(i+1), None, d_model, *stride, groups)?; + block1_4.push(block_i); + } + let fc_mu = WNCausalConv1d::new(vb.pp("fc_mu"), d_model, laten_dim, 3, 1, 1, 1, 1)?; + let fc_logvar = WNCausalConv1d::new(vb.pp("fc_logvar"), d_model, laten_dim, 3, 1, 1, 1, 1)?; + Ok(Self { + block0, + block1_4, + fc_mu, + fc_logvar, + }) + } + + pub fn forward(&self, x: &Tensor) -> Result<(Tensor, Tensor, Tensor)> { + let mut hidden_state = self.block0.forward(x)?; + for block_i in &self.block1_4 { + hidden_state = block_i.forward(&hidden_state)?; + } + let mu = self.fc_mu.forward(&hidden_state)?; + let logvar = self.fc_logvar.forward(&hidden_state)?; + Ok((hidden_state, mu, logvar)) + } +} + +pub struct NoiseBlock { + linear: WNCausalConv1d, +} + +impl NoiseBlock { + pub fn new(vb: VarBuilder, dim: usize) -> Result { + let linear = WNCausalConv1d::new(vb.pp("linear"), dim, dim, 1, 1, 0, 1, 1)?; + Ok(Self { linear }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let (bs, _, t) = x.dims3()?; + let noise = Tensor::randn(0.0_f32, 1.0, (bs, 1, t), x.device())?.to_dtype(x.dtype())?; + let h = self.linear.forward(x)?; + let n = h.broadcast_mul(&noise)?; + let x = x.add(&n)?; + Ok(x) + } +} + +pub struct CausalDecoderBlock { + block0: Snake1d, + block1: WNCausalConvTranspose1d, + block2: CausalResidualUnit, + block3: CausalResidualUnit, + block4: CausalResidualUnit, +} + +impl CausalDecoderBlock { + pub fn new( + vb: VarBuilder, + input_dim: usize, + output_dim: usize, + stride: usize, + groups: usize, + ) -> Result { + let block0 = Snake1d::new(vb.pp("block.0"), input_dim)?; + let padding = (stride as f32 / 2.0).ceil() as usize; + let block1 = WNCausalConvTranspose1d::new( + vb.pp("block.1"), + input_dim, + output_dim, + 1, + 2 * stride, + padding, + stride % 2, + 1, + stride, + )?; + let block2 = CausalResidualUnit::new(vb.pp("block.2"), output_dim, 1, 7, groups)?; + let block3 = CausalResidualUnit::new(vb.pp("block.3"), output_dim, 3, 7, groups)?; + let block4 = CausalResidualUnit::new(vb.pp("block.4"), output_dim, 9, 7, groups)?; + Ok(Self { + block0, + block1, + block2, + block3, + block4, + }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + println!("decoder block x : {:?}", x); + let x = self.block0.forward(x)?; + println!("decoder block0 x : {:?}", x); + let x = self.block1.forward(&x)?; + println!("decoder block1 x : {:?}", x); + let x = self.block2.forward(&x)?; + println!("decoder block2 x : {:?}", x); + let x = self.block3.forward(&x)?; + println!("decoder block3 x : {:?}", x); + let x = self.block4.forward(&x)?; + println!("decoder block4 x : {:?}", x); + Ok(x) + } +} + +pub struct CausalDecoder { + model0: WNCausalConv1d, + model1: WNCausalConv1d, + model2_5: Vec, + model6: Snake1d, + model7: WNCausalConv1d, +} + +impl CausalDecoder { + pub fn new( + vb: VarBuilder, + input_channel: usize, + channels: usize, + rates: Vec, + d_out: usize, + ) -> Result { + let model0 = WNCausalConv1d::new( + vb.pp("model.0"), + input_channel, + input_channel, + 7, + 1, + 3, + input_channel, + 1, + )?; + let model1 = WNCausalConv1d::new(vb.pp("model.1"), input_channel, channels, 1, 1, 1, 1, 1)?; + let vb_model = vb.pp("model"); + let mut output_dim = channels; + let mut model2_5 = Vec::new(); + for (i, stride) in rates.iter().enumerate() { + let input_dim = channels / 2_usize.pow(i as u32); + output_dim = channels / 2_usize.pow((i + 1) as u32); + let groups = output_dim; + let model_i = CausalDecoderBlock::new( + vb_model.pp(i + 2), + input_dim, + output_dim, + *stride, + groups, + )?; + model2_5.push(model_i); + } + let model6 = Snake1d::new(vb.pp("model.6"), output_dim)?; + let model7 = WNCausalConv1d::new(vb.pp("model.7"), output_dim, d_out, 7, 1, 3, 1, 1)?; + Ok(Self { + model0, + model1, + model2_5, + model6, + model7, + }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + print!("audio_vae decoder input x shape: {:?}", x); + let x = self.model0.forward(x)?; + print!("audio_vae decoder model0 x shape: {:?}", x); + let mut x = self.model1.forward(&x)?; + for model_i in &self.model2_5 { + x = model_i.forward(&x)?; + } + let x = self.model6.forward(&x)?; + let x = self.model7.forward(&x)?; + let x = x.tanh()?; + Ok(x) + } +} + +pub struct AudioVAE { + encoder_dim: usize, + encoder_rates: Vec, + decoder_dim: usize, + decoder_rates: Vec, + pub latent_dim: usize, + hop_length: usize, + encoder: CausalEncoder, + decoder: CausalDecoder, + pub sample_rate: usize, + pub chunk_size: usize, +} + +impl AudioVAE { + pub fn new( + vb: VarBuilder, + encoder_dim: usize, + encoder_rates: Vec, + laten_dim: Option, + decoder_dim: usize, + decoder_rates: Vec, + sample_rate: usize, + ) -> Result { + let latent_dim = match laten_dim { + Some(d) => d, + None => encoder_dim * (2_usize.pow(encoder_rates.len() as u32)), + }; + let hop_length = encoder_rates.iter().product(); + let encoder = CausalEncoder::new( + vb.pp("encoder"), + encoder_dim, + latent_dim, + encoder_rates.clone(), + true, + )?; + let decoder = CausalDecoder::new( + vb.pp("decoder"), + latent_dim, + decoder_dim, + decoder_rates.clone(), + 1, + )?; + let chunk_size = hop_length; + Ok(Self { + encoder_dim, + encoder_rates, + decoder_dim, + decoder_rates, + latent_dim, + hop_length, + encoder, + decoder, + sample_rate, + chunk_size, + }) + } + + pub fn preprocess(&self, audio_data: &Tensor, sample_rate: Option) -> Result{ + let sample_rate = match sample_rate { + Some(r) => r, + None => self.sample_rate + }; + assert_eq!(sample_rate, self.sample_rate); + let pad_to = self.hop_length; + let length = audio_data.dim(D::Minus1)?; + let right_pad = (length as f32 / pad_to as f32).ceil() as usize * pad_to - length; + let audio_data = audio_data.pad_with_zeros(D::Minus1, 0, right_pad)?; + Ok(audio_data) + } + + pub fn decode(&self, z: &Tensor) -> Result { + let x = self.decoder.forward(z)?; + Ok(x) + } + + pub fn encode(&self, audio_data: &Tensor, sample_rate: Option) -> Result { + let audio_data = match audio_data.rank() { + 2 => audio_data.unsqueeze(1)?, + _ => audio_data.clone() + }; + let audio_data = self.preprocess(&audio_data, sample_rate)?; + let (_, mu, _) = self.encoder.forward(&audio_data)?; + Ok(mu) + } +} diff --git a/src/models/voxcpm/config.rs b/src/models/voxcpm/config.rs index ebc296b..b37ebf4 100644 --- a/src/models/voxcpm/config.rs +++ b/src/models/voxcpm/config.rs @@ -1,15 +1,14 @@ -use candle_nn::Activation; #[derive(Debug, Clone, PartialEq, serde::Deserialize)] -pub struct RopeScalingConfig { - pub rope_type: String, +pub struct VoxRopeScalingConfig { + pub r#type: String, pub long_factor: Vec, pub short_factor: Vec, pub original_max_position_embeddings: usize, } #[derive(Debug, Clone, PartialEq, serde::Deserialize)] -pub struct MiniCPM4Config { +pub struct VoxMiniCPM4Config { pub bos_token_id: u32, pub eos_token_id: u32, pub hidden_size: usize, @@ -19,38 +18,50 @@ pub struct MiniCPM4Config { pub num_hidden_layers: usize, pub num_key_value_heads: usize, pub rms_norm_eps: f64, - pub rope_scaling: RopeScalingConfig, - pub torch_dtype: String, + pub rope_theta: f32, + pub rope_scaling: VoxRopeScalingConfig, pub vocab_size: usize, - // pub use_mup: bool, pub scale_emb:f32, pub dim_model_base: usize, pub scale_depth: f32, - // pub rope_theta: f32, - // pub kv_channels: i32, + pub use_mup: bool, } #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct VoxCPMEncoderConfig { - hidden_dim: usize, - ffn_dim: usize, - num_heads: usize, - num_layers: usize, + pub hidden_dim: usize, + pub ffn_dim: usize, + pub num_heads: usize, + pub num_layers: usize, } #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct CfmConfig { - sigma_min: f32, - solver: String, - t_scheduler: String, - inference_cfg_rate: f32, + pub sigma_min: f32, + pub solver: String, + pub t_scheduler: String, + pub inference_cfg_rate: f32, } #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct VoxCPMDitConfig { - hidden_dim: usize, - ffn_dim: usize, - num_heads: usize, - num_layers: usize, - cfm_config: CfmConfig, + pub hidden_dim: usize, + pub ffn_dim: usize, + pub num_heads: usize, + pub num_layers: usize, + pub cfm_config: CfmConfig, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct VoxCPMConfig { + pub lm_config: VoxMiniCPM4Config, + pub patch_size: usize, + pub feat_dim: usize, + pub scalar_quantization_latent_dim: usize, + pub scalar_quantization_scale: usize, + pub residual_lm_num_layers: usize, + pub encoder_config: VoxCPMEncoderConfig, + pub dit_config: VoxCPMDitConfig, + pub max_length: usize, + pub dtype: String, } \ No newline at end of file diff --git a/src/models/voxcpm/minicpm4.rs b/src/models/voxcpm/minicpm4.rs new file mode 100644 index 0000000..b454010 --- /dev/null +++ b/src/models/voxcpm/minicpm4.rs @@ -0,0 +1,330 @@ +use std::{thread, time}; + +use crate::{ + models::{ + base_modules::{AttentionNobias, MLPNoBias}, + voxcpm::config::VoxMiniCPM4Config, + }, + position_embed::rope::compute_default_rope_parameters, + utils::tensor_utils::prepare_causal_attention_mask, +}; +use anyhow::{anyhow, Ok, Result}; +use candle_core::{DType, Device, Tensor, D}; +use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_norm}; + +pub struct MiniCPMLongRoPE { + short_factor: Vec, + long_factor: Vec, + original_max_position_embeddings: usize, + max_seq_len_cached: usize, + scaling_factor: f64, + inv_freq: Tensor, + cos_cached: Tensor, + sin_cached: Tensor, + device: Device, + dtype: DType, +} +impl MiniCPMLongRoPE { + pub fn new(cfg: &VoxMiniCPM4Config, device: &Device, dtype: DType) -> Result { + let head_dim = cfg.hidden_size / cfg.num_attention_heads; + let rope_theta = cfg.rope_theta; + let short_factor = cfg.rope_scaling.short_factor.clone(); + let long_factor = cfg.rope_scaling.short_factor.clone(); + let original_max_position_embeddings = cfg.rope_scaling.original_max_position_embeddings; + let max_position_embeddings = cfg.max_position_embeddings; + let scale = max_position_embeddings as f64 / original_max_position_embeddings as f64; + let scaling_factor = + (1.0 + scale.ln() / (original_max_position_embeddings as f64).ln()).sqrt(); + let inv_freq = compute_default_rope_parameters(head_dim, rope_theta); + let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?; + let max_seq_len_cached = max_position_embeddings; + let t = Tensor::arange(0.0_f32, max_position_embeddings as f32, device)? + .reshape((max_position_embeddings, 1))?; + // short_factor.len() = 32 + // head_dim = 1024 / 16 = 64, inv_freq.len() = 32 + let ext_factors = Tensor::from_slice(&short_factor, (1, short_factor.len()), device)?; + let ext_factors = Tensor::ones_like(&ext_factors)?.div(&ext_factors)?; + // (seq_len, 1) matmul (1, 32) -> (seq_len, 32) * (1, 32)-> (seq_len, 32) + let freqs = t.matmul(&ext_factors)?.broadcast_mul(&inv_freq)?; + + let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?; + let cos_cached = emb.cos()?.affine(scaling_factor, 0.0)?.to_dtype(dtype)?; + let sin_cached = emb.sin()?.affine(scaling_factor, 0.0)?.to_dtype(dtype)?; + Ok(Self { + short_factor, + long_factor, + original_max_position_embeddings, + max_seq_len_cached, + scaling_factor, + inv_freq, + cos_cached, + sin_cached, + device: device.clone(), + dtype, + }) + } + pub fn update_cos_sin_cache(&mut self, seqlen: usize) -> Result<()> { + self.max_seq_len_cached = seqlen; + let t = Tensor::arange(0.0_f32, seqlen as f32, &self.device)?.reshape((seqlen, 1))?; + let mut ext_factors = Tensor::from_slice( + &self.short_factor, + (1, self.short_factor.len()), + &self.device, + )?; + if seqlen > self.original_max_position_embeddings { + ext_factors = + Tensor::from_slice(&self.long_factor, (1, self.long_factor.len()), &self.device)?; + } + let ext_factors = Tensor::ones_like(&ext_factors)?.div(&ext_factors)?; + let freqs = t.matmul(&ext_factors)?.broadcast_mul(&self.inv_freq)?; + let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?; + let cos_cached = emb.cos()?.affine(self.scaling_factor, 0.0)?.to_dtype(self.dtype)?; + let sin_cached = emb.sin()?.affine(self.scaling_factor, 0.0)?.to_dtype(self.dtype)?; + self.cos_cached = cos_cached; + self.sin_cached = sin_cached; + Ok(()) + } + pub fn forward(&mut self, pos_offset: usize, seqlen: usize) -> Result<(Tensor, Tensor)> { + if pos_offset + seqlen > self.max_seq_len_cached { + let _ = self.update_cos_sin_cache(pos_offset + seqlen)?; + } + let cos = self.cos_cached.narrow(0, pos_offset, seqlen)?; + let sin = self.sin_cached.narrow(0, pos_offset, seqlen)?; + + Ok((cos, sin)) + } +} + +pub struct MiniCPMDecoderLayer { + self_attn: AttentionNobias, + mlp: MLPNoBias, + input_layernorm: RmsNorm, + post_attention_layernorm: RmsNorm, + scale_depth: f32, + num_hidden_layers: usize, + use_mup: bool, +} + +impl MiniCPMDecoderLayer { + pub fn new(vb: VarBuilder, cfg: &VoxMiniCPM4Config) -> Result { + let self_attn = AttentionNobias::new( + vb.pp("self_attn"), + cfg.hidden_size, + cfg.num_attention_heads, + cfg.num_key_value_heads, + )?; + let mlp = MLPNoBias::new( + vb.pp("mlp"), + cfg.hidden_size, + cfg.intermediate_size, + candle_nn::Activation::Silu, + )?; + let input_layernorm = + rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; + let post_attention_layernorm = rms_norm( + cfg.hidden_size, + cfg.rms_norm_eps, + vb.pp("post_attention_layernorm"), + )?; + Ok(Self { + self_attn, + mlp, + input_layernorm, + post_attention_layernorm, + scale_depth: cfg.scale_depth, + num_hidden_layers: cfg.num_hidden_layers, + use_mup: cfg.use_mup, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let xs = self.input_layernorm.forward(xs)?; + let xs = self + .self_attn + .forward(&xs, cos, sin, attention_mask, true)?; + let xs = if self.use_mup { + let res_add = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; + res_add + } else { + let res_add = (residual + xs)?; + res_add + }; + let residual = xs.clone(); + let xs = xs.apply(&self.post_attention_layernorm)?; + let xs = xs.apply(&self.mlp)?; + let xs = if self.use_mup { + let res_add = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; + res_add + } else { + let res_add = (residual + xs)?; + res_add + }; + Ok(xs) + } + + pub fn forward_step( + &mut self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let xs = self.input_layernorm.forward(xs)?; + let xs = self + .self_attn + .forward_step(&xs, cos, sin, attention_mask, true)?; + let xs = if self.use_mup { + let res_add = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; + res_add + } else { + let res_add = (residual + xs)?; + res_add + }; + let residual = &xs; + let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?; + let xs = if self.use_mup { + let res_add = (residual + + xs.affine( + self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), + 0.0, + ))?; + res_add + } else { + let res_add = (residual + xs)?; + res_add + }; + Ok(xs) + } + pub fn clear_kv_cache(&mut self) { + self.self_attn.clear_kv_cache(); + } +} + +pub struct MiniCPMModel { + cfg: VoxMiniCPM4Config, + pub embed_tokens: Option, + layers: Vec, + norm: RmsNorm, + rope_emb: MiniCPMLongRoPE, + // lm_head: Linear, +} + +impl MiniCPMModel { + pub fn new(vb: VarBuilder, cfg: VoxMiniCPM4Config) -> Result { + // let vb = vb.pp("model"); + let embed_tokens = if cfg.vocab_size > 0 { + Some(embedding( + cfg.vocab_size, + cfg.hidden_size, + vb.pp("embed_tokens"), + )?) + } else { + None + }; + + let mut layers = Vec::with_capacity(cfg.num_hidden_layers); + let vb_layers = vb.pp("layers"); + for i in 0..cfg.num_hidden_layers { + let layer = MiniCPMDecoderLayer::new(vb_layers.pp(i), &cfg)?; + layers.push(layer); + } + let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?; + let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device(), vb.dtype())?; + // let lm_head = Linear::new(embed_tokens.embeddings().clone(), None); + Ok(Self { + cfg, + embed_tokens, + layers, + norm, + rope_emb, + // lm_head, + }) + } + + pub fn forward(&mut self, input_embeds: &Tensor, position_id: usize, is_causal: bool) -> Result { + let (bs, seq_len, _) = input_embeds.dims3()?; + // let input_embeds = self + // .embed_tokens + // .forward(&input_ids)? + // .affine(self.cfg.scale_emb, 0.0)?; + let attention_mask: Option<&Tensor> = { + if !is_causal || seq_len <= 1 { + None + } else { + Some(&prepare_causal_attention_mask( + bs, + seq_len, + position_id, + input_embeds.device(), + )?) + } + }; + let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; + let mut hidden_states = input_embeds.clone(); + for decode_layer in &self.layers { + hidden_states = decode_layer.forward(&hidden_states, &cos, &sin, attention_mask)?; + } + hidden_states = self.norm.forward(&hidden_states)?; + Ok(hidden_states) + } + + pub fn forward_step(&mut self, input_embeds: &Tensor, position_id: usize) -> Result { + let input_embeds = match input_embeds.rank() { + 2 => input_embeds.unsqueeze(1)?, + 3 => input_embeds.clone(), + _ => return Err(anyhow!("MiniCPMModelinput_embeds illigal")) + }; + let (bs, seq_len, _) = input_embeds.dims3()?; + // let input_embeds = self + // .embed_tokens + // .forward(&input_ids)? + // .affine(self.cfg.scale_emb, 0.0)?; + let attention_mask: Option<&Tensor> = { + if seq_len <= 1 { + None + } else { + Some(&prepare_causal_attention_mask( + bs, + seq_len, + position_id, + input_embeds.device(), + )?) + } + }; + let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; + let mut hidden_states = input_embeds.clone(); + for decode_layer in &mut self.layers { + hidden_states = + decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?; + } + hidden_states = self.norm.forward(&hidden_states)?; + + Ok(hidden_states) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.layers.iter_mut() { + layer.clear_kv_cache() + } + } +} diff --git a/src/models/voxcpm/mod.rs b/src/models/voxcpm/mod.rs index e69de29..1695295 100644 --- a/src/models/voxcpm/mod.rs +++ b/src/models/voxcpm/mod.rs @@ -0,0 +1,5 @@ +pub mod config; +pub mod audio_vae; +pub mod minicpm4; +pub mod tokenizer; +pub mod model; \ No newline at end of file diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs new file mode 100644 index 0000000..08a48ab --- /dev/null +++ b/src/models/voxcpm/model.rs @@ -0,0 +1,730 @@ +use std::{cmp::max, f64, thread, time}; + +use anyhow::{Ok, Result}; +use candle_core::{D, DType, Device, IndexOp, Tensor}; +use candle_nn::{Linear, Module, VarBuilder, linear, linear_no_bias}; +use candle_transformers::models::deepseek2::SplitOp; + +use crate::{ + models::voxcpm::{ + audio_vae::{self, AudioVAE}, + config::{CfmConfig, VoxCPMConfig, VoxMiniCPM4Config}, + minicpm4::MiniCPMModel, + tokenizer::SingleChineseTokenizer, + }, + utils::{audio_utils::load_audio_with_resample, tensor_utils::linspace}, +}; + +pub struct ScalarQuantizationLayer { + scale: usize, + in_proj: Linear, + out_proj: Linear, +} + +impl ScalarQuantizationLayer { + pub fn new( + vb: VarBuilder, + in_dim: usize, + out_dim: usize, + laten_dim: usize, + scale: usize, + ) -> Result { + let in_proj = linear(in_dim, laten_dim, vb.pp("in_proj"))?; + let out_proj = linear(laten_dim, out_dim, vb.pp("out_proj"))?; + Ok(Self { + scale, + in_proj, + out_proj, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.in_proj.forward(xs)?; + let xs = xs.tanh()?; + let xs = xs + .affine(self.scale as f64, 0.0)? + .round()? + .affine(1.0 / self.scale as f64, 0.0)?; + let xs = self.out_proj.forward(&xs)?; + Ok(xs) + } +} + +pub struct SinusoidalPosEmb { + dim: usize, +} + +impl SinusoidalPosEmb { + pub fn new(dim: usize) -> Result { + assert_eq!(dim % 2, 0, "SinusoidalPosEmb requires dim to be even"); + Ok(Self { dim }) + } + pub fn forward(&self, x: &Tensor, scale: usize) -> Result { + let x = if x.rank() < 1 { + x.unsqueeze(0)? + } else { + x.clone() + }; + let half_dim = self.dim / 2; + let dif = 10000.0_f64.ln() / (half_dim - 1) as f64; + let emb = Tensor::arange(0.0, half_dim as f32, x.device())? + .affine(-1.0 * dif, 0.0)? + .exp()? + .to_dtype(x.dtype())?; + + let emb = x + .unsqueeze(D::Minus1)? + .contiguous()? + .matmul(&emb.unsqueeze(0)?.contiguous()?)? + .affine(scale as f64, 0.0)?; + let emb = Tensor::cat(&[emb.sin()?, emb.cos()?], D::Minus1)?; + Ok(emb) + } +} + +pub struct TimestepEmbedding { + linear_1: Linear, + linear_2: Linear, +} + +impl TimestepEmbedding { + pub fn new( + vb: VarBuilder, + in_channels: usize, + time_embed_dim: usize, + out_dim: Option, + ) -> Result { + let linear_1 = linear(in_channels, time_embed_dim, vb.pp("linear_1"))?; + let time_embed_dim_out = if out_dim.is_some() { + out_dim.unwrap() + } else { + time_embed_dim + }; + let linear_2 = linear(time_embed_dim, time_embed_dim_out, vb.pp("linear_2"))?; + Ok(Self { linear_1, linear_2 }) + } + + pub fn forward(&self, sample: &Tensor) -> Result { + let sample = self.linear_1.forward(&sample)?.silu()?; + let sample = self.linear_2.forward(&sample)?; + Ok(sample) + } +} + +pub struct VoxCPMLocDiT { + in_proj: Linear, + cond_proj: Linear, + out_proj: Linear, + time_embeddings: SinusoidalPosEmb, + time_mlp: TimestepEmbedding, + delta_time_mlp: TimestepEmbedding, + decoder: MiniCPMModel, + config: VoxMiniCPM4Config, + in_channels: usize, +} + +impl VoxCPMLocDiT { + pub fn new(vb: VarBuilder, config: VoxMiniCPM4Config, in_channels: usize) -> Result { + let in_proj = linear(in_channels, config.hidden_size, vb.pp("in_proj"))?; + let cond_proj = linear(in_channels, config.hidden_size, vb.pp("cond_proj"))?; + let out_proj = linear(config.hidden_size, in_channels, vb.pp("out_proj"))?; + let time_embeddings = SinusoidalPosEmb::new(config.hidden_size)?; + let time_mlp = TimestepEmbedding::new( + vb.pp("time_mlp"), + config.hidden_size, + config.hidden_size, + None, + )?; + let delta_time_mlp = TimestepEmbedding::new( + vb.pp("delta_time_mlp"), + config.hidden_size, + config.hidden_size, + None, + )?; + assert_eq!(config.vocab_size, 0, "vocab_size must be 0 for local DiT"); + let decoder = MiniCPMModel::new(vb.pp("decoder"), config.clone())?; + Ok(Self { + in_proj, + cond_proj, + out_proj, + time_embeddings, + time_mlp, + delta_time_mlp, + decoder, + config, + in_channels, + }) + } + + pub fn forward( + &mut self, + x: &Tensor, + mu: &Tensor, + t: &Tensor, + cond: &Tensor, + dt: &Tensor, + ) -> Result { + let x = self.in_proj.forward(&x.transpose(1, 2)?.contiguous()?)?; + let cond = self + .cond_proj + .forward(&cond.transpose(1, 2)?.contiguous()?)?; + let prefix = cond.dims()[1]; + let t = self.time_embeddings.forward(t, 1000)?.to_dtype(x.dtype())?; + let t = self.time_mlp.forward(&t)?; + let dt = self + .time_embeddings + .forward(dt, 1000)? + .to_dtype(x.dtype())?; + let dt = self.delta_time_mlp.forward(&dt)?; + let t = t.add(&dt)?; + + let x = Tensor::cat(&[mu.add(&t)?.unsqueeze(1)?, cond, x], 1)?; + let hidden = self.decoder.forward(&x, 0, false)?; + let select_len = hidden.dims()[1] - (prefix + 1); + let hidden = hidden.narrow(1, prefix + 1, select_len)?; + let hidden = self.out_proj.forward(&hidden)?; + let hidden = hidden.transpose(1, 2)?.contiguous()?; + Ok(hidden) + } +} + +pub struct UnifiedCFM { + solver: String, + sigma_min: f32, + t_scheduler: String, + in_channels: usize, + mean_mode: bool, + estimator: VoxCPMLocDiT, +} + +impl UnifiedCFM { + pub fn new( + in_channels: usize, + cfm_params: CfmConfig, + estimator: VoxCPMLocDiT, + mean_mode: bool, + ) -> Result { + let solver = cfm_params.solver; + let sigma_min = cfm_params.sigma_min; + let t_scheduler = cfm_params.t_scheduler; + Ok(Self { + solver, + sigma_min, + t_scheduler, + in_channels, + mean_mode, + estimator, + }) + } + + pub fn forward( + &mut self, + mu: &Tensor, + n_timesteps: usize, + patch_size: usize, + cond: &Tensor, + temperature: f64, + cfg_value: f64, + sway_sampling_coef: f64, + use_cfg_zero_star: bool, + ) -> Result { + let (b, c) = mu.dims2()?; + let t = patch_size; + let dtype = mu.dtype(); + let z = Tensor::randn(0.0f32, 1.0, (b, self.in_channels, t), mu.device())? + .to_dtype(dtype)? + .affine(temperature, 0.0)?; + println!("z: {}", z); + let t_span = linspace(1.0, 0.0, n_timesteps + 1, mu.device())?.to_dtype(dtype)?; + let t_span = t_span + .affine(f64::consts::PI / 2.0, 0.0)? + .cos()? + .affine(1.0, -1.0)? + .add(&t_span)? + .affine(sway_sampling_coef, 0.0)? + .add(&t_span)?; + println!("t_span: {}", t_span); + println!("mu: {}", mu); + println!("cond: {}", cond); + println!("cfg_value: {}", cfg_value); + println!("use_cfg_zero_star: {}", use_cfg_zero_star); + let x = self.solve_euler(&z, &t_span, mu, cond, cfg_value, use_cfg_zero_star)?; + Ok(x) + } + + pub fn optimized_scale( + &self, + positive_flat: &Tensor, + negative_flat: &Tensor, + ) -> Result { + let dot_product = positive_flat.mul(negative_flat)?.sum_keepdim(1)?; + let squared_norm = negative_flat.powf(2.0)?.sum_keepdim(1)?.affine(1.0, 1e-8)?; + let st_star = dot_product.div(&squared_norm)?; + Ok(st_star) + } + + pub fn solve_euler( + &mut self, + x: &Tensor, + t_span: &Tensor, + mu: &Tensor, + cond: &Tensor, + cfg_value: f64, + use_cfg_zero_star: bool, + ) -> Result { + let mut t = t_span.i(0)?; + let mut dt = t.sub(&t_span.i(1)?)?; + let mut sol = Vec::new(); + let t_span_len = t_span.dims1()?; + let zero_init_steps = max(1, (t_span_len as f32 * 0.04) as usize); + let mut dphi_dt = Tensor::zeros(1, t_span.dtype(), t_span.device())?; + let mut x = x.clone(); + for step in 1..t_span_len { + if use_cfg_zero_star && step <= zero_init_steps { + dphi_dt = Tensor::zeros(1, t_span.dtype(), t_span.device())?; + } else { + let b = x.dim(0)?; + // let x_in = Tensor::zeros((2*b, self.in_channels, x.dim(2)?), x.dtype(), x.device())?; + let x_in = Tensor::cat(&[x.clone(), x.clone()], 0)?; + let mu_in = Tensor::zeros((b, mu.dim(1)?), x.dtype(), x.device())?; + let mu_in = Tensor::cat(&[mu.clone(), mu_in], 0)?; + let t_in = t.broadcast_as(2 * b)?; + let dt_in = if self.mean_mode { + dt.broadcast_as(2 * b)? + } else { + Tensor::zeros(2 * b, x.dtype(), x.device())? + }; + let cond_in = Tensor::cat(&[cond, cond], 0)?; + dphi_dt = self + .estimator + .forward(&x_in, &mu_in, &t_in, &cond_in, &dt_in)?; + let split = dphi_dt.split(&[b, b], 0)?; + dphi_dt = split[0].clone(); + let cfg_dphi_dt = split[1].clone(); + let mut st_star = Tensor::ones(1, x.dtype(), x.device())?; + if use_cfg_zero_star { + let positive_flat = dphi_dt.reshape((b, ()))?; + let negative_flat = cfg_dphi_dt.reshape((b, ()))?; + st_star = self.optimized_scale(&positive_flat, &negative_flat)?; + let mut vec_shape = vec![b]; + let vec_shape1 = vec![1; dphi_dt.rank() - 1]; + vec_shape.extend_from_slice(&vec_shape1); + st_star = st_star.reshape(vec_shape)?; + } + let cfg = cfg_dphi_dt.broadcast_mul(&st_star)?; + dphi_dt = cfg.add(&dphi_dt.sub(&cfg)?.affine(cfg_value, 0.0)?)?; + } + x = x.broadcast_sub(&dphi_dt.broadcast_mul(&dt)?)?; + t = t.sub(&dt)?; + sol.push(x.clone()); + if step < t_span_len - 1 { + dt = t.sub(&t_span.i(step + 1)?)?; + } + } + Ok(sol[sol.len() - 1].clone()) + } +} + +pub struct VoxCPMLocEnc { + special_token: Tensor, + in_proj: Linear, + encoder: MiniCPMModel, + hidden_size: usize, +} + +impl VoxCPMLocEnc { + pub fn new(vb: VarBuilder, config: VoxMiniCPM4Config, input_dim: usize) -> Result { + // let special_token = Tensor::randn(0.0f32, 1.0, (1, 1, 1, config.hidden_size), vb.device())? + // .to_dtype(vb.dtype())?; + let special_token = vb.get((1, 1, 1, config.hidden_size), "special_token")?; + let in_proj = linear(input_dim, config.hidden_size, vb.pp("in_proj"))?; + assert_eq!( + config.vocab_size, 0, + "vocab_size must be 0 for local encoder" + ); + let hidden_size = config.hidden_size; + let encoder = MiniCPMModel::new(vb.pp("encoder"), config)?; + Ok(Self { + special_token, + in_proj, + encoder, + hidden_size, + }) + } + + pub fn forward(&mut self, x: &Tensor) -> Result { + let (b, t, p, d) = x.dims4()?; + let x = self.in_proj.forward(x)?; + println!("VoxCPMLocEnc: in_proj: {}", x); + let special_tokens = self.special_token.expand((b, t, 1, self.hidden_size))?; + let x = Tensor::cat(&[special_tokens, x], 2)?; + println!("VoxCPMLocEnc: cat: {}", x); + let (b, t, p, c) = x.dims4()?; + let x = x.reshape((b * t, p, c))?; + let outputs = self.encoder.forward(&x, 0, false)?; + println!("VoxCPMLocEnc: encoder: {}", outputs); + let cls_output = outputs.i((.., 0, ..))?; + println!("VoxCPMLocEnc: cls_output: {}", cls_output); + let cls_output = cls_output.reshape((b, t, c))?; + Ok(cls_output) + } +} + +pub struct VoxCPMModel { + config: VoxCPMConfig, + patch_size: usize, + audio_start_token: usize, + audio_end_token: usize, + chunk_size: usize, + sample_rate: usize, + tokenizer: SingleChineseTokenizer, + audio_vae: AudioVAE, + base_lm: MiniCPMModel, + residual_lm: MiniCPMModel, + feat_encoder: VoxCPMLocEnc, + feat_decoder: UnifiedCFM, + fsq_layer: ScalarQuantizationLayer, + enc_to_lm_proj: Linear, + lm_to_dit_proj: Linear, + res_to_dit_proj: Linear, + stop_proj: Linear, + stop_head: Linear, + device: Device, + dtype: DType, +} + +impl VoxCPMModel { + pub fn new( + vb: VarBuilder, + config: VoxCPMConfig, + tokenizer: SingleChineseTokenizer, + audio_vae: AudioVAE, + ) -> Result { + let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?; + let audio_start_token = 101usize; + let audio_end_token = 102usize; + let mut residual_lm_config = config.lm_config.clone(); + residual_lm_config.num_hidden_layers = config.residual_lm_num_layers; + residual_lm_config.vocab_size = 0; + let residual_lm = MiniCPMModel::new(vb.pp("residual_lm"), residual_lm_config)?; + let mut encoder_config = config.lm_config.clone(); + encoder_config.hidden_size = config.encoder_config.hidden_dim; + encoder_config.intermediate_size = config.encoder_config.ffn_dim; + encoder_config.num_attention_heads = config.encoder_config.num_heads; + encoder_config.num_hidden_layers = config.encoder_config.num_layers; + encoder_config.vocab_size = 0; + let feat_encoder = + VoxCPMLocEnc::new(vb.pp("feat_encoder"), encoder_config, config.feat_dim)?; + + let mut decoder_config = config.lm_config.clone(); + decoder_config.hidden_size = config.dit_config.hidden_dim; + decoder_config.intermediate_size = config.dit_config.ffn_dim; + decoder_config.num_attention_heads = config.dit_config.num_heads; + decoder_config.num_hidden_layers = config.dit_config.num_layers; + decoder_config.vocab_size = 0; + let estimator = VoxCPMLocDiT::new( + vb.pp("feat_decoder.estimator"), + decoder_config, + config.feat_dim, + )?; + let feat_decoder = UnifiedCFM::new( + config.feat_dim, + config.dit_config.cfm_config.clone(), + estimator, + false, + )?; + let fsq_layer = ScalarQuantizationLayer::new( + vb.pp("fsq_layer"), + config.lm_config.hidden_size, + config.lm_config.hidden_size, + config.scalar_quantization_latent_dim, + config.scalar_quantization_scale, + )?; + let enc_to_lm_proj = linear( + config.encoder_config.hidden_dim, + config.lm_config.hidden_size, + vb.pp("enc_to_lm_proj"), + )?; + let lm_to_dit_proj = linear( + config.lm_config.hidden_size, + config.dit_config.hidden_dim, + vb.pp("lm_to_dit_proj"), + )?; + let res_to_dit_proj = linear( + config.lm_config.hidden_size, + config.dit_config.hidden_dim, + vb.pp("res_to_dit_proj"), + )?; + + let stop_proj = linear( + config.lm_config.hidden_size, + config.lm_config.hidden_size, + vb.pp("stop_proj"), + )?; + let stop_head = linear_no_bias(config.lm_config.hidden_size, 2, vb.pp("stop_head"))?; + + let patch_size = config.patch_size; + Ok(Self { + config, + patch_size, + audio_start_token, + audio_end_token, + chunk_size: audio_vae.chunk_size, + sample_rate: audio_vae.sample_rate, + tokenizer, + audio_vae, + base_lm, + residual_lm, + feat_encoder, + feat_decoder, + fsq_layer, + enc_to_lm_proj, + lm_to_dit_proj, + res_to_dit_proj, + stop_proj, + stop_head, + device: vb.device().clone(), + dtype: vb.dtype(), + }) + } + + pub fn generate( + &mut self, + target_text: String, + prompt_text: Option, + prompt_wav_path: Option, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + retry_badcase: bool, + retry_badcase_max_times: usize, + retry_badcase_ratio_threshold: f64, + ) -> Result { + let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path { + None => { + let text_token = self.tokenizer.encode(target_text.clone())?; + let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; + let audio_start = Tensor::new(vec![self.audio_start_token as u32], &self.device)?; + let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; + let text_length = text_token.dim(0)?; + let audio_feat = Tensor::zeros( + (text_length, self.patch_size, self.audio_vae.latent_dim), + DType::F32, + &self.device, + )?; + let text_mask = Tensor::ones(text_length, self.dtype, &self.device)?; + let audio_mask = Tensor::zeros(text_length, self.dtype, &self.device)?; + (text_token, text_mask, audio_feat, audio_mask) + } + Some(path) => { + let text = prompt_text.unwrap_or("".to_string()) + &target_text; + let text_token = self.tokenizer.encode(text)?; + let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; + let audio_start = Tensor::new(vec![self.audio_start_token as u32], &self.device)?; + let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; + let text_length = text_token.dim(0)?; + let mut audio = + load_audio_with_resample(path, self.device.clone(), 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, + )?; + } + let audio_feat = self.audio_vae.encode(&audio, Some(self.sample_rate))?; + let audio_feat = audio_feat + .reshape((self.audio_vae.latent_dim, (), self.patch_size))? + .permute((1, 2, 0))?; + let dim0 = audio_feat.dim(0)?; + println!("audio_feat: {:?}", audio_feat); + let audio_feat = audio_feat.i(..dim0)?; + println!("audio_feat --: {:?}", audio_feat); + let audio_length = audio_feat.dim(0)?; + let text_pad_token = Tensor::zeros(audio_length, DType::U32, &self.device)?; + let text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?; + let audio_pad_feat = Tensor::zeros( + (text_length, self.patch_size, self.audio_vae.latent_dim), + audio_feat.dtype(), + &self.device, + )?; + let audio_feat = Tensor::cat(&[audio_pad_feat, audio_feat], 0)?; + let text_mask = Tensor::cat( + &[ + Tensor::ones(text_length, self.dtype, &self.device)?, + Tensor::zeros(audio_length, self.dtype, &self.device)?, + ], + D::Minus1, + )?; + let audio_mask = Tensor::cat( + &[ + Tensor::zeros(text_length, self.dtype, &self.device)?, + Tensor::ones(audio_length, self.dtype, &self.device)?, + ], + D::Minus1, + )?; + (text_token, text_mask, audio_feat, audio_mask) + } + }; + + let text_token = text_token.unsqueeze(0)?; + let text_mask = text_mask.unsqueeze(0)?; + let audio_feat = audio_feat.unsqueeze(0)?.to_dtype(DType::BF16)?; + let audio_mask = audio_mask.unsqueeze(0)?; + let target_text_length = self.tokenizer.encode(target_text)?.len(); + let max_len = if retry_badcase { + (target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize + } else { + max_len + }; + let latent_pred = self.inference( + &text_token, + &text_mask, + &audio_feat, + &audio_mask, + min_len, + max_len, + inference_timesteps, + cfg_value, + )?; + let decode_audio = self + .audio_vae + .decode(&latent_pred.to_dtype(DType::F32)?)? + .squeeze(1)?; + let decode_audio_len = decode_audio.dim(D::Minus1)? - 640 - 640; + let decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?; + println!("decode_audio: {}", decode_audio); + Ok(decode_audio) + } + + pub fn inference( + &mut self, + text: &Tensor, + text_mask: &Tensor, + feat: &Tensor, + feat_mask: &Tensor, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + ) -> Result { + println!("text: {}", text); + println!("text_mask: {}", text_mask); + println!("feat: {}", feat); + println!("feat_mask: {}", feat_mask); + let (b, t, p, d) = feat.dims4()?; + let feat_embed = self.feat_encoder.forward(feat)?; // [b, t, h_feat] + println!("feat_embed: {}", feat_embed); + let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?; + println!("feat_embed: {}", feat_embed); + + let scale_emb = if self.config.lm_config.use_mup { + self.config.lm_config.scale_emb + } else { + 1.0 + }; + let text_embed = self + .base_lm + .embed_tokens + .as_ref() + .unwrap() + .forward(text)? + .affine(scale_emb as f64, 0.0)?; + println!("text_embed: {}", text_embed); + let combined_embed = text_mask + .unsqueeze(D::Minus1)? + .broadcast_mul(&text_embed)? + .add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?; + println!("combined_embed: {}", combined_embed); + + let mut prefix_feat_cond = feat.i((.., t - 1, ..))?; + let mut pred_feat_seq = Vec::new(); + let mut position_id = 0; + let mut seq_len = t; + let enc_outputs = self.base_lm.forward_step(&combined_embed, position_id)?; + println!("base_lm enc_outputs: {}", enc_outputs); + let enc_outputs = self + .fsq_layer + .forward(&enc_outputs)? + .broadcast_mul(&feat_mask.unsqueeze(D::Minus1)?)? + .add(&enc_outputs.broadcast_mul(&text_mask.unsqueeze(D::Minus1)?)?)?; + println!("fsq_layer enc_outputs: {}", enc_outputs); + let mut lm_hidden = enc_outputs.i((.., t - 1, ..))?; + println!("lm_hidden shape: {:?}", lm_hidden); + + let input_embeds = + enc_outputs.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?; + let residual_enc_outputs = self.residual_lm.forward_step(&input_embeds, position_id)?; + println!("residual_lm residual_enc_outputs: {}", residual_enc_outputs); + let mut residual_hidden = residual_enc_outputs.i((.., t - 1, ..))?; + + for i in 0..max_len { + let dit_hidden_1 = self.lm_to_dit_proj.forward(&lm_hidden)?; // [b, h_dit] + println!("dit_hidden_1: {}", dit_hidden_1); + let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_hidden)?; // [b, h_dit] + println!("dit_hidden_2: {}", dit_hidden_2); + let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?; + println!("dit_hidden: {}", dit_hidden); + let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?; + + let pred_feat = self + .feat_decoder + .forward( + &dit_hidden, + inference_timesteps, + self.patch_size, + &cond, + 1.0, + cfg_value, + 1.0, + true, + )? + .transpose(1, 2)?; // [b, p, d] + println!("pred_feat: {}", pred_feat); + + let curr_embed = self.feat_encoder.forward(&pred_feat.unsqueeze(1)?)?; // [b, 1, c] + let curr_embed = self.enc_to_lm_proj.forward(&curr_embed)?; + println!("curr_embed: {}", curr_embed); + pred_feat_seq.push(pred_feat.unsqueeze(1)?); + + prefix_feat_cond = pred_feat; + println!("lm_hidden: {}", lm_hidden); + let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?; + println!("stop_flag: {}", stop_flag); + let stop_flag = self + .stop_head + .forward(&stop_flag)? + .argmax(D::Minus1)? + .i(0)? + .to_scalar::()?; + println!("i: {}, stop_flag: {}", i, stop_flag); + if i > min_len && stop_flag == 1 { + break; + } + position_id += seq_len; + seq_len = 1; + lm_hidden = self + .base_lm + .forward_step(&curr_embed.i((.., 0, ..))?, position_id)? + .squeeze(1)?; + lm_hidden = self.fsq_layer.forward(&lm_hidden)?; + residual_hidden = self + .residual_lm + .forward_step(&lm_hidden.add(&curr_embed.i((.., 0, ..))?)?, position_id)? + .squeeze(1)?; + } + let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d) + let (b, t, p, d) = pred_seq.dims4()?; + println!("pred_seq: {:?}", pred_seq); + let feat_pred = pred_seq + .permute((0, 3, 1, 2))? + .reshape((b, d, ()))? + .contiguous()?; + println!("feat_pred: {:?}", feat_pred); + self.base_lm.clear_kv_cache(); + self.residual_lm.clear_kv_cache(); + + Ok(feat_pred) + } +} diff --git a/src/models/voxcpm/tokenizer.rs b/src/models/voxcpm/tokenizer.rs new file mode 100644 index 0000000..ca0ee8c --- /dev/null +++ b/src/models/voxcpm/tokenizer.rs @@ -0,0 +1,66 @@ +use anyhow::{Ok, Result, anyhow}; +use candle_core::Tensor; +use tokenizers::Tokenizer; + +pub struct SingleChineseTokenizer { + tokenizer: Tokenizer, + multichar_tokens: Vec, +} + +impl SingleChineseTokenizer { + pub fn new(path: &str) -> Result { + let path = path.to_string(); + assert!( + std::path::Path::new(&path).exists(), + "model path file not exists" + ); + let tokenizer_file = path.clone() + "/tokenizer.json"; + assert!( + std::path::Path::new(&tokenizer_file).exists(), + "tokenizer.json not exists in model path" + ); + let tokenizer = Tokenizer::from_file(tokenizer_file) + .map_err(|e| anyhow!(format!("tokenizer from file error{}", e)))?; + let mut multichar_tokens = Vec::new(); + for (token, _) in tokenizer.get_vocab(false) { + let len = token.chars().count(); + if len >= 2 { + let is_chinese = token.chars().all(|c| { + let c_ = c as u32; + 0x4E00 <= c_ && c_ <= 0x9FFF + }); + if is_chinese { + multichar_tokens.push(token); + } + } + } + Ok(Self { + tokenizer, + multichar_tokens, + }) + } + pub fn encode(&self, text: String) -> Result> { + let encode = self + .tokenizer + .encode(text, false) + .map_err(|e| anyhow!(format!("tokenizer encode error: {}", e)))?; + let tokens = encode.get_tokens(); + println!("tokens: {:?}", tokens); + let mut split_character = Vec::new(); + for token in tokens { + let clean_token = token.replace("▁", "to"); + if self.multichar_tokens.contains(&clean_token) { + let chars: Vec = clean_token.chars().map(|c| c.to_string()).collect(); + split_character.extend(chars); + } else { + split_character.push(token.clone()); + } + } + println!("split_character: {:?}", split_character); + let ids: Vec = split_character + .iter() + .filter_map(|c| self.tokenizer.token_to_id(c)) + .collect(); + Ok(ids) + } +} diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index cc70efc..c60f8e8 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -64,6 +64,7 @@ pub fn apply_rotary_pos_emb_vision( // cos, sin -> (seq_len, head_dim) -> (seq_len, 1, head_dim) let cos = cos.unsqueeze(D::Minus2)?; let sin = sin.unsqueeze(D::Minus2)?; + let q_embed = q .broadcast_mul(&cos)? .add(&rotate_half(q)?.broadcast_mul(&sin)?)?; @@ -78,15 +79,42 @@ pub fn apply_rotary_pos_emb( k: &Tensor, cos: &Tensor, sin: &Tensor, + tof32: bool, ) -> Result<(Tensor, Tensor)> { - // sin/cos: (bs, 1, seq_len, head_dim) + // sin/cos: to (bs, 1, seq_len, head_dim) // q/k: (bs, n_head, seq_len, head_dim) + let mut cos = cos.clone(); + let mut sin = sin.clone(); + if cos.rank() == 2 { + // (seq_len, head_dim) -> (1, 1, seq_len, head_dim) + cos = cos.unsqueeze(0)?.unsqueeze(0)?; + sin = sin.unsqueeze(0)?.unsqueeze(0)?; + } + if cos.rank() == 3 { + // (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim) + cos = cos.unsqueeze(1)?; + sin = sin.unsqueeze(1)?; + } + let orig_dtype = q.dtype(); + let q = if tof32 { + &q.to_dtype(DType::F32)? + } else { + q + }; + let k = if tof32 { + &k.to_dtype(DType::F32)? + } else { + k + }; + let cos = cos.to_dtype(q.dtype())?; + let sin = sin.to_dtype(q.dtype())?; + let q_embed = q .broadcast_mul(&cos)? - .add(&rotate_half(q)?.broadcast_mul(&sin)?)?; + .add(&rotate_half(q)?.broadcast_mul(&sin)?)?.to_dtype(orig_dtype)?; let k_embed = k .broadcast_mul(&cos)? - .add(&rotate_half(k)?.broadcast_mul(&sin)?)?; + .add(&rotate_half(k)?.broadcast_mul(&sin)?)?.to_dtype(orig_dtype)?; Ok((q_embed, k_embed)) } diff --git a/src/tokenizer/tokenizer.rs b/src/tokenizer/tokenizer.rs index 8040a76..38157a4 100644 --- a/src/tokenizer/tokenizer.rs +++ b/src/tokenizer/tokenizer.rs @@ -41,4 +41,5 @@ impl TokenizerModel { .map_err(|e| anyhow!(format!("tokenizer encode error{}", e)))?; Ok(decode) } + } diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index 8b13789..40c5cd3 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -1 +1,279 @@ +use anyhow::{Result, anyhow}; +use candle_core::{D, DType, Device, Tensor}; +use candle_nn::{conv1d_no_bias, Conv1d, Conv1dConfig, Module}; +use hound::{SampleFormat, WavReader}; +use rocket::futures::future::ok; +use rubato::{ + Resampler, SincFixedIn, SincInterpolationParameters, SincInterpolationType, WindowFunction, +}; +use std::f64::consts::PI; +use std::path::Path; +// 重采样方法枚举 +#[derive(Debug, Clone, Copy)] +pub enum ResamplingMethod { + SincInterpHann, + SincInterpKaiser, +} + +// 计算最大公约数 +fn gcd(a: i64, b: i64) -> i64 { + if b == 0 { a } else { gcd(b, a % b) } +} + +// 零阶修正贝塞尔函数 I0 +fn i0(x: f32) -> f32 { + let mut result = 1.0; + let mut term = 1.0; + let half_x_sq = x * x / 4.0; + + for k in 1..50 { + term = term * half_x_sq / (k * k) as f32; + result += term; + + if term < 1e-12 { + break; + } + } + + result +} + +// 获取sinc重采样核 +pub fn get_sinc_resample_kernel( + orig_freq: i64, + new_freq: i64, + gcd_val: i64, + lowpass_filter_width: i64, + rolloff: f64, + resampling_method: ResamplingMethod, + beta: Option, + device: &Device, +) -> Result<(Tensor, i64)> { + if orig_freq <= 0 || new_freq <= 0 { + return Err(anyhow!("Frequencies must be positive".to_string())); + } + + if lowpass_filter_width <= 0 { + return Err(anyhow!( + "Low pass filter width should be positive".to_string() + )); + } + + let orig_freq = orig_freq / gcd_val; + let new_freq = new_freq / gcd_val; + + let base_freq = (orig_freq.min(new_freq) as f64) * rolloff; + + let width_f = (lowpass_filter_width as f64) * (orig_freq as f64) / base_freq; + let width = width_f.ceil() as i64; + // 创建索引数组 [1, 1, 2*width + orig_freq_reduced] + let idx = Tensor::arange(-width as f32, (width + orig_freq) as f32, device)? + .affine(1.0 / orig_freq as f64, 0.0)? + .unsqueeze(0)? + .unsqueeze(0)?; + // 创建时间数组 t [new_freq_reduced, 1, idx_len] + let t = Tensor::arange_step(0.0, -new_freq as f32, -1.0, device)? + .affine(1.0 / new_freq as f64, 0.0)? + .unsqueeze(D::Minus1)? + .unsqueeze(D::Minus1)? + .broadcast_add(&idx)? + .affine(base_freq, 0.0)?; + let t = t.clamp(-lowpass_filter_width as f32, lowpass_filter_width as f32)?; + // 计算窗口函数 + let window = match resampling_method { + ResamplingMethod::SincInterpHann => { + let window_arg = t.affine(PI / (lowpass_filter_width as f64) / 2.0, 0.0)?; + window_arg.cos()?.sqr()? + } + ResamplingMethod::SincInterpKaiser => { + let beta_val = beta.unwrap_or(14.769656459379492); + let i0_beta = i0(beta_val); + + let normalized_t = t.affine(1.0 / lowpass_filter_width as f64, 0.0)?; + let arg = (1.0 - normalized_t.sqr()?)?; + // 处理arg为负数的情况 + let sqrt_arg = arg.relu()?.sqrt()?; + let sqrt_dims = sqrt_arg.dims(); + let sqrt_arg_vec = sqrt_arg.flatten_all()?.to_vec1::()?; + + let window_val:Vec = sqrt_arg_vec.iter().map(|x| i0(beta_val * x) / i0_beta).collect(); + let window = Tensor::new(window_val, device)?.reshape(sqrt_dims)?; + window + } + }; + + // 计算sinc核 + let scale = base_freq / (orig_freq as f64); + let t_scaled = t.affine(PI, 0.0)?; + + let t_zeros = Tensor::zeros_like(&t_scaled)?; + let t_ones = Tensor::ones_like(&t_scaled)?; + let mask = t_scaled.eq(&t_zeros)?; + let sinc = mask.where_cond(&t_ones, &t_scaled.sin()?.div(&t_scaled)?)?; + let kernels = sinc.mul(&window)?.affine(scale, 0.0)?; + + Ok((kernels, width)) +} + +// 应用sinc重采样核 +pub fn apply_sinc_resample_kernel( + waveform: &Tensor, + orig_freq: i64, + new_freq: i64, + gcd_val: i64, + kernel: &Tensor, + width: i64, +) -> Result { + let orig_freq = orig_freq / gcd_val; + let new_freq = new_freq / gcd_val; + + // 获取波形形状 + let dims = waveform.dims(); + let waveform_flat = waveform.reshape(((), dims[dims.len()-1]))?; + + let (num_wavs, length) = waveform_flat.dims2()?; + let padded_waveform = waveform.pad_with_zeros(D::Minus1, width as usize, (width+orig_freq) as usize)?; + + // 添加通道维度 [batch_size, 1, padded_length] + let waveform_3d = padded_waveform.unsqueeze(1)?; + let config = Conv1dConfig { + padding: 0, + stride: orig_freq as usize, + dilation: 1, + groups: 1, + cudnn_fwd_algo: None, + }; + + let conv1d = Conv1d::new(kernel.clone(), None, config); + // 执行卷积 + // kernel形状: [new_freq_reduced, 1, kernel_len] + // 输出形状: [batch_size, new_freq_reduced, output_length] + let conv_output = conv1d.forward(&waveform_3d)?; + + // 转置并重塑 [batch_size, output_length * new_freq_reduced] + let conv_transposed = conv_output.transpose(1, 2)?.reshape((num_wavs, ()))?; + + // 计算目标长度 + let target_length = + ((new_freq as f64 * length as f64) / orig_freq as f64).ceil() as usize; + + // 截取目标长度 + let resampled_flat = + conv_transposed.narrow(1, 0, target_length.min(conv_transposed.dim(1)?))?; + let mut new_dims = dims.to_vec(); + let last_dim = new_dims.len()-1; + new_dims[last_dim] = resampled_flat.dim(1)?; + // 恢复原始批次形状 + + let resampled = resampled_flat.reshape(new_dims)?; + + Ok(resampled) +} + +// 主要的重采样函数 +pub fn resample( + waveform: &Tensor, + orig_freq: i64, + new_freq: i64, + lowpass_filter_width: i64, + rolloff: f64, + resampling_method: ResamplingMethod, + beta: Option, +) -> Result { + if orig_freq <= 0 || new_freq <= 0 { + return Err(anyhow!( + "Frequencies must be positive".to_string(), + )); + } + + if orig_freq == new_freq { + return Ok(waveform.clone()); + } + + let gcd_val = gcd(orig_freq, new_freq); + let device = waveform.device(); + + let (kernel, width) = get_sinc_resample_kernel( + orig_freq, + new_freq, + gcd_val, + lowpass_filter_width, + rolloff, + resampling_method, + beta, + &device, + )?; + let t = apply_sinc_resample_kernel(waveform, orig_freq, new_freq, gcd_val, &kernel, width)?; + Ok(t) +} + +// 为方便使用提供的简化版本 +pub fn resample_simple(waveform: &Tensor, orig_freq: i64, new_freq: i64) -> Result { + resample( + waveform, + orig_freq, + new_freq, + 6, + 0.99, + ResamplingMethod::SincInterpHann, + None, + ) +} + +pub fn load_audio>(path: P, device: Device) -> Result<(Tensor, usize)> { + let mut reader = WavReader::open(path)?; + let spec = reader.spec(); + let samples: Vec = match spec.sample_format { + SampleFormat::Int => { + // 将整数样本转换为浮点数 [-1.0, 1.0] + let max_value = match spec.bits_per_sample { + 8 => i8::MAX as f32, + 16 => i16::MAX as f32, + 24 => 8388607.0, + _ => { + return Err(anyhow::anyhow!( + "Unsupported bit depth: {}", + spec.bits_per_sample + )); + } + }; + reader + .samples::() + .map(|s| s.map(|sample| sample as f32 / max_value)) + .collect::, _>>()? + } + SampleFormat::Float => { + // 直接读取浮点数样本 + reader.samples::().collect::, _>>()? + } + }; + let sample_rate = spec.sample_rate; + let mut audio_tensor = Tensor::from_slice( + &samples, + ( + samples.len() / spec.channels as usize, + spec.channels as usize, + ), + &device, + )? + .t()?; + if spec.channels > 1 { + // 对channel通道求平均, channel维度变为1 + audio_tensor = audio_tensor.mean_keepdim(0)?; + } + Ok((audio_tensor, sample_rate as usize)) +} + +pub fn load_audio_with_resample>( + path: P, + device: Device, + target_sample_rate: Option, +) -> Result { + let (mut audio, sr) = load_audio(path, device)?; + if target_sample_rate.is_some() && target_sample_rate.unwrap() as usize != sr { + let target_sample_rate = target_sample_rate.unwrap(); + audio = resample_simple(&audio, sr as i64, target_sample_rate as i64)?; + } + Ok(audio) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 70da618..c313084 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -2,3 +2,4 @@ pub mod img_utils; pub mod tensor_utils; pub mod utils; pub mod video_utils; +pub mod audio_utils; \ No newline at end of file diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index 41a851c..f715399 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -1,4 +1,4 @@ -use anyhow::{Result, anyhow}; +use anyhow::{anyhow, Ok, Result}; use candle_core::{D, DType, Device, IndexOp, Tensor, shape::Dim}; pub fn prepare_causal_attention_mask( @@ -249,3 +249,18 @@ pub fn get_vision_next_indices(input_ids: &Tensor, token_id: u32) -> Result Result { + assert!(steps > 0, "steps must be > 0"); + if steps == 1 { + let t = Tensor::from_slice(&[start], 1, device)?; + return Ok(t); + } + let step_size = (end - start) / (steps-1) as f32; + let data: Vec = (0..steps) + .map(|i| start + i as f32 * step_size) + .collect(); + + let t = Tensor::from_slice(&data, steps, device)?; + Ok(t) +} \ No newline at end of file diff --git a/src/utils/utils.rs b/src/utils/utils.rs index 2b8a70b..62edd7e 100644 --- a/src/utils/utils.rs +++ b/src/utils/utils.rs @@ -61,7 +61,7 @@ pub fn string_to_static_str(s: String) -> &'static str { Box::leak(s.into_boxed_str()) } -pub fn find_safetensors_files(path: &str) -> Result> { +pub fn find_type_files(path: &str, extension_type: &str) -> Result> { let mut files = Vec::new(); for entry in std::fs::read_dir(path)? { @@ -70,7 +70,7 @@ pub fn find_safetensors_files(path: &str) -> Result> { if file_path.is_file() { if let Some(extension) = file_path.extension() { - if extension == "safetensors" { + if extension == extension_type { files.push(file_path.to_string_lossy().to_string()); } } @@ -80,6 +80,7 @@ pub fn find_safetensors_files(path: &str) -> Result> { Ok(files) } + pub fn round_by_factor(num: u32, factor: u32) -> u32 { let round = (num as f32 / factor as f32).round() as u32; round * factor diff --git a/tests/config_tests.rs b/tests/config_tests.rs index 3c30f03..68366a3 100644 --- a/tests/config_tests.rs +++ b/tests/config_tests.rs @@ -1,4 +1,4 @@ -use aha::models::{minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig}; +use aha::models::{minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig, voxcpm::config::VoxCPMConfig}; use anyhow::Result; #[test] @@ -20,4 +20,14 @@ fn minicpm4_config() -> Result<()> { let config: MiniCPM4Config = serde_json::from_slice(&std::fs::read(config_path)?)?; println!("{:?}", config); Ok(()) +} + +#[test] +fn voxcpm_config() -> Result<()> { + // cargo test -F cuda,flash-attn minicpm4_config -- --nocapture + let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; + let config_path = model_path.to_string() + "/config.json"; + let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + println!("{:?}", config); + Ok(()) } \ No newline at end of file diff --git a/tests/messy_test.rs b/tests/messy_test.rs new file mode 100644 index 0000000..4b10d68 --- /dev/null +++ b/tests/messy_test.rs @@ -0,0 +1,20 @@ +use aha::utils::audio_utils::{load_audio_with_resample}; +use anyhow::Result; +use candle_core::Tensor; + +#[test] +fn messy_test() -> Result<()> { + let device = candle_core::Device::Cpu; + let wav_path = "./assets/audio/example.wav"; + let audio_tensor = load_audio_with_resample(wav_path, device,Some(16000))?; + + println!("audio_tensor: {}", audio_tensor); + // let string = "你好啊".to_string(); + // let vec_str: Vec= string.chars().map(|c| c.to_string()).collect(); + // println!("vec_str: {:?}", vec_str); + // let t = Tensor::rand(-1.0, 1.0, (2, 2), &device)?; + // println!("t: {}", t); + // let re_t = t.recip()?; + // println!("re_t: {}", re_t); + Ok(()) +} diff --git a/tests/test_minicpm4.rs b/tests/test_minicpm4.rs index 74bebd4..f8a4308 100644 --- a/tests/test_minicpm4.rs +++ b/tests/test_minicpm4.rs @@ -1,39 +1,34 @@ -use std::time::Instant; +use std::{pin::pin, time::Instant}; +use aha::models::{minicpm4::generate::MiniCPMGenerateModel, GenerateModel}; use anyhow::Result; use candle_core::{DType, Device}; use openai_dive::v1::resources::chat::ChatCompletionParameters; +use rocket::futures::StreamExt; #[test] -fn qwen2_5vl_generate() -> Result<()> { - // test with cpu :(太慢了, : RUST_BACKTRACE=1 cargo test qwen2_5vl_generate -- --nocapture - // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen2_5vl_generate -- --nocapture - // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen2_5vl_generate -- --nocapture - let device = Device::cuda_if_available(0)?; - let dtype = DType::BF16; - +fn minicpm_generate() -> Result<()> { + // test with cpu :(太慢了, : RUST_BACKTRACE=1 cargo test minicpm_generate -- --nocapture + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda minicpm_generate -- --nocapture + // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn minicpm_generate -- --nocapture + let model_path = "/home/jhq/huggingface_model/OpenBMB/MiniCPM4-0.5B/"; - let message = r#" { + "temperature": 0.3, + "top_p": 0.8, "model": "minicpm4", "messages": [ { "role": "user", - "content": [ - { - "type": "text", - "text": "你是谁" - } - ] + "content": "贾宝玉和孙悟空有什么关系" } ] } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - // let mut model = Qwen2_5VLGenerateModel::init(model_path, &device, dtype)?; - let mut model = ModelType::init(ModelType::Qwen2_5VL, model_path, None, None)?; + let mut model = MiniCPMGenerateModel::init(model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); @@ -47,40 +42,25 @@ fn qwen2_5vl_generate() -> Result<()> { } #[tokio::test] -async fn qwen2_5vl_stream() -> Result<()> { - // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen2_5vl_generate -- --nocapture - let device = Device::cuda_if_available(0)?; - let dtype = DType::BF16; - - let model_path = "/home/jhq/huggingface_model/Qwen/Qwen2.5-VL-3B-Instruct/"; +async fn minicpm_stream() -> Result<()> { + // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn minicpm_stream -- --nocapture + + let model_path = "/home/jhq/huggingface_model/OpenBMB/MiniCPM4-0.5B/"; let message = r#" { - "model": "qwen2.5vl", + "model": "minicpm4", "messages": [ { "role": "user", - "content": [ - { - "type": "image", - "image_url": - { - "url": "file://./assets/img/ocr_test.png" - } - }, - { - "type": "text", - "text": "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本" - } - ] + "content": "你是谁" } ] } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - // let mut model = Qwen2_5VLGenerateModel::init(model_path, &device, dtype)?; - let mut model = ModelType::init(ModelType::Qwen2_5VL, model_path, None, None)?; + let mut model = MiniCPMGenerateModel::init(model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); diff --git a/tests/test_qwen2_5vl.rs b/tests/test_qwen2_5vl.rs index 1c797e3..ce798ef 100644 --- a/tests/test_qwen2_5vl.rs +++ b/tests/test_qwen2_5vl.rs @@ -1,7 +1,6 @@ use std::{pin::pin, time::Instant}; use aha::{ - ModelType, models::{GenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel}, }; use anyhow::Result; @@ -14,8 +13,8 @@ fn qwen2_5vl_generate() -> Result<()> { // test with cpu :(太慢了, : RUST_BACKTRACE=1 cargo test qwen2_5vl_generate -- --nocapture // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen2_5vl_generate -- --nocapture // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen2_5vl_generate -- --nocapture - let device = Device::cuda_if_available(0)?; - let dtype = DType::BF16; + // let device = Device::cuda_if_available(0)?; + // let dtype = DType::BF16; let model_path = "/home/jhq/huggingface_model/Qwen/Qwen2.5-VL-3B-Instruct/"; @@ -30,7 +29,7 @@ fn qwen2_5vl_generate() -> Result<()> { "type": "image", "image_url": { - "url": "file://./assets/img/ocr_test.png" + "url": "file://./assets/img/ocr_test1.png" } }, { @@ -44,8 +43,7 @@ fn qwen2_5vl_generate() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - // let mut model = Qwen2_5VLGenerateModel::init(model_path, &device, dtype)?; - let mut model = ModelType::init(ModelType::Qwen2_5VL, model_path, None, None)?; + let mut model = Qwen2_5VLGenerateModel::init(model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); @@ -61,8 +59,8 @@ fn qwen2_5vl_generate() -> Result<()> { #[tokio::test] async fn qwen2_5vl_stream() -> Result<()> { // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen2_5vl_generate -- --nocapture - let device = Device::cuda_if_available(0)?; - let dtype = DType::BF16; + // let device = Device::cuda_if_available(0)?; + // let dtype = DType::BF16; let model_path = "/home/jhq/huggingface_model/Qwen/Qwen2.5-VL-3B-Instruct/"; @@ -91,8 +89,7 @@ async fn qwen2_5vl_stream() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - // let mut model = Qwen2_5VLGenerateModel::init(model_path, &device, dtype)?; - let mut model = ModelType::init(ModelType::Qwen2_5VL, model_path, None, None)?; + let mut model = Qwen2_5VLGenerateModel::init(model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs new file mode 100644 index 0000000..af8f617 --- /dev/null +++ b/tests/test_voxcpm.rs @@ -0,0 +1,57 @@ +use std::collections::HashMap; +use anyhow::{Ok, Result}; + +use aha::{models::voxcpm::{audio_vae::AudioVAE, config::VoxCPMConfig, model::VoxCPMModel, tokenizer::SingleChineseTokenizer}, utils::utils::{find_type_files, get_device}}; +use candle_core::pickle::read_all_with_key; +use candle_nn::VarBuilder; + + +#[test] +fn voxcpm_generate() -> Result<()> { + let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; + let model_list = find_type_files(&model_path, "pth")?; + println!(" pth model_list: {:?}", model_list); + let dev = get_device(None); + let mut dict_to_hashmap = HashMap::new(); + let mut dtype = candle_core::DType::F32; + for m in model_list { + let dict = read_all_with_key(m, Some("state_dict"))?; + dtype = dict[0].1.dtype(); + for (k, v) in dict { + // println!("key: {}, tensor shape: {:?}", k, v); + dict_to_hashmap.insert(k, v); + } + } + let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev); + let audio_vae = AudioVAE::new(vb, 128, vec![2, 5, 8, 8], Some(64), 1536, vec![8, 8, 5, 2], 16000)?; + println!("audio vae load down"); + let model_list = find_type_files(&model_path, "bin")?; + println!(" bin model_list: {:?}", model_list); + dict_to_hashmap = HashMap::new(); + for m in model_list { + let dict = read_all_with_key(m, Some("state_dict"))?; + dtype = dict[0].1.dtype(); + for (k, v) in dict { + // println!("key: {}, tensor shape: {:?}", k, v); + dict_to_hashmap.insert(k, v); + } + } + let vb_vox = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev); + let config_path = model_path.to_string() + "/config.json"; + let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + let tokenizer = SingleChineseTokenizer::new(model_path)?; + let mut voxcpm = VoxCPMModel::new(vb_vox, config, tokenizer, audio_vae)?; + let generate = voxcpm.generate("你好啊,这是初始测试语句".to_string(), None, None, 2, 30, 10, 2.0, false, 3, 6.0)?; + // let audio_path = "./assets/audio/example.wav"; + + Ok(()) +} + +#[test] +fn voxcpm_tokenizer() -> Result<()> { + let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; + let tokenizer = SingleChineseTokenizer::new(model_path)?; + let ids = tokenizer.encode("你好啊,你吃饭了吗".to_string())?; + println!("ids: {:?}", ids); + Ok(()) +} \ No newline at end of file diff --git a/tests/weight_test.rs b/tests/weight_test.rs index dd8761b..117d2fa 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -1,10 +1,14 @@ -use aha::utils::utils::find_safetensors_files; +use std::collections::HashMap; + +use aha::utils::utils::{find_type_files, get_device}; use anyhow::Result; -use candle_core::{safetensors, Device}; +use candle_core::{pickle::{read_all_with_key, read_pth_tensor_info, PthTensors}, safetensors, Device, Tensor}; +use candle_nn::VarBuilder; + #[test] fn minicpm4_weight() -> Result<()> { let model_path = "/home/jhq/huggingface_model/OpenBMB/MiniCPM4-0.5B/"; - let model_list = find_safetensors_files(&model_path)?; + let model_list = find_type_files(&model_path, "safetensors")?; let device = Device::Cpu; for m in model_list { let weights = safetensors::load(m, &device)?; @@ -15,4 +19,26 @@ fn minicpm4_weight() -> Result<()> { } } Ok(()) +} + +#[test] +fn voxcpm_weight() -> Result<()> { + let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; + let model_list = find_type_files(&model_path, "pth")?; + println!("model_list: {:?}", model_list); + let dev = get_device(None); + let mut dict_to_hashmap = HashMap::new(); + let mut dtype = candle_core::DType::F16; + for m in model_list { + let dict = read_all_with_key(m, Some("state_dict"))?; + dtype = dict[0].1.dtype(); + for (k, v) in dict { + println!("key: {}, tensor shape: {:?}", k, v); + dict_to_hashmap.insert(k, v); + } + } + let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev); + let contain_key = vb.contains_tensor("encoder.block.4.block.2.block.3.weight_g"); + println!("contain encoder.block.4.block.2.block.3.weight_g: {}", contain_key); + Ok(()) } \ No newline at end of file