add voxcpm with some bug
This commit is contained in:
Generated
+68
@@ -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"
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Binary file not shown.
Binary file not shown.
+23
-23
@@ -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<DType>,
|
||||
) -> Result<Box<dyn GenerateModel>> {
|
||||
match model_type {
|
||||
ModelType::Qwen2_5VL => {
|
||||
let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?;
|
||||
Ok(Box::new(model) as Box<dyn GenerateModel>)
|
||||
},
|
||||
ModelType::MiniCPM4 => {
|
||||
let model = MiniCPMGenerateModel::init(model_path, device, dtype)?;
|
||||
Ok(Box::new(model)as Box<dyn GenerateModel>)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// impl ModelType {
|
||||
// pub fn init(
|
||||
// model_type: ModelType,
|
||||
// model_path: &str,
|
||||
// device: Option<&Device>,
|
||||
// dtype: Option<DType>,
|
||||
// ) -> Result<Box<dyn GenerateModel>> {
|
||||
// match model_type {
|
||||
// ModelType::Qwen2_5VL => {
|
||||
// let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?;
|
||||
// Ok(Box::new(model) as Box<dyn GenerateModel>)
|
||||
// },
|
||||
// ModelType::MiniCPM4 => {
|
||||
// let model = MiniCPMGenerateModel::init(model_path, device, dtype)?;
|
||||
// Ok(Box::new(model)as Box<dyn GenerateModel>)
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
|
||||
@@ -116,6 +116,7 @@ impl AttentionNobias {
|
||||
cos: &Tensor,
|
||||
sin: &Tensor,
|
||||
attention_mask: Option<&Tensor>,
|
||||
tof32: bool,
|
||||
) -> Result<Tensor> {
|
||||
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<Tensor> {
|
||||
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)) => {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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<DType>) -> Result<Self> {
|
||||
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)?;
|
||||
|
||||
@@ -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<f32>,
|
||||
long_factor: Vec<f32>,
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
pub fn forward(&mut self, input_ids: &Tensor, position_id: usize) -> Result<Tensor> {
|
||||
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
|
||||
@@ -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<Tensor> {
|
||||
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)
|
||||
}
|
||||
|
||||
+1
-4
@@ -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<DType>) -> Result<Self>
|
||||
// where
|
||||
// Self: Sized;
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse>;
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)) => {
|
||||
|
||||
@@ -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<Tensor>,
|
||||
// in_c: usize,
|
||||
// out_c: usize,
|
||||
// kernel_size: usize,
|
||||
padding: usize,
|
||||
dilation: usize,
|
||||
groups: usize,
|
||||
stride: usize,
|
||||
) -> Result<Self> {
|
||||
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<Tensor> {
|
||||
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<Tensor>,
|
||||
padding: usize,
|
||||
dilation: usize,
|
||||
output_padding: usize,
|
||||
groups: usize,
|
||||
stride: usize,
|
||||
) -> Result<Self> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
let x = self.conv_transpose.forward(x)?;
|
||||
Ok(x)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct Snake1d {
|
||||
alpha: Tensor,
|
||||
}
|
||||
impl Snake1d {
|
||||
pub fn new(vb: VarBuilder, channels: usize) -> Result<Self> {
|
||||
let alpha = vb.get((1, channels, 1), "alpha")?;
|
||||
Ok(Self { alpha })
|
||||
}
|
||||
|
||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<usize>,
|
||||
out_dim: usize,
|
||||
stride: usize,
|
||||
groups: usize,
|
||||
) -> Result<Self> {
|
||||
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<Tensor> {
|
||||
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<CausalEncoderBlock>,
|
||||
fc_mu: WNCausalConv1d,
|
||||
fc_logvar: WNCausalConv1d,
|
||||
}
|
||||
|
||||
impl CausalEncoder {
|
||||
pub fn new(
|
||||
vb: VarBuilder,
|
||||
d_model: usize,
|
||||
laten_dim: usize,
|
||||
strides: Vec<usize>,
|
||||
depthwise: bool,
|
||||
) -> Result<Self> {
|
||||
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<Self> {
|
||||
let linear = WNCausalConv1d::new(vb.pp("linear"), dim, dim, 1, 1, 0, 1, 1)?;
|
||||
Ok(Self { linear })
|
||||
}
|
||||
|
||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<CausalDecoderBlock>,
|
||||
model6: Snake1d,
|
||||
model7: WNCausalConv1d,
|
||||
}
|
||||
|
||||
impl CausalDecoder {
|
||||
pub fn new(
|
||||
vb: VarBuilder,
|
||||
input_channel: usize,
|
||||
channels: usize,
|
||||
rates: Vec<usize>,
|
||||
d_out: usize,
|
||||
) -> Result<Self> {
|
||||
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<Tensor> {
|
||||
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<usize>,
|
||||
decoder_dim: usize,
|
||||
decoder_rates: Vec<usize>,
|
||||
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<usize>,
|
||||
laten_dim: Option<usize>,
|
||||
decoder_dim: usize,
|
||||
decoder_rates: Vec<usize>,
|
||||
sample_rate: usize,
|
||||
) -> Result<Self> {
|
||||
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<usize>) -> Result<Tensor>{
|
||||
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<Tensor> {
|
||||
let x = self.decoder.forward(z)?;
|
||||
Ok(x)
|
||||
}
|
||||
|
||||
pub fn encode(&self, audio_data: &Tensor, sample_rate: Option<usize>) -> Result<Tensor> {
|
||||
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)
|
||||
}
|
||||
}
|
||||
+33
-22
@@ -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<f32>,
|
||||
pub short_factor: Vec<f32>,
|
||||
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,
|
||||
}
|
||||
@@ -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<f32>,
|
||||
long_factor: Vec<f32>,
|
||||
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<Self> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<Embedding>,
|
||||
layers: Vec<MiniCPMDecoderLayer>,
|
||||
norm: RmsNorm,
|
||||
rope_emb: MiniCPMLongRoPE,
|
||||
// lm_head: Linear,
|
||||
}
|
||||
|
||||
impl MiniCPMModel {
|
||||
pub fn new(vb: VarBuilder, cfg: VoxMiniCPM4Config) -> Result<Self> {
|
||||
// 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<Tensor> {
|
||||
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<Tensor> {
|
||||
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()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
pub mod config;
|
||||
pub mod audio_vae;
|
||||
pub mod minicpm4;
|
||||
pub mod tokenizer;
|
||||
pub mod model;
|
||||
@@ -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<Self> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
assert_eq!(dim % 2, 0, "SinusoidalPosEmb requires dim to be even");
|
||||
Ok(Self { dim })
|
||||
}
|
||||
pub fn forward(&self, x: &Tensor, scale: usize) -> Result<Tensor> {
|
||||
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<usize>,
|
||||
) -> Result<Self> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
// 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<Tensor> {
|
||||
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<Self> {
|
||||
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<String>,
|
||||
prompt_wav_path: Option<String>,
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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::<u32>()?;
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
use anyhow::{Ok, Result, anyhow};
|
||||
use candle_core::Tensor;
|
||||
use tokenizers::Tokenizer;
|
||||
|
||||
pub struct SingleChineseTokenizer {
|
||||
tokenizer: Tokenizer,
|
||||
multichar_tokens: Vec<String>,
|
||||
}
|
||||
|
||||
impl SingleChineseTokenizer {
|
||||
pub fn new(path: &str) -> Result<Self> {
|
||||
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<Vec<u32>> {
|
||||
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<String> = 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<u32> = split_character
|
||||
.iter()
|
||||
.filter_map(|c| self.tokenizer.token_to_id(c))
|
||||
.collect();
|
||||
Ok(ids)
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
@@ -41,4 +41,5 @@ impl TokenizerModel {
|
||||
.map_err(|e| anyhow!(format!("tokenizer encode error{}", e)))?;
|
||||
Ok(decode)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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<f32>,
|
||||
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::<f32>()?;
|
||||
|
||||
let window_val:Vec<f32> = 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<Tensor> {
|
||||
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<f32>,
|
||||
) -> Result<Tensor> {
|
||||
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<Tensor> {
|
||||
resample(
|
||||
waveform,
|
||||
orig_freq,
|
||||
new_freq,
|
||||
6,
|
||||
0.99,
|
||||
ResamplingMethod::SincInterpHann,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn load_audio<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, usize)> {
|
||||
let mut reader = WavReader::open(path)?;
|
||||
let spec = reader.spec();
|
||||
let samples: Vec<f32> = 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::<i16>()
|
||||
.map(|s| s.map(|sample| sample as f32 / max_value))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
}
|
||||
SampleFormat::Float => {
|
||||
// 直接读取浮点数样本
|
||||
reader.samples::<f32>().collect::<Result<Vec<_>, _>>()?
|
||||
}
|
||||
};
|
||||
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<P: AsRef<Path>>(
|
||||
path: P,
|
||||
device: Device,
|
||||
target_sample_rate: Option<usize>,
|
||||
) -> Result<Tensor> {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -2,3 +2,4 @@ pub mod img_utils;
|
||||
pub mod tensor_utils;
|
||||
pub mod utils;
|
||||
pub mod video_utils;
|
||||
pub mod audio_utils;
|
||||
@@ -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<Tens
|
||||
let indices = indices.broadcast_add(&Tensor::new(vec![1u32], input_ids.device())?)?;
|
||||
Ok(indices)
|
||||
}
|
||||
|
||||
pub fn linspace(start: f32, end: f32, steps: usize, device: &Device) -> Result<Tensor> {
|
||||
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<f32> = (0..steps)
|
||||
.map(|i| start + i as f32 * step_size)
|
||||
.collect();
|
||||
|
||||
let t = Tensor::from_slice(&data, steps, device)?;
|
||||
Ok(t)
|
||||
}
|
||||
+3
-2
@@ -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<Vec<String>> {
|
||||
pub fn find_type_files(path: &str, extension_type: &str) -> Result<Vec<String>> {
|
||||
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<Vec<String>> {
|
||||
|
||||
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<Vec<String>> {
|
||||
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
|
||||
|
||||
+11
-1
@@ -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]
|
||||
@@ -21,3 +21,13 @@ fn minicpm4_config() -> Result<()> {
|
||||
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(())
|
||||
}
|
||||
@@ -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>= 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(())
|
||||
}
|
||||
+17
-37
@@ -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;
|
||||
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/Qwen/Qwen2.5-VL-3B-Instruct/";
|
||||
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);
|
||||
|
||||
|
||||
+7
-10
@@ -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);
|
||||
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
+29
-3
@@ -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)?;
|
||||
@@ -16,3 +20,25 @@ 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(())
|
||||
}
|
||||
Reference in New Issue
Block a user