diff --git a/Cargo.lock b/Cargo.lock index 216a541..178d23d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -50,6 +50,7 @@ dependencies = [ "minijinja", "modelscope", "num", + "rand 0.9.2", "rayon", "realfft", "regex", diff --git a/Cargo.toml b/Cargo.toml index 6a3e13c..06fcf81 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -42,6 +42,7 @@ half = "2.7.1" byteorder = "1.5.0" sentencepiece = "0.13.1" regex = "1.12.3" +rand = "0.9.2" [features] flash-attn = ["candle-flash-attn"] diff --git a/docs/installation.md b/docs/installation.md index bf74ccb..4146ed3 100644 --- a/docs/installation.md +++ b/docs/installation.md @@ -85,7 +85,7 @@ cargo build --release --features ffmpeg ```bash # Install build dependencies sudo apt-get update -sudo apt-get install -y build-essential pkg-config git clang +sudo apt-get install -y build-essential pkg-config git clang cmake # For FFmpeg feature sudo apt-get install -y ffmpeg libavutil-dev libavcodec-dev \ @@ -166,7 +166,7 @@ cargo build --release # Follow Linux instructions inside WSL2 wsl sudo apt-get update -sudo apt-get install -y build-essential pkg-config git clang +sudo apt-get install -y build-essential pkg-config git clang cmake ``` ## Feature Flags diff --git a/docs/installation.zh-CN.md b/docs/installation.zh-CN.md index 3a68bb6..b9470ea 100644 --- a/docs/installation.zh-CN.md +++ b/docs/installation.zh-CN.md @@ -85,7 +85,7 @@ cargo build --release --features ffmpeg ```bash # 安装构建依赖 sudo apt-get update -sudo apt-get install -y build-essential pkg-config git clang +sudo apt-get install -y build-essential pkg-config git clang cmake # FFmpeg 功能所需 sudo apt-get install -y ffmpeg libavutil-dev libavcodec-dev \ @@ -165,7 +165,7 @@ cargo build --release # 在 WSL2 中按照 Linux 说明操作 wsl sudo apt-get update -sudo apt-get install -y build-essential pkg-config git clang +sudo apt-get install -y build-essential pkg-config git clang cmake ``` ## 功能特性 diff --git a/src/exec/qwen3_asr.rs b/src/exec/qwen3_asr.rs index 44e2687..2b6bde0 100644 --- a/src/exec/qwen3_asr.rs +++ b/src/exec/qwen3_asr.rs @@ -5,14 +5,13 @@ use std::time::Instant; use anyhow::{Ok, Result}; use crate::exec::ExecModel; +use crate::models::GenerateModel; use crate::models::qwen3_asr::generate::Qwen3AsrGenerateModel; -use crate::models::{GenerateModel}; pub struct Qwen3ASRExec; impl ExecModel for Qwen3ASRExec { fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> { - let i_start = Instant::now(); let mut model = Qwen3AsrGenerateModel::init(weight_path, None, None)?; let i_duration = i_start.elapsed(); diff --git a/src/models/bigvgan/config.rs b/src/models/bigvgan/config.rs new file mode 100644 index 0000000..a21164c --- /dev/null +++ b/src/models/bigvgan/config.rs @@ -0,0 +1,46 @@ +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct BigVGANConfig { + pub resblock: String, + pub num_gpus: usize, + pub batch_size: usize, + pub learning_rate: f64, + pub adam_b1: f64, + pub adam_b2: f64, + pub lr_decay: f64, + pub seed: u32, + pub upsample_rates: Vec, + pub upsample_kernel_sizes: Vec, + pub upsample_initial_channel: usize, + pub resblock_kernel_sizes: Vec, + pub resblock_dilation_sizes: Vec>, + pub use_tanh_at_final: bool, + pub use_bias_at_final: bool, + pub activation: String, + pub snake_logscale: bool, + pub use_cqtd_instead_of_mrd: bool, + pub cqtd_filters: usize, + pub cqtd_max_filters: usize, + pub cqtd_filters_scale: usize, + pub cqtd_dilations: Vec, + pub cqtd_hop_lengths: Vec, + pub cqtd_n_octaves: Vec, + pub cqtd_bins_per_octaves: Vec, + pub mpd_reshapes: Vec, + pub use_spectral_norm: bool, + pub discriminator_channel_mult: usize, + pub use_multiscale_melloss: bool, + pub lambda_melloss: f64, + pub clip_grad_norm: f64, + pub segment_size: usize, + pub num_mels: usize, + pub num_freq: usize, + pub n_fft: usize, + pub hop_size: usize, + pub win_size: usize, + pub sampling_rate: usize, + pub fmin: usize, + pub fmax: Option, + pub fmax_for_loss: Option, + pub normalize_volume: bool, + pub num_workers: usize, +} diff --git a/src/models/bigvgan/mod.rs b/src/models/bigvgan/mod.rs new file mode 100644 index 0000000..c152f9c --- /dev/null +++ b/src/models/bigvgan/mod.rs @@ -0,0 +1,333 @@ +use anyhow::Result; +use candle_core::{D, Tensor}; +use candle_nn::{Init, VarBuilder}; + +use crate::{ + models::{ + bigvgan::config::BigVGANConfig, + common::{WNConv1d, WNConvTranspose1d}, + }, + utils::tensor_utils::pad_replicate_last_dim, +}; + +pub mod config; + +pub struct UpSample1d { + // ratio: usize, + stride: usize, + pad: usize, + pad_left: usize, + pad_right: usize, + filter: Tensor, +} + +impl UpSample1d { + pub fn new(vb: VarBuilder, ratio: usize, kernel_size: Option) -> Result { + let stride = ratio; + let kernel_size = kernel_size.unwrap_or(6 * ratio / 2 * 2); + let pad = kernel_size / ratio - 1; + let pad_left = pad * stride + (kernel_size - stride) / 2; + let pad_right = pad * stride + (kernel_size - stride + 1) / 2; + let filter = vb.get_with_hints((1, 1, kernel_size), "filter", Init::Const(0.0))?; + + Ok(Self { + // ratio, + stride, + pad, + pad_left, + pad_right, + filter, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let c = xs.dim(1)?; + let xs = pad_replicate_last_dim(xs, (self.pad, self.pad))?; + let xs = xs.conv_transpose1d(&self.filter.repeat((c, 1, 1))?, 0, 0, self.stride, 1, c)?; + let xs_last_dim = xs.dim(D::Minus1)?; + let xs_len = xs_last_dim - self.pad_left - self.pad_right; + let xs = xs.narrow(D::Minus1, self.pad_left, xs_len)?; + Ok(xs) + } +} + +pub struct DownSample1d { + stride: usize, + // kernel_size: usize, + pad_left: usize, + pad_right: usize, + filter: Tensor, +} + +impl DownSample1d { + pub fn new(vb: VarBuilder, ratio: usize, kernel_size: Option) -> Result { + let stride = ratio; + let kernel_size = kernel_size.unwrap_or(6 * ratio / 2 * 2); + let even = if kernel_size.is_multiple_of(2) { 1 } else { 0 }; + let pad_left = kernel_size / 2 - even; + let pad_right = kernel_size / 2; + let filter = vb.get_with_hints((1, 1, kernel_size), "lowpass.filter", Init::Const(0.0))?; + Ok(Self { + stride, + // kernel_size, + pad_left, + pad_right, + filter, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let c = xs.dim(1)?; + let xs = pad_replicate_last_dim(xs, (self.pad_left, self.pad_right))?; + let xs = xs.conv1d(&self.filter.repeat((c, 1, 1))?, 0, self.stride, 1, c)?; + + Ok(xs) + } +} + +pub struct SnakeBeta { + alpha: Tensor, + beta: Tensor, + no_div_by_zero: f64, +} + +impl SnakeBeta { + pub fn new(vb: VarBuilder, in_features: usize) -> Result { + let no_div_by_zero = 0.000000001f64; + let alpha = vb + .get_with_hints(in_features, "alpha", Init::Const(0.0))? + .unsqueeze(0)? + .unsqueeze(D::Minus1)? + .contiguous()? + .exp()?; + let beta = vb + .get_with_hints(in_features, "beta", Init::Const(0.0))? + .unsqueeze(0)? + .unsqueeze(D::Minus1)? + .contiguous()? + .exp()?; + Ok(Self { + alpha, + beta, + no_div_by_zero, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let beta = (1.0 / self.beta.affine(1.0, self.no_div_by_zero)?)?; + let xs = xs + .broadcast_mul(&self.alpha)? + .sin()? + .powf(2.0)? + .broadcast_mul(&beta)? + .add(xs)?; + Ok(xs) + } +} + +pub struct TorchActivation1d { + upsample: UpSample1d, + downsample: DownSample1d, + act: SnakeBeta, +} + +impl TorchActivation1d { + pub fn new( + vb: VarBuilder, + up_ratio: usize, + down_ratio: usize, + up_kernel_size: usize, + down_kernel_size: usize, + channels: usize, + ) -> Result { + let upsample = UpSample1d::new(vb.pp("upsample"), up_ratio, Some(up_kernel_size))?; + let downsample = + DownSample1d::new(vb.pp("downsample"), down_ratio, Some(down_kernel_size))?; + let act = SnakeBeta::new(vb.pp("act"), channels)?; + Ok(Self { + upsample, + downsample, + act, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.upsample.forward(xs)?; + let xs = self.act.forward(&xs)?; + let xs = self.downsample.forward(&xs)?; + Ok(xs) + } +} + +pub struct AMPBlock1 { + convs1: Vec, + convs2: Vec, + activations: Vec, +} + +impl AMPBlock1 { + pub fn new( + vb: VarBuilder, + channels: usize, + kernel_size: usize, + dilation: Vec, + ) -> Result { + let vb_convs1 = vb.pp("convs1"); + let mut convs1 = vec![]; + for (i, &d) in dilation.iter().enumerate() { + let pad = ((kernel_size * d - d) as f32 / 2.0).round() as usize; + let layer = WNConv1d::new( + vb_convs1.pp(i), + channels, + channels, + kernel_size, + d, + pad, + 1, + 1, + true, + )?; + convs1.push(layer); + } + let vb_convs2 = vb.pp("convs2"); + let mut convs2 = vec![]; + for (i, _) in dilation.iter().enumerate() { + let pad = ((kernel_size - 1) as f32 / 2.0).round() as usize; + let layer = WNConv1d::new( + vb_convs2.pp(i), + channels, + channels, + kernel_size, + 1, + pad, + 1, + 1, + true, + )?; + convs2.push(layer); + } + + let num_layer = convs1.len() + convs2.len(); + let act_vb = vb.pp("activations"); + let mut activations = vec![]; + for i in 0..num_layer { + let layer = TorchActivation1d::new(act_vb.pp(i), 2, 2, 12, 12, channels)?; + activations.push(layer); + } + + Ok(Self { + convs1, + convs2, + activations, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let lens = self.convs1.len(); + let mut xs = xs.clone(); + for i in 0..lens { + let xt = self.activations[i * 2].forward(&xs)?; + let xt = self.convs1[i].forward(&xt)?; + let xt = self.activations[i * 2 + 1].forward(&xt)?; + let xt = self.convs2[i].forward(&xt)?; + xs = xs.add(&xt)?; + } + Ok(xs) + } +} + +pub struct BigVGAN { + num_kernels: usize, + num_upsamples: usize, + conv_pre: WNConv1d, + ups: Vec, + resblocks: Vec, + activation_post: TorchActivation1d, + conv_post: WNConv1d, + use_tanh_at_final: bool, +} + +impl BigVGAN { + pub fn new(vb: VarBuilder, cfg: &BigVGANConfig) -> Result { + let num_kernels = cfg.resblock_kernel_sizes.len(); + let num_upsamples = cfg.upsample_rates.len(); + let conv_pre = WNConv1d::new( + vb.pp("conv_pre"), + cfg.num_mels, + cfg.upsample_initial_channel, + 7, + 1, + 3, + 1, + 1, + true, + )?; + + let vb_ups = vb.pp("ups"); + let mut ups = vec![]; + for (i, (&u, &k)) in cfg + .upsample_rates + .iter() + .zip(cfg.upsample_kernel_sizes.iter()) + .enumerate() + { + let in_c = cfg.upsample_initial_channel / (2_i32.pow(i as u32) as usize); + let out_c = cfg.upsample_initial_channel / (2_i32.pow(i as u32 + 1) as usize); + let pad = (k - u) / 2; + let layer = + WNConvTranspose1d::new(vb_ups.pp(i).pp("0"), in_c, out_c, 1, k, pad, 0, 1, u)?; + ups.push(layer); + } + let vb_resblocks = vb.pp("resblocks"); + let mut resblocks = vec![]; + let ups_len = ups.len(); + let res_len = cfg.resblock_kernel_sizes.len(); + let mut ch = 0; + for i in 0..ups_len { + ch = cfg.upsample_initial_channel / (2_i32.pow(i as u32 + 1) as usize); + for (j, (&k, d)) in cfg + .resblock_kernel_sizes + .iter() + .zip(cfg.resblock_dilation_sizes.iter()) + .enumerate() + { + let layer = AMPBlock1::new(vb_resblocks.pp(i * res_len + j), ch, k, d.clone())?; + resblocks.push(layer); + } + } + let activation_post = TorchActivation1d::new(vb.pp("activation_post"), 2, 2, 12, 12, ch)?; + let conv_post = WNConv1d::new(vb.pp("conv_post"), ch, 1, 7, 1, 3, 1, 1, false)?; + + Ok(Self { + num_kernels, + num_upsamples, + conv_pre, + ups, + resblocks, + activation_post, + conv_post, + use_tanh_at_final: cfg.use_tanh_at_final, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let mut xs = self.conv_pre.forward(xs)?; + for i in 0..self.num_upsamples { + xs = self.ups[i].forward(&xs)?; + let mut xs_j = xs.zeros_like()?; + for j in 0..self.num_kernels { + let j_ = self.resblocks[i * self.num_kernels + j].forward(&xs)?; + xs_j = xs_j.add(&j_)?; + } + xs = xs_j.affine(1.0 / (self.num_kernels as f64), 0.0)?; + } + xs = self.activation_post.forward(&xs)?; + xs = self.conv_post.forward(&xs)?; + xs = if self.use_tanh_at_final { + xs.tanh()? + } else { + xs.clamp(-1.0, 1.0)? + }; + Ok(xs) + } +} diff --git a/src/models/campplus/mod.rs b/src/models/campplus/mod.rs index d5638e5..a523dbd 100644 --- a/src/models/campplus/mod.rs +++ b/src/models/campplus/mod.rs @@ -25,7 +25,11 @@ impl Shortcut { ) -> Result { let conv_0 = get_conv2d(vb.pp("0"), in_c, out_c, ks, padding, 1, 1, 1, bias)?; let bn_1 = get_batch_norm(vb.pp("1"), 1e-5, out_c, true)?; - Ok(Self { conv_0, bn_1, stride }) + Ok(Self { + conv_0, + bn_1, + stride, + }) } pub fn forward(&self, x: &Tensor) -> Result { @@ -36,7 +40,7 @@ impl Shortcut { let indices = Tensor::arange(0u32, half_h as u32, x.device())?.affine(2.0, 0.0)?; x = x.index_select(&indices, 2)?; } - x = self.bn_1.forward_t(&x, false)?; + x = self.bn_1.forward_t(&x, false)?; Ok(x) } } @@ -106,7 +110,7 @@ impl BasicResBlock { } else { xs = xs.add(&residual)?; } - xs = xs.relu()?; + xs = xs.relu()?; Ok(xs) } } @@ -435,7 +439,7 @@ impl DenseLayer { .forward(&xs.unsqueeze(D::Minus1)?)? .squeeze(D::Minus1)? } else { - self.linear.forward(&xs)? + self.linear.forward(xs)? }; let xs = self.nonlinear.forward_t(&xs, false)?; Ok(xs) @@ -463,7 +467,7 @@ impl XVector { let mut channels = init_channels; let mut blocks = vec![]; let mut transits = vec![]; - let params = vec![(12, 3, 1), (24, 3, 2), (16, 3, 2)]; + let params = [(12, 3, 1), (24, 3, 2), (16, 3, 2)]; for (i, (num_layers, ks, dilation)) in params.iter().enumerate() { let block = CAMDenseTDNNBlock::new( vb.pp(format!("block{}", i + 1)), @@ -477,7 +481,7 @@ impl XVector { false, )?; blocks.push(block); - channels = channels + num_layers * growth_rate; + channels += num_layers * growth_rate; let transit = TransitLayer::new( vb.pp(format!("transit{}", i + 1)), channels, diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index c3222c3..415161e 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -2,8 +2,8 @@ use anyhow::{Result, anyhow}; use candle_core::{D, IndexOp, Tensor}; use candle_nn::{ Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, - ConvTranspose1d, ConvTranspose1dConfig, Embedding, GroupNorm, Init, LayerNorm, LayerNormConfig, - Linear, Module, ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, + ConvTranspose1d, ConvTranspose1dConfig, Embedding, Init, LayerNorm, LayerNormConfig, Linear, + Module, ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm, }; use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator}; @@ -700,11 +700,7 @@ pub fn get_layer_norm(vb: VarBuilder, eps: f64, dim: usize, affine: bool) -> Res Ok(norm) } -pub fn get_layer_norm_without_weight( - vb: VarBuilder, - eps: f64, - dim: usize, -) -> Result { +pub fn get_layer_norm_without_weight(vb: VarBuilder, eps: f64, dim: usize) -> Result { let weight = Tensor::ones(dim, vb.dtype(), vb.device())?; let bias = Tensor::zeros(dim, vb.dtype(), vb.device())?; Ok(LayerNorm::new(weight, bias, eps)) @@ -1027,11 +1023,25 @@ impl GLU { Ok(Self { dim }) } pub fn forward(&self, xs: &Tensor) -> Result { - let half_dim = xs.dim(self.dim)? / 2; - let a = xs.narrow(self.dim, 0, half_dim)?; - let b = xs.narrow(self.dim, half_dim, half_dim)?; - let b = sigmoid(&b)?; - let xs = a.mul(&b)?; + let x_ = xs.chunk(2, self.dim)?; + let x_1 = sigmoid(x_[1].as_ref())?; + let xs = x_1.mul(x_[0].as_ref())?; + Ok(xs) + } +} + +pub struct GEGLU { + dim: usize, +} + +impl GEGLU { + pub fn new(dim: usize) -> Result { + Ok(Self { dim }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let x_ = xs.chunk(2, self.dim)?; + let x_1 = x_[1].as_ref().gelu()?; + let xs = x_1.mul(x_[0].as_ref())?; Ok(xs) } } @@ -1103,8 +1113,8 @@ impl WNConvTranspose1d { let normalized_weight = weight_v.broadcast_div(&weight_norm)?; let scaled_weight = normalized_weight.broadcast_mul(&weight_g)?; let config = ConvTranspose1dConfig { - padding: padding, - output_padding: output_padding, + padding, + output_padding, stride, dilation, groups, @@ -1184,3 +1194,236 @@ pub fn mish(xs: &Tensor) -> Result { let xs = xs.mul(&tanh)?; Ok(xs) } + +pub struct GPT2Attention { + num_heads: usize, + head_dim: usize, + c_attn: Linear, + c_proj: Linear, + kv_cache: Option<(Tensor, Tensor)>, +} + +impl GPT2Attention { + pub fn new(vb: VarBuilder, hidden_size: usize, num_heads: usize) -> Result { + let c_attn_weight = vb + .get_with_hints( + (hidden_size, 3 * hidden_size), + "c_attn.weight", + Init::Const(1.0), + )? + .t()?; + let c_attn_bias = vb.get_with_hints(3 * hidden_size, "c_attn.bias", Init::Const(0.0))?; + let c_attn = Linear::new(c_attn_weight, Some(c_attn_bias)); + // let c_attn = linear_b(3 * hidden_size, hidden_size, true, vb.pp("c_attn"))?; + let c_proj_weight = vb + .get_with_hints( + (hidden_size, hidden_size), + "c_proj.weight", + Init::Const(1.0), + )? + .t()?; + let c_proj_bias = vb.get_with_hints(hidden_size, "c_proj.bias", Init::Const(0.0))?; + let c_proj = Linear::new(c_proj_weight, Some(c_proj_bias)); + // let c_proj = linear_b(hidden_size, hidden_size, true, vb.pp("c_proj"))?; + let head_dim = hidden_size / num_heads; + Ok(Self { + num_heads, + head_dim, + c_attn, + c_proj, + kv_cache: None, + }) + } + + pub fn forward(&mut self, xs: &Tensor, attention_mask: Option<&Tensor>) -> Result { + let (b, seq_len, _) = xs.dims3()?; + let xs = self.c_attn.forward(xs)?; + let xs_splits = xs.chunk(3, 2)?; + let query_states = xs_splits[0] + .as_ref() + .reshape((b, seq_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let key_states = xs_splits[1] + .as_ref() + .reshape((b, seq_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let value_states = xs_splits[2] + .as_ref() + .reshape((b, seq_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let (key_states, value_states) = match &self.kv_cache { + None => (key_states, value_states), + Some((prev_k, prev_v)) => { + let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; + let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; + (key_states, value_states) + } + }; + + self.kv_cache = Some((key_states.clone(), value_states.clone())); + let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let attn_output = eager_attention_forward( + &query_states, + &key_states, + &value_states, + None, + attention_mask, + scale, + )?; + let attn_output = attn_output.reshape((b, seq_len, self.num_heads * self.head_dim))?; + let attn_output = attn_output.apply(&self.c_proj)?; + Ok(attn_output) + } + + pub fn clear_kv_cache(&mut self) { + self.kv_cache = None + } +} + +pub struct GPT2MLP { + linear1: Linear, + linear2: Linear, + act: Activation, +} + +impl GPT2MLP { + pub fn new( + vb: VarBuilder, + in_dim: usize, + middle_dim: usize, + out_dim: usize, + act: Activation, + ) -> Result { + let c_fc_weight = vb + .get_with_hints((in_dim, middle_dim), "c_fc.weight", Init::Const(1.0))? + .t()?; + let c_fc_bias = vb.get_with_hints(middle_dim, "c_fc.bias", Init::Const(0.0))?; + let c_fc = Linear::new(c_fc_weight, Some(c_fc_bias)); + + let c_proj_weight = vb + .get_with_hints((middle_dim, out_dim), "c_proj.weight", Init::Const(1.0))? + .t()?; + let c_proj_bias = vb.get_with_hints(out_dim, "c_proj.bias", Init::Const(0.0))?; + let c_proj = Linear::new(c_proj_weight, Some(c_proj_bias)); + + Ok(Self { + linear1: c_fc, + linear2: c_proj, + act, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = xs + .apply(&self.linear1)? + .apply(&self.act)? + .apply(&self.linear2)?; + Ok(xs) + } +} + +pub struct GPT2Block { + ln_1: LayerNorm, + attn: GPT2Attention, + ln_2: LayerNorm, + mlp: GPT2MLP, +} + +impl GPT2Block { + pub fn new( + vb: VarBuilder, + hidden_size: usize, + num_heads: usize, + inner_dim: Option, + ) -> Result { + let inner_dim = inner_dim.unwrap_or(4 * hidden_size); + let ln_1 = get_layer_norm(vb.pp("ln_1"), 1e-5, hidden_size, true)?; + let attn = GPT2Attention::new(vb.pp("attn"), hidden_size, num_heads)?; + let ln_2 = get_layer_norm(vb.pp("ln_2"), 1e-5, hidden_size, true)?; + let mlp = GPT2MLP::new( + vb.pp("mlp"), + hidden_size, + inner_dim, + hidden_size, + Activation::NewGelu, + )?; + Ok(Self { + ln_1, + attn, + ln_2, + mlp, + }) + } + + pub fn forward(&mut self, xs: &Tensor, attention_mask: Option<&Tensor>) -> Result { + let residual = xs.clone(); + let xs = self.ln_1.forward(xs)?; + let xs = self.attn.forward(&xs, attention_mask)?; + let residual = xs.add(&residual)?; + let xs = self.ln_2.forward(&residual)?; + let xs = self.mlp.forward(&xs)?; + let xs = xs.add(&residual)?; + Ok(xs) + } + pub fn clear_kv_cache(&mut self) { + self.attn.clear_kv_cache() + } +} + +#[allow(unused)] +pub struct GPT2Model { + wte: Embedding, + // wpe: Embedding, + h: Vec, + ln_f: LayerNorm, +} + +#[allow(unused)] +impl GPT2Model { + pub fn new( + vb: VarBuilder, + hidden_size: usize, + num_heads: usize, + num_hidden_layers: usize, + wte_embeddings: &Tensor, + ) -> Result { + // let wte = embedding(vocab_size, hidden_size, vb.pp("wte"))?; + let wte = Embedding::new(wte_embeddings.clone(), hidden_size); + // let wpe = embedding(max_position_embeddings, hidden_size, vb.pp("wpe"))?; + let vb_layers = vb.pp("h"); + let mut h = vec![]; + for i in 0..num_hidden_layers { + let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?; + h.push(block); + } + let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?; + Ok(Self { wte, h, ln_f }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor) -> Result { + let (b_size, seq_len, _) = inputs_embeds.dims3()?; + let mut xs = inputs_embeds.clone(); + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + b_size, + seq_len, + 0, + xs.device(), + )?) + } + }; + for block in &mut self.h { + xs = block.forward(&xs, attention_mask.as_ref())?; + } + xs = self.ln_f.forward(&xs)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.h.iter_mut() { + layer.clear_kv_cache() + } + } +} diff --git a/src/models/feature_extractor/config.rs b/src/models/feature_extractor/config.rs index b79a94d..8334c30 100644 --- a/src/models/feature_extractor/config.rs +++ b/src/models/feature_extractor/config.rs @@ -18,4 +18,4 @@ pub struct FeatureExtractor { fn default_sampling_rate() -> usize { 16000 -} \ No newline at end of file +} diff --git a/src/models/feature_extractor/mod.rs b/src/models/feature_extractor/mod.rs index 80016f5..6e3fd69 100644 --- a/src/models/feature_extractor/mod.rs +++ b/src/models/feature_extractor/mod.rs @@ -1,3 +1,3 @@ -pub mod seamless_m4t_feature_extractor; +pub mod config; pub mod feature_extraction_whisper; -pub mod config; \ No newline at end of file +pub mod seamless_m4t_feature_extractor; diff --git a/src/models/feature_extractor/seamless_m4t_feature_extractor.rs b/src/models/feature_extractor/seamless_m4t_feature_extractor.rs index a399fb7..12e23e9 100644 --- a/src/models/feature_extractor/seamless_m4t_feature_extractor.rs +++ b/src/models/feature_extractor/seamless_m4t_feature_extractor.rs @@ -82,7 +82,7 @@ impl SeamlessM4TFeatureExtractor { 0.97, Some(&self.mel_filters), Some("log"), - 1.192092955078125e-07, + 1.192_092_9e-7, true, )? .transpose(D::Minus1, D::Minus2)?; diff --git a/src/models/glm_asr_nano/config.rs b/src/models/glm_asr_nano/config.rs index ea90498..cf2a3dd 100644 --- a/src/models/glm_asr_nano/config.rs +++ b/src/models/glm_asr_nano/config.rs @@ -11,8 +11,6 @@ pub struct GlmAsrNanoProcessorConfig { pub max_audio_len: usize, } - - #[derive(Debug, Clone, PartialEq, Deserialize)] pub struct GlmAsrNanoConfig { pub audio_config: GlmAsrAudioConfig, diff --git a/src/models/glm_asr_nano/processor.rs b/src/models/glm_asr_nano/processor.rs index e20222d..2ec3985 100644 --- a/src/models/glm_asr_nano/processor.rs +++ b/src/models/glm_asr_nano/processor.rs @@ -1,4 +1,3 @@ - use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::Result; use candle_core::{D, DType, Device, IndexOp, Tensor}; @@ -27,6 +26,7 @@ pub struct GlmAsrNanoProcessor { whisper_feature_extrator: WhisperFeatureExtractor, } +#[allow(unused)] impl GlmAsrNanoProcessor { pub fn new(path: &str, device: &Device, dtype: DType) -> Result { let path = path.to_string(); diff --git a/src/models/index_tts2/config.rs b/src/models/index_tts2/config.rs index e6affb6..67fc52f 100644 --- a/src/models/index_tts2/config.rs +++ b/src/models/index_tts2/config.rs @@ -5,7 +5,7 @@ use crate::models::mask_gct::config::SemanticCodec; #[derive(Debug, Clone, PartialEq, Deserialize)] pub struct IndexTTS2Config { pub dataset: Dataset, - pub gpt: Gpt, + pub gpt: GptConfig, pub semantic_codec: SemanticCodec, pub s2mel: S2MelConfig, pub gpt_checkpoint: String, @@ -39,7 +39,7 @@ pub struct Mel { } #[derive(Debug, Clone, PartialEq, Deserialize)] -pub struct Gpt { +pub struct GptConfig { pub model_dim: usize, pub max_mel_tokens: usize, pub max_text_tokens: usize, @@ -206,7 +206,7 @@ impl DiTModelArgs { has_cross_attention: false, context_dim: 0, uvit_skip_connection: config.uvit_skip_connection, - time_as_token: config.uvit_skip_connection, + time_as_token: config.time_as_token, } } } diff --git a/src/models/index_tts2/generate.rs b/src/models/index_tts2/generate.rs index 5c4f63d..35a3ce8 100644 --- a/src/models/index_tts2/generate.rs +++ b/src/models/index_tts2/generate.rs @@ -1,78 +1,59 @@ -use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use aha_openai_dive::v1::resources::chat::{ChatCompletionParameters, ChatCompletionResponse}; use anyhow::{Result, anyhow}; +use base64::{Engine, prelude::BASE64_STANDARD}; use candle_core::{DType, Device}; use sentencepiece::SentencePieceProcessor; use crate::{ models::index_tts2::{ - config::IndexTTS2Config, - model::IndexTTS2Model, - processor::IndexTTS2Processor, - utils::{TextNormalizer, tokenize_by_cjk_char}, + config::IndexTTS2Config, model::IndexTTS2Model, utils::tokenize_by_cjk_char, }, tokenizer::sentencepiece_encode, utils::{ - audio_utils::extract_audio_url, extract_user_text, get_default_save_dir, get_device, - get_dtype, + audio_utils::get_audio_wav_u8, build_audio_completion_response, extract_user_text, + get_default_save_dir, get_device, }, }; pub struct IndexTTS2Generate { - processor: IndexTTS2Processor, tokenizer: SentencePieceProcessor, - config: IndexTTS2Config, - cache_spk_audio_prompt: Option, + // config: IndexTTS2Config, model: IndexTTS2Model, device: Device, + sample_rate: u32, + model_name: String, } +#[allow(unused)] impl IndexTTS2Generate { pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { let config_path = path.to_string() + "/config.yaml"; let save_dir = get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; let config: IndexTTS2Config = serde_yaml::from_slice(&std::fs::read(config_path)?)?; let device = get_device(device); - let dtype = get_dtype(dtype, "bf16"); - let processor = IndexTTS2Processor::new(&device)?; let bpe_path = path.to_string() + "/bpe.model"; let tokenizer = SentencePieceProcessor::open(bpe_path) .map_err(|e| anyhow!(format!("load bpe,model file error:{}", e)))?; - let model = IndexTTS2Model::new(path, &save_dir, &config, &device, dtype)?; + let model = IndexTTS2Model::new(path, &save_dir, &config, &device)?; Ok(Self { - processor, tokenizer, - config, - cache_spk_audio_prompt: None, + // config, model, device, + sample_rate: 22050, + model_name: "index-tts2".to_string(), }) } - pub fn use_prompt(&self, mes: &ChatCompletionParameters) -> bool { - if let Some(cache) = &self.cache_spk_audio_prompt { - let audio_vec = extract_audio_url(mes); - if audio_vec.len() == 0 { - true - } else { - if cache.eq(&audio_vec[0]) { true } else { false } - } - } else { - false - } - } - - pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<()> { + pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result { let text = extract_user_text(&mes)?; let text = tokenize_by_cjk_char(&text, true); let input_ids = sentencepiece_encode(&text, &self.tokenizer, &self.device)?; - let (audio_22k, audio_16k) = if self.use_prompt(&mes) { - (None, None) - } else { - let (audio_22k, audio_16k, prompt) = self.processor.process_info(&mes)?; - self.cache_spk_audio_prompt = Some(prompt); - (Some(audio_22k), Some(audio_16k)) - }; - let _ = self.model.forward(&input_ids, audio_22k.as_ref(), audio_16k.as_ref())?; - Ok(()) + + let audio = self.model.forward(&input_ids, &mes)?; + let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate)?; + let base64_audio = BASE64_STANDARD.encode(wav_u8); + let response = build_audio_completion_response(&base64_audio, &self.model_name); + Ok(response) } } diff --git a/src/models/index_tts2/mod.rs b/src/models/index_tts2/mod.rs index 6634bcd..13cbcd3 100644 --- a/src/models/index_tts2/mod.rs +++ b/src/models/index_tts2/mod.rs @@ -1,5 +1,5 @@ pub mod config; pub mod generate; pub mod model; -pub mod processor; -pub mod utils; \ No newline at end of file +// pub mod processor; +pub mod utils; diff --git a/src/models/index_tts2/model.rs b/src/models/index_tts2/model.rs index a87473f..19847bc 100644 --- a/src/models/index_tts2/model.rs +++ b/src/models/index_tts2/model.rs @@ -1,38 +1,48 @@ +use std::collections::HashMap; + +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::{Result, anyhow}; -use candle_core::{D, DType, Device, IndexOp, Tensor, pickle::read_all_with_key}; +use candle_core::{D, DType, Device, IndexOp, Shape, Tensor, pickle::read_all_with_key}; use candle_nn::{ - Activation, Conv1d, Embedding, GroupNorm, Init, LayerNorm, Linear, Module, RmsNorm, VarBuilder, - embedding, group_norm, linear, linear_b, ops::sigmoid, rms_norm, + Activation, Conv1d, Conv2d, Embedding, GroupNorm, Init, LayerNorm, Linear, Module, RmsNorm, + VarBuilder, embedding, group_norm, linear, linear_b, ops::sigmoid, rms_norm, }; +use rand::Rng; use crate::{ models::{ + bigvgan::{BigVGAN, config::BigVGANConfig}, campplus::CAMPPlus, common::{ - GateUpDownMLP, QKVCatAttention, TwoLinearMLP, WNConv1d, WNLinear, get_conv1d, - get_layer_norm, get_layer_norm_without_weight, mish, + GEGLU, GLU, GPT2Model, GateUpDownMLP, QKVCatAttention, TwoLinearMLP, WNConv1d, + WNLinear, eager_attention_forward, get_conv1d, get_conv2d, get_layer_norm, + get_layer_norm_without_weight, mish, }, feature_extractor::seamless_m4t_feature_extractor::SeamlessM4TFeatureExtractor, - index_tts2::config::{DiTModelArgs, IndexTTS2Config, PreprocessParams, S2MelConfig}, + index_tts2::config::{ + DiTModelArgs, GptConfig, IndexTTS2Config, PreprocessParams, S2MelConfig, + }, mask_gct::model::RepCodec, w2v_bert_2_0::model::W2VBert2_0Model, }, - position_embed::rope::RoPE, + position_embed::rope::{RoPE, compute_default_rope_parameters}, utils::{ audio_utils::{ - create_hann_window, get_waveform_and_window_properties, kaldi_fbank, - kaldi_get_mel_banks, mel_filter_bank, torch_stft, + create_hann_window, extract_audio_url, get_waveform_and_window_properties, kaldi_fbank, + kaldi_get_mel_banks, load_audio, mel_filter_bank, resample_simple, torch_stft, }, - get_vb_model_path, read_pth_tensor_info_cycle, + get_dtype, get_logit_processor, get_vb_model_path, load_tensor_from_pt, + read_pth_tensor_info_cycle, tensor_utils::{ - interpolate_nearest_1d, pad_reflect_last_dim, sequence_mask, split_tensor_with_size, + cosine_similarity, interpolate_nearest_1d, l2_normalize, linspace, masked_fill_zeros, + pad_reflect_last_dim, sequence_mask, split_tensor, }, }, }; pub struct AdaptiveLayerNorm { project_layer: Linear, norm: RmsNorm, - d_model: usize, + // d_model: usize, } impl AdaptiveLayerNorm { @@ -42,14 +52,15 @@ impl AdaptiveLayerNorm { Ok(Self { project_layer, norm, - d_model, + // d_model, }) } pub fn forward(&self, xs: &Tensor, embedding: Option<&Tensor>) -> Result { if let Some(embedding) = embedding { let emb = self.project_layer.forward(embedding)?; - let emb_split = split_tensor_with_size(&emb, 2, D::Minus1)?; + // let emb_split = split_tensor_with_size(&emb, 2, D::Minus1)?; + let emb_split = emb.chunk(2, D::Minus1)?; let weight = &emb_split[0]; let bias = &emb_split[1]; Ok(self @@ -170,7 +181,7 @@ impl DiTTransformer { layers.push(layer); } let norm = AdaptiveLayerNorm::new(vb.pp("norm"), config.dim, config.norm_eps)?; - let rope = RoPE::new(config.dim, 10000.0, vb.device())?; + let rope = RoPE::new(config.head_dim, 10000.0, vb.device())?; let mut layers_emit_skip: Vec = vec![]; let mut layers_receive_skip: Vec = vec![]; if config.uvit_skip_connection { @@ -195,7 +206,7 @@ impl DiTTransformer { let (cos, sin) = self.rope.forward(0, seq_len, xs.device())?; let mut skip_in_x_list = vec![]; let mut xs = xs.clone(); - for (i, layer) in (&self.layers).iter().enumerate() { + for (i, layer) in self.layers.iter().enumerate() { let skip_in_x = if self.uvit_skip_connection && self.layers_receive_skip.contains(&i) { skip_in_x_list.pop() } else { @@ -252,7 +263,7 @@ impl TimestepEmbedder { .unsqueeze(D::Minus1)? .broadcast_matmul(&self.freqs.unsqueeze(0)?)?; let mut embedding = Tensor::cat(&[args.cos()?, args.sin()?], D::Minus1)?; - if self.frequency_embedding_size % 2 > 0 { + if !self.frequency_embedding_size.is_multiple_of(2) { embedding = embedding.pad_with_zeros(D::Minus1, 0, 1)?; } embedding = self.mlp.forward(&embedding)?; @@ -390,15 +401,16 @@ impl Wavenet { input_a: &Tensor, input_b: &Tensor, ) -> Result { - let in_act = input_a.add(&input_b)?; - let parts = split_tensor_with_size(&in_act, 2, 1)?; - let t_act = (&parts[0]).tanh()?; + let in_act = input_a.broadcast_add(input_b)?; + let parts = in_act.chunk(2, 1)?; + let t_act = parts[0].tanh()?; let s_act = sigmoid(&parts[1])?; let acts = t_act.mul(&s_act)?; Ok(acts) } pub fn forward(&self, xs: &Tensor, x_mask: &Tensor, g: Option<&Tensor>) -> Result { + let x_mask = x_mask.to_dtype(xs.dtype())?; let mut output = xs.zeros_like()?; let g = if let Some(g) = g && let Some(cond_layer) = &self.cond_layer @@ -409,7 +421,7 @@ impl Wavenet { }; let mut xs = xs.clone(); for i in 0..self.n_layers { - let xs_in = &self.in_layers[i].forward(&xs)?; + let xs_in = self.in_layers[i].forward(&xs)?; let g_l = if let Some(g) = &g { let cond_offset = i * 2 * self.hidden_c; g.narrow(1, cond_offset, 2 * self.hidden_c)? @@ -421,13 +433,13 @@ impl Wavenet { if i < self.n_layers - 1 { let res_acts = res_skip_act.narrow(1, 0, self.hidden_c)?; let out_acts = res_skip_act.narrow(1, self.hidden_c, self.hidden_c)?; - xs = xs.add(&res_acts)?.mul(x_mask)?; + xs = xs.add(&res_acts)?.broadcast_mul(&x_mask)?; output = output.add(&out_acts)?; } else { - output = output.add(&res_skip_act)?; + output = output.add(res_skip_act)?; } } - output = output.mul(x_mask)?; + output = output.broadcast_mul(&x_mask)?; Ok(output) } } @@ -472,12 +484,13 @@ impl FinalLayer { .unsqueeze(1)? .affine(1.0, 1.0)? .broadcast_mul(&xs)? - .add(&linear_c[0].unsqueeze(1)?)?; + .broadcast_add(&linear_c[0].unsqueeze(1)?)?; let xs = self.linear.forward(&xs)?; Ok(xs) } } +#[allow(unused)] pub struct DiT { transformer: DiTTransformer, x_embedder: WNLinear, @@ -634,42 +647,45 @@ impl DiT { && let Some(style) = style { let style = style.unsqueeze(1)?.repeat((1, t_dim, 1))?; - x_in = Tensor::cat(&[&x_in, &style], D::Minus1)?; + x_in = Tensor::cat(&[&x_in, &style], D::Minus1)?.contiguous()?; } x_in = self.cond_x_merge_linear.forward(&x_in)?; - // if self.style_as_token - // && let Some(style_in) = self.style_in.as_ref() - // { - // let style = style_in.forward(style)?.unsqueeze(1)?; - // x_in = Tensor::cat(&[&style, &x_in], 1)?; - // } - // if self.time_as_token { - // let t1 = t1.unsqueeze(1)?; - // x_in = Tensor::cat(&[&t1, &x_in], 1)?; - // } - // let mut x_lens = x_lens.clone(); - // if self.style_as_token { - // x_lens = x_lens.affine(1.0, 1.0)?; - // } - // if self.time_as_token { - // x_lens = x_lens.affine(1.0, 0.0)?; - // } + if self.style_as_token + && let Some(style_in) = self.style_in.as_ref() + && let Some(style) = style + { + let style = style_in.forward(style)?.unsqueeze(1)?; + x_in = Tensor::cat(&[&style, &x_in], 1)?.contiguous()?; + } + if self.time_as_token { + let t1 = t1.unsqueeze(1)?; + x_in = Tensor::cat(&[&t1, &x_in], 1)?.contiguous()?; + } + let mut x_lens = x_lens.clone(); + if self.style_as_token { + x_lens = x_lens.affine(1.0, 1.0)?; + } + if self.time_as_token { + x_lens = x_lens.affine(1.0, 1.0)?; + } let x_mask = sequence_mask(&x_lens, Some(x_in.dim(1)? as u32))? .to_device(xs.device())? .unsqueeze(1)?; - let mut x_res = self.transformer.forward(&x_in, &t1.unsqueeze(1)?, None)?; - // if self.time_as_token { - // let last_dim = x_res.dim(D::Minus1)?; - // x_res = x_res.narrow(D::Minus1, 1, last_dim - 1)?; - // } - // if self.style_as_token { - // let last_dim = x_res.dim(D::Minus1)?; - // x_res = x_res.narrow(D::Minus1, 1, last_dim - 1)?; - // } + let mut x_res = self + .transformer + .forward(&x_in, &t1.unsqueeze(1)?, Some(&x_mask))?; + if self.time_as_token { + let last_dim = x_res.dim(D::Minus1)?; + x_res = x_res.narrow(D::Minus1, 1, last_dim - 1)?.contiguous()?; + } + if self.style_as_token { + let last_dim = x_res.dim(D::Minus1)?; + x_res = x_res.narrow(D::Minus1, 1, last_dim - 1)?.contiguous()?; + } if self.long_skip_connection { x_res = self .skip_linear - .forward(&Tensor::cat(&[&x_res, &xs], D::Minus1)?)?; + .forward(&Tensor::cat(&[&x_res, &xs], D::Minus1)?.contiguous()?)?; } let xs = self.conv1.forward(&x_res)?; let xs = xs.transpose(1, 2)?; @@ -689,20 +705,90 @@ pub struct CFM { in_channels: usize, estimator: DiT, // criterion: l1Loss - sigma_min: f32, + // sigma_min: f32, } impl CFM { pub fn new(vb: VarBuilder, config: &S2MelConfig) -> Result { let in_channels = config.di_t.in_channels; - let sigma_min = 1e-6; + // let sigma_min = 1e-6; let estimator = DiT::new(vb.pp("estimator"), config)?; Ok(Self { in_channels, estimator, - sigma_min, + // sigma_min, }) } + + pub fn inference( + &self, + mu: &Tensor, + x_lens: &Tensor, + prompt: &Tensor, + style: &Tensor, + n_timesteps: usize, + inference_cfg_rate: f64, + ) -> Result { + let (b, t, _) = mu.dims3()?; + let z = Tensor::randn(0.0f32, 1.0f32, (b, self.in_channels, t), mu.device())?; + let t_span = linspace(0.0, 1.0, n_timesteps + 1, mu.device())?; + let res = self.solve_euler(&z, x_lens, prompt, mu, style, &t_span, inference_cfg_rate)?; + Ok(res) + } + + pub fn solve_euler( + &self, + x: &Tensor, + x_lens: &Tensor, + prompt: &Tensor, + mu: &Tensor, + style: &Tensor, + t_span: &Tensor, + inference_cfg_rate: f64, + ) -> Result { + let mut x: Tensor = x.clone(); + let prompt_len = prompt.dim(D::Minus1)?; + let prompt_x = Tensor::zeros_like(&x)?; + prompt_x.slice_set(prompt, D::Minus1, 0)?; + let x_0 = Tensor::zeros_like(&x)? + .narrow(D::Minus1, 0, prompt_len)? + .contiguous()?; + x.slice_set(&x_0, D::Minus1, 0)?; + let mut t = t_span.i(0)?; + let t_len = t_span.dim(0)?; + let mut res_x = x.clone(); + for step in 1..t_len { + let dt = t_span.i(step)?.sub(&t_span.i(step - 1)?)?; + let dphi_dt = if inference_cfg_rate > 0.0 { + let stacked_prompt_x = + Tensor::cat(&[&prompt_x, &Tensor::zeros_like(&prompt_x)?], 0)?; + let stacked_style = Tensor::cat(&[style, &Tensor::zeros_like(style)?], 0)?; + let stacked_mu = Tensor::cat(&[mu, &Tensor::zeros_like(mu)?], 0)?; + let stacked_x = Tensor::cat(&[&x, &x], 0)?; + let stacked_t = Tensor::cat(&[&t.unsqueeze(0)?, &t.unsqueeze(0)?], 0)?; + let stacked_dphi_dt = self.estimator.forward( + &stacked_x, + &stacked_prompt_x, + x_lens, + &stacked_t, + Some(&stacked_style), + &stacked_mu, + )?; + let dphi = stacked_dphi_dt.chunk(2, 0)?; + dphi[0] + .affine(1.0 + inference_cfg_rate, 0.0)? + .sub(&dphi[1].affine(inference_cfg_rate, 0.0)?)? + } else { + self.estimator + .forward(&x, &prompt_x, x_lens, &t.unsqueeze(0)?, Some(style), mu)? + }; + x = x.add(&dphi_dt.broadcast_mul(&dt)?)?; + res_x = x.clone(); + t = t.add(&dt)?; + x.slice_set(&x_0, D::Minus1, 0)?; + } + Ok(res_x) + } } pub struct InterpolateModule { @@ -726,6 +812,7 @@ impl InterpolateModule { } } +#[allow(unused)] pub struct InterpolateRegulator { sampling_ratios: Vec, out_channels: usize, @@ -733,12 +820,13 @@ pub struct InterpolateRegulator { model_12: Conv1d, embedding: Embedding, mask_token: Tensor, - quantizer_dropout: f32, + // quantizer_dropout: f32, content_in_proj: Linear, n_codebooks: usize, interpolate: bool, } +#[allow(unused)] impl InterpolateRegulator { pub fn new( vb: VarBuilder, @@ -784,7 +872,7 @@ impl InterpolateRegulator { model_12, embedding, mask_token, - quantizer_dropout, + // quantizer_dropout, content_in_proj, n_codebooks, interpolate, @@ -805,9 +893,36 @@ impl InterpolateRegulator { } } +pub struct GptLayer { + layer_0: Linear, + layer_1: Linear, + layer_2: Linear, +} + +impl GptLayer { + pub fn new(vb: VarBuilder) -> Result { + let layer_0 = linear(1280, 256, vb.pp("0"))?; + let layer_1 = linear(256, 128, vb.pp("1"))?; + let layer_2 = linear(128, 1024, vb.pp("2"))?; + Ok(Self { + layer_0, + layer_1, + layer_2, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.layer_0.forward(xs)?; + let xs = self.layer_1.forward(&xs)?; + let xs = self.layer_2.forward(&xs)?; + Ok(xs) + } +} + pub struct MyModel { cfm: CFM, length_regulator: InterpolateRegulator, + gpt_layer: GptLayer, } impl MyModel { @@ -839,9 +954,13 @@ impl MyModel { let cfm_dict = read_pth_tensor_info_cycle(s2mel_path.clone(), Some("net.cfm"))?; let cfm_vb = VarBuilder::from_tensors(cfm_dict, dtype, device); let cfm = CFM::new(cfm_vb, config)?; + let gpt_layer_dict = read_pth_tensor_info_cycle(s2mel_path.clone(), Some("net.gpt_layer"))?; + let gpt_layer_vb = VarBuilder::from_tensors(gpt_layer_dict, dtype, device); + let gpt_layer = GptLayer::new(gpt_layer_vb)?; Ok(Self { cfm, length_regulator, + gpt_layer, }) } @@ -855,17 +974,1022 @@ impl MyModel { } } -pub struct IndexTTS2Cache { +pub struct RelPositionMultiHeadedAttention { + q_proj: Linear, + k_proj: Linear, + v_proj: Linear, + o_proj: Linear, + num_heads: usize, + head_dim: usize, + linear_pos: Linear, + pos_bias_u: Tensor, + pos_bias_v: Tensor, + kv_cache: Option<(Tensor, Tensor)>, +} + +impl RelPositionMultiHeadedAttention { + pub fn new(vb: VarBuilder, n_head: usize, n_feat: usize) -> Result { + let head_dim = n_feat / n_head; + let q_proj = linear_b(n_feat, n_feat, true, vb.pp("linear_q"))?; + let k_proj = linear_b(n_feat, n_feat, true, vb.pp("linear_k"))?; + let v_proj = linear_b(n_feat, n_feat, true, vb.pp("linear_v"))?; + let o_proj = linear_b(n_feat, n_feat, true, vb.pp("linear_out"))?; + let linear_pos = linear_b(n_feat, n_feat, false, vb.pp("linear_pos"))?; + let pos_bias_u = vb + .get_with_hints((n_head, head_dim), "pos_bias_u", Init::Const(0.0))? + .unsqueeze(0)? + .unsqueeze(0)?; + let pos_bias_v = vb + .get_with_hints((n_head, head_dim), "pos_bias_v", Init::Const(0.0))? + .unsqueeze(0)? + .unsqueeze(0)?; + Ok(Self { + q_proj, + k_proj, + v_proj, + o_proj, + num_heads: n_head, + head_dim, + linear_pos, + pos_bias_u, + pos_bias_v, + kv_cache: None, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + pos_emb: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let (b_sz, q_len, _) = xs.dims3()?; + let query_states = self.q_proj.forward(xs)?; + let key_states = self.k_proj.forward(xs)?; + let value_states = self.v_proj.forward(xs)?; + let query_states = query_states.reshape((b_sz, q_len, self.num_heads, self.head_dim))?; + let q_with_bias_u = query_states + .broadcast_add(&self.pos_bias_u)? + .transpose(1, 2)? + .contiguous()?; + let q_with_bias_v = query_states + .broadcast_add(&self.pos_bias_v)? + .transpose(1, 2)? + .contiguous()?; + let key_states = key_states + .reshape((b_sz, q_len, self.num_heads, self.head_dim))? + .transpose(1, 2)? + .contiguous()?; + let value_states = value_states + .reshape((b_sz, q_len, self.num_heads, self.head_dim))? + .transpose(1, 2)? + .contiguous()?; + let n_batch_pos = pos_emb.dim(0)?; + let p = self + .linear_pos + .forward(pos_emb)? + .reshape((n_batch_pos, (), self.num_heads, self.head_dim))? + .transpose(1, 2)? + .contiguous()?; + let matrix_ac = + q_with_bias_u.matmul(&key_states.transpose(D::Minus2, D::Minus1)?.contiguous()?)?; + let matrix_bd = q_with_bias_v.matmul(&p.transpose(D::Minus2, D::Minus1)?.contiguous()?)?; + let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let scores = matrix_ac.add(&matrix_bd)?.affine(scale, 0.0)?; + let attn_weights = match attention_mask { + None => scores, + Some(mask) => scores.broadcast_add(&mask.to_dtype(scores.dtype())?)?, + }; + let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?; + let attn_output = attn_weights.matmul(&value_states)?; + let attn_output = attn_output.transpose(1, 2)?.contiguous()?; + let attn_output = attn_output + .reshape((b_sz, q_len, self.num_heads * self.head_dim))? + .contiguous()?; + let attn_output = attn_output.apply(&self.o_proj)?; + Ok(attn_output) + } + + pub fn forward_with_cache( + &mut self, + xs: &Tensor, + pos_emb: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let (b_sz, q_len, _) = xs.dims3()?; + let query_states = self.q_proj.forward(xs)?; + let key_states = self.k_proj.forward(xs)?; + let value_states = self.v_proj.forward(xs)?; + let query_states = query_states.reshape((b_sz, q_len, self.num_heads, self.head_dim))?; + let q_with_bias_u = query_states + .broadcast_add(&self.pos_bias_u)? + .transpose(1, 2)?; + let q_with_bias_v = query_states + .broadcast_add(&self.pos_bias_v)? + .transpose(1, 2)?; + let key_states = key_states + .reshape((b_sz, q_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let value_states = value_states + .reshape((b_sz, q_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let (key_states, value_states) = match &self.kv_cache { + None => (key_states, value_states), + Some((prev_k, prev_v)) => { + let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; + let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; + (key_states, value_states) + } + }; + + self.kv_cache = Some((key_states.clone(), value_states.clone())); + let n_batch_pos = pos_emb.dim(0)?; + let p = self.linear_pos.forward(pos_emb)?.reshape(( + n_batch_pos, + (), + self.num_heads, + self.head_dim, + ))?; + let matrix_ac = q_with_bias_u.matmul(&key_states.transpose(D::Minus2, D::Minus1)?)?; + let matrix_bd = q_with_bias_v.matmul(&p.transpose(D::Minus2, D::Minus1)?)?; + let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let scores = matrix_ac.add(&matrix_bd)?.affine(scale, 0.0)?; + let attn_weights = match attention_mask { + None => scores, + Some(mask) => scores.broadcast_add(&mask.to_dtype(scores.dtype())?)?, + }; + let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?; + let attn_output = attn_weights.matmul(&value_states)?; + let attn_output = attn_output.transpose(1, 2)?.contiguous()?; + let attn_output = attn_output.reshape((b_sz, q_len, self.num_heads * self.head_dim))?; + let attn_output = attn_output.apply(&self.o_proj)?; + Ok(attn_output) + } + + pub fn clear_kv_cache(&mut self) { + self.kv_cache = None + } +} + +pub struct ConvolutionModule { + pointwise_conv1: Conv1d, + glu: GLU, + depthwise_conv: Conv1d, + norm: LayerNorm, + pointwise_conv2: Conv1d, + activation: Activation, +} + +impl ConvolutionModule { + pub fn new( + vb: VarBuilder, + channels: usize, + kernel_size: usize, + activation: Activation, + bias: bool, + ) -> Result { + let pointwise_conv1 = get_conv1d( + vb.pp("pointwise_conv1"), + channels, + 2 * channels, + 1, + 0, + 1, + 1, + 1, + bias, + )?; + let padding = (kernel_size - 1) / 2; + let depthwise_conv = get_conv1d( + vb.pp("depthwise_conv"), + channels, + channels, + kernel_size, + padding, + 1, + 1, + channels, + bias, + )?; + let norm = get_layer_norm(vb.pp("norm"), 1e-5, channels, true)?; + let pointwise_conv2 = get_conv1d( + vb.pp("pointwise_conv2"), + channels, + channels, + 1, + 0, + 1, + 1, + 1, + bias, + )?; + let glu = GLU::new(1)?; + Ok(Self { + pointwise_conv1, + glu, + depthwise_conv, + norm, + pointwise_conv2, + activation, + }) + } + + pub fn forward(&self, xs: &Tensor, mask_pad: Option<&Tensor>) -> Result { + let mut xs = xs.transpose(1, 2)?; + if let Some(mask_pad) = mask_pad { + xs = masked_fill_zeros(&xs, mask_pad)?; + } + xs = self.pointwise_conv1.forward(&xs)?; + xs = self.glu.forward(&xs)?; + xs = self.depthwise_conv.forward(&xs)?; + xs = xs.transpose(1, 2)?; + xs = self.norm.forward(&xs)?.apply(&self.activation)?; + xs = xs.transpose(1, 2)?; + xs = self.pointwise_conv2.forward(&xs)?; + if let Some(mask_pad) = mask_pad { + xs = masked_fill_zeros(&xs, mask_pad)?; + } + xs = xs.transpose(1, 2)?; + Ok(xs) + } +} + +pub struct ConformerEncoderLayer { + self_attn: RelPositionMultiHeadedAttention, + feed_forward: TwoLinearMLP, // silu, + // feed_forward_macaron: None, + ff_scale: f32, + conv_module: ConvolutionModule, + norm_ff: LayerNorm, + norm_mha: LayerNorm, + norm_conv: LayerNorm, + norm_final: LayerNorm, + // concat_linear: None, +} + +impl ConformerEncoderLayer { + pub fn new( + vb: VarBuilder, + attention_heads: usize, + output_size: usize, + linear_units: usize, + ) -> Result { + let self_attn = + RelPositionMultiHeadedAttention::new(vb.pp("self_attn"), attention_heads, output_size)?; + let feed_forward = TwoLinearMLP::new( + vb.pp("feed_forward"), + output_size, + linear_units, + output_size, + Activation::Silu, + true, + "w_1", + "w_2", + )?; + let ff_scale = 1.0; + let conv_module = ConvolutionModule::new( + vb.pp("conv_module"), + output_size, + 15, + Activation::Silu, + true, + )?; + let norm_ff = get_layer_norm(vb.pp("norm_ff"), 1e-5, output_size, true)?; + let norm_mha = get_layer_norm(vb.pp("norm_mha"), 1e-5, output_size, true)?; + let norm_conv = get_layer_norm(vb.pp("norm_conv"), 1e-5, output_size, true)?; + let norm_final = get_layer_norm(vb.pp("norm_final"), 1e-5, output_size, true)?; + Ok(Self { + self_attn, + feed_forward, + ff_scale, + conv_module, + norm_ff, + norm_mha, + norm_conv, + norm_final, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + mask: Option<&Tensor>, + pos_emb: &Tensor, + mask_pad: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let mut xs = self.norm_mha.forward(xs)?; + xs = self.self_attn.forward(&xs, pos_emb, mask)?; + xs = residual.add(&xs)?; + let residual = xs.clone(); + xs = self.norm_conv.forward(&xs)?; + xs = self.conv_module.forward(&xs, mask_pad)?; + xs = residual.add(&xs)?; + let residual = xs.clone(); + xs = self.norm_ff.forward(&xs)?; + xs = self + .feed_forward + .forward(&xs)? + .affine(self.ff_scale as f64, 1.0)?; + xs = residual.add(&xs)?; + xs = self.norm_final.forward(&xs)?; + Ok(xs) + } +} + +pub struct RelPositionalEncoding { + xscale: f64, + inv_freq: Tensor, +} + +impl RelPositionalEncoding { + pub fn new(d_model: usize, device: &Device) -> Result { + let xscale = (d_model as f64).sqrt(); + let inv_freq = compute_default_rope_parameters(d_model, 10000.0); + let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?; + + Ok(Self { xscale, inv_freq }) + } + pub fn forward(&self, xs: &Tensor, offset: usize) -> Result<(Tensor, Tensor)> { + let xs = xs.affine(self.xscale, 0.0)?; + let seq_len = xs.dim(1)?; + let positions = + Tensor::arange(offset as f32, seq_len as f32, xs.device())?.reshape((seq_len, 1))?; // (max_len, 1) + let freqs = positions.matmul(&self.inv_freq.to_device(xs.device())?)?; // (max_len, dim / 2) + let sin = freqs.sin()?; + let cos = freqs.cos()?; + let pe = Tensor::stack(&[sin, cos], D::Minus1)? + .flatten(1, 2)? + .unsqueeze(0)? + .to_dtype(xs.dtype())?; + Ok((xs, pe)) + } +} + +pub struct Conv2dSubsampling2 { + conv_0: Conv2d, // conv+relu + out_0: Linear, + pos_enc: RelPositionalEncoding, +} + +impl Conv2dSubsampling2 { + pub fn new(vb: VarBuilder, in_dim: usize, out_dim: usize) -> Result { + let conv_0 = get_conv2d(vb.pp("conv.0"), 1, out_dim, 3, 0, 2, 1, 1, true)?; + let out_0 = linear_b(out_dim * ((in_dim - 1) / 2), out_dim, true, vb.pp("out.0"))?; + let pos_enc = RelPositionalEncoding::new(out_dim, vb.device())?; + Ok(Self { + conv_0, + out_0, + pos_enc, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + mask: Option<&Tensor>, + offset: usize, + ) -> Result<(Tensor, Tensor, Option)> { + let xs = xs.unsqueeze(1)?; + let xs = self.conv_0.forward(&xs)?.relu()?; + let (b, c, t, f) = xs.dims4()?; + let xs = xs + .transpose(1, 2)? + .contiguous()? + .reshape((b, t, c * f))? + .apply(&self.out_0)?; + let (xs, pos_emb) = self.pos_enc.forward(&xs, offset)?; + let mask = if let Some(mask) = mask { + let t = mask.dim(2)?; + let mask_index = Tensor::arange_step(2u32, t as u32, 2u32, xs.device())?; + let mask = mask.index_select(&mask_index, 2)?; + Some(mask) + } else { + None + }; + + Ok((xs, pos_emb, mask)) + } +} + +pub struct ConformerEncoder { + encoders: Vec, + embed: Conv2dSubsampling2, + after_norm: LayerNorm, +} + +impl ConformerEncoder { + pub fn new( + vb: VarBuilder, + input_size: usize, + attention_heads: usize, + output_size: usize, + linear_units: usize, + num_blocks: usize, + ) -> Result { + let mut encoders = vec![]; + let vb_encoders = vb.pp("encoders"); + for i in 0..num_blocks { + let layer = ConformerEncoderLayer::new( + vb_encoders.pp(i), + attention_heads, + output_size, + linear_units, + )?; + encoders.push(layer); + } + let embed = Conv2dSubsampling2::new(vb.pp("embed"), input_size, output_size)?; + let after_norm = get_layer_norm(vb.pp("after_norm"), 1e-5, output_size, true)?; + Ok(Self { + encoders, + embed, + after_norm, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result<(Tensor, Option)> { + let (mut xs, pos_emb, mask) = self.embed.forward(xs, None, 0)?; + for layer in self.encoders.iter() { + xs = layer.forward(&xs, mask.as_ref(), &pos_emb, mask.as_ref())?; + } + xs = self.after_norm.forward(&xs)?; + Ok((xs, None)) + } +} + +pub struct PerceiverResamplerAttention { + to_q: Linear, + to_kv: Linear, + to_out: Linear, + n_head: usize, + head_dim: usize, + cross_attn_include_queries: bool, +} + +impl PerceiverResamplerAttention { + pub fn new( + vb: VarBuilder, + dim: usize, + dim_context: Option, + head_dim: usize, + n_head: usize, + cross_attn_include_queries: bool, + ) -> Result { + let dim_inner = head_dim * n_head; + let dim_context = dim_context.unwrap_or(dim); + let to_q = linear_b(dim, dim_inner, false, vb.pp("to_q"))?; + let to_kv = linear_b(dim_context, dim_inner * 2, false, vb.pp("to_kv"))?; + let to_out = linear_b(dim_inner, dim, false, vb.pp("to_out"))?; + Ok(Self { + to_q, + to_kv, + to_out, + n_head, + head_dim, + cross_attn_include_queries, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + context: Option<&Tensor>, + mask: Option<&Tensor>, + ) -> Result { + let b_sz = xs.dim(0)?; + let query_states = self.to_q.forward(xs)?; + let context = if let Some(context) = context + && self.cross_attn_include_queries + { + Tensor::cat(&[xs, context], D::Minus2)? + } else { + xs.clone() + }; + let key_value = self.to_kv.forward(&context)?.chunk(2, D::Minus1)?; + let query_states = query_states + .reshape((b_sz, (), self.n_head, self.head_dim))? + .transpose(1, 2)?; + let key_states = key_value[0] + .reshape((b_sz, (), self.n_head, self.head_dim))? + .transpose(1, 2)?; + let value_states = key_value[1] + .reshape((b_sz, (), self.n_head, self.head_dim))? + .transpose(1, 2)?; + let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let attn_output = + eager_attention_forward(&query_states, &key_states, &value_states, None, mask, scale)?; + let attn_output = attn_output.reshape((b_sz, (), self.n_head * self.head_dim))?; + let attn_output = attn_output.apply(&self.to_out)?; + Ok(attn_output) + } +} + +pub struct PerceiverResamplerFeedForward { + net_0: Linear, + net_1: GEGLU, + net_2: Linear, +} + +impl PerceiverResamplerFeedForward { + pub fn new(vb: VarBuilder, dim: usize, mult: usize) -> Result { + let dim_inner = dim * mult * 2 / 3; + let net_0 = linear(dim, dim_inner * 2, vb.pp("0"))?; + let net_1 = GEGLU::new(2)?; + let net_2 = linear(dim_inner, dim, vb.pp("2"))?; + Ok(Self { + net_0, + net_1, + net_2, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.net_0.forward(xs)?; + let xs = self.net_1.forward(&xs)?; + let xs = self.net_2.forward(&xs)?; + Ok(xs) + } +} + +pub struct PerceiverResamplerLayer { + layer_0: PerceiverResamplerAttention, + layer_1: PerceiverResamplerFeedForward, +} + +impl PerceiverResamplerLayer { + pub fn new( + vb: VarBuilder, + dim: usize, + head_dim: usize, + n_head: usize, + ff_mult: usize, + ) -> Result { + let layer_0 = + PerceiverResamplerAttention::new(vb.pp("0"), dim, None, head_dim, n_head, true)?; + let layer_1 = PerceiverResamplerFeedForward::new(vb.pp("1"), dim, ff_mult)?; + Ok(Self { layer_0, layer_1 }) + } + + pub fn forward( + &self, + xs: &Tensor, + context: Option<&Tensor>, + mask: Option<&Tensor>, + ) -> Result { + let latents = self.layer_0.forward(xs, context, mask)?.add(xs)?; + let latents = self.layer_1.forward(&latents)?.add(&latents)?; + Ok(latents) + } +} + +pub struct PerceiverRmsNorm { + // dim: usize, + scale: f64, + gamma: Tensor, +} + +impl PerceiverRmsNorm { + pub fn new(vb: VarBuilder, dim: usize) -> Result { + let scale = (dim as f64).sqrt(); + let gamma = vb.get_with_hints(dim, "gamma", Init::Const(1.0))?; + Ok(Self { scale, gamma }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let out = l2_normalize(xs, xs.rank() - 1)? + .affine(self.scale, 0.0)? + .broadcast_mul(&self.gamma)?; + Ok(out) + } +} + +#[allow(unused)] +pub struct PerceiverResampler { + dim_context: usize, + proj_context: Linear, + latents: Tensor, + layers: Vec, + norm: PerceiverRmsNorm, +} + +impl PerceiverResampler { + pub fn new( + vb: VarBuilder, + dim: usize, + depth: usize, + dim_context: Option, + num_latents: usize, + head_dim: usize, + n_head: usize, + ff_mult: usize, + ) -> Result { + let dim_context = dim_context.unwrap_or(dim); + let proj_context = linear(dim_context, dim, vb.pp("proj_context"))?; + let latents = vb.get_with_hints((num_latents, dim), "latents", Init::Const(1.0))?; + let mut layers = vec![]; + let vb_layers = vb.pp("layers"); + for i in 0..depth { + let layer = + PerceiverResamplerLayer::new(vb_layers.pp(i), dim, head_dim, n_head, ff_mult)?; + layers.push(layer); + } + let norm = PerceiverRmsNorm::new(vb.pp("norm"), dim)?; + Ok(Self { + dim_context, + proj_context, + latents, + layers, + norm, + }) + } + + pub fn forward(&self, xs: &Tensor, mask: Option<&Tensor>) -> Result { + let b = xs.dim(0)?; + let xs = self.proj_context.forward(xs)?; + let mut latents = self.latents.unsqueeze(0)?.repeat((b, 1, 1))?; + for layer in self.layers.iter() { + latents = layer.forward(&latents, Some(&xs), mask)?; + } + let xs = self.norm.forward(&latents)?; + Ok(xs) + } +} + +pub struct LearnedPositionEmbeddings { + emb: Embedding, +} + +impl LearnedPositionEmbeddings { + pub fn new(vb: VarBuilder, seq_len: usize, model_dim: usize) -> Result { + let emb = embedding(seq_len, model_dim, vb.pp("emb"))?; + Ok(Self { emb }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let sl = match xs.rank() { + 1 => xs.dim(0)?, + 2 => xs.dim(1)?, + 3 => xs.dim(1)?, + _ => return Err(anyhow!(" only support xs rank 1, 2, 3")), + }; + let id = Tensor::arange(0, sl as u32, xs.device())?; + let embed = self.emb.forward(&id)?.unsqueeze(0)?; + Ok(embed) + } + pub fn get_fixed_embedding(&self, id: u32, device: &Device) -> Result { + let id = Tensor::new(id, device)?; + let embed = self.emb.forward(&id)?.unsqueeze(0)?; + Ok(embed) + } +} + +#[allow(unused)] +pub struct UnifiedVoice { + conditioning_encoder: ConformerEncoder, + perceiver_encoder: PerceiverResampler, + emo_conditioning_encoder: ConformerEncoder, + emo_perceiver_encoder: PerceiverResampler, + text_embedding: Embedding, + emo_layer: Linear, + emovec_layer: Linear, + mel_embedding: Embedding, + gpt: GPT2Model, + mel_pos_embedding: LearnedPositionEmbeddings, + text_pos_embedding: LearnedPositionEmbeddings, + final_norm: LayerNorm, + text_head: Linear, + mel_head: Linear, + speed_emb: Embedding, + start_text_token: u32, + stop_text_token: u32, + start_mel_token: u32, + stop_mel_token: u32, + max_mel_tokens: usize, + dtype: DType, + device: Device, +} + +impl UnifiedVoice { + pub fn new(vb: VarBuilder, cfg: &GptConfig) -> Result { + let conditioning_encoder = ConformerEncoder::new( + vb.pp("conditioning_encoder"), + 1024, + cfg.condition_module.attention_heads, + cfg.condition_module.output_size, + cfg.condition_module.linear_units, + cfg.condition_module.num_blocks, + )?; + let perceiver_encoder = PerceiverResampler::new( + vb.pp("perceiver_encoder"), + cfg.model_dim, + 2, + Some(cfg.condition_module.output_size), + 32, + 64, + cfg.condition_module.attention_heads, + cfg.condition_module.perceiver_mult, + )?; + let emo_conditioning_encoder = ConformerEncoder::new( + vb.pp("emo_conditioning_encoder"), + 1024, + cfg.emo_condition_module.attention_heads, + cfg.emo_condition_module.output_size, + cfg.emo_condition_module.linear_units, + cfg.emo_condition_module.num_blocks, + )?; + let emo_perceiver_encoder = PerceiverResampler::new( + vb.pp("emo_perceiver_encoder"), + 1024, + 2, + Some(cfg.emo_condition_module.output_size), + 1, + 64, + cfg.emo_condition_module.attention_heads, + cfg.emo_condition_module.perceiver_mult, + )?; + + let text_embedding = embedding( + cfg.number_text_tokens + 1, + cfg.model_dim, + vb.pp("text_embedding"), + )?; + let emo_layer = linear(cfg.model_dim, cfg.model_dim, vb.pp("emo_layer"))?; + let emovec_layer = linear(1024, cfg.model_dim, vb.pp("emovec_layer"))?; + let mel_embedding = embedding(cfg.number_mel_codes, cfg.model_dim, vb.pp("mel_embedding"))?; + + let gpt = GPT2Model::new( + vb.pp("gpt"), + cfg.model_dim, + cfg.heads, + cfg.layers, + mel_embedding.embeddings(), + )?; + let max_mel_seq_len = cfg.max_mel_tokens + 3; + let max_text_seq_len = cfg.max_text_tokens + 2; + let mel_pos_embedding = LearnedPositionEmbeddings::new( + vb.pp("mel_pos_embedding"), + max_mel_seq_len, + cfg.model_dim, + )?; + let text_pos_embedding = LearnedPositionEmbeddings::new( + vb.pp("text_pos_embedding"), + max_text_seq_len, + cfg.model_dim, + )?; + let final_norm = get_layer_norm(vb.pp("final_norm"), 1e-5, cfg.model_dim, true)?; + let text_head = linear( + cfg.model_dim, + cfg.number_text_tokens + 1, + vb.pp("text_head"), + )?; + let mel_head = linear(cfg.model_dim, cfg.number_mel_codes, vb.pp("mel_head"))?; + let speed_emb = embedding(2, cfg.model_dim, vb.pp("speed_emb"))?; + Ok(Self { + conditioning_encoder, + perceiver_encoder, + emo_conditioning_encoder, + emo_perceiver_encoder, + text_embedding, + emo_layer, + emovec_layer, + mel_embedding, + gpt, + mel_pos_embedding, + text_pos_embedding, + final_norm, + text_head, + mel_head, + speed_emb, + start_text_token: 0, + stop_text_token: 1, + start_mel_token: 8192, + stop_mel_token: 8193, + max_mel_tokens: cfg.max_mel_tokens, + dtype: vb.dtype(), + device: vb.device().clone(), + }) + } + + pub fn get_conditioning(&self, speech_conditioning_input: &Tensor) -> Result { + let (speech_conditioning_input, mask) = self + .conditioning_encoder + .forward(speech_conditioning_input)?; + let conds = self + .perceiver_encoder + .forward(&speech_conditioning_input, mask.as_ref())?; + Ok(conds) + } + + pub fn prepare_inputs( + &self, + conditional_latents: &Tensor, + text_inputs: &Tensor, + ) -> Result<(Tensor, Tensor)> { + let pad_left = + Tensor::new(vec![self.start_text_token], text_inputs.device())?.unsqueeze(0)?; + let pad_right = + Tensor::new(vec![self.stop_text_token], text_inputs.device())?.unsqueeze(0)?; + let text_input = Tensor::cat(&[&pad_left, text_inputs, &pad_right], D::Minus1)?; + let seq_len = text_input.dim(D::Minus1)?; + let text_input_pos = + Tensor::arange(0u32, seq_len as u32, text_inputs.device())?.unsqueeze(0)?; + let text_embedding = self.text_embedding.forward(&text_input)?; + let text_pos_embedding = self.text_pos_embedding.forward(&text_input_pos)?; + let text_emb = text_embedding.add(&text_pos_embedding)?; + let conds_text_emb = Tensor::cat(&[conditional_latents, &text_emb], 1)?; + let (bs, len, _) = conds_text_emb.dims3()?; + let fake_inputs = Tensor::ones((bs, len), DType::U32, conds_text_emb.device())?; + let start_mel_token = + Tensor::new(vec![self.start_mel_token], conds_text_emb.device())?.unsqueeze(0)?; + let fake_inputs = Tensor::cat(&[fake_inputs, start_mel_token], 1)?; + Ok((fake_inputs, conds_text_emb)) + } + + pub fn gpt_forward( + &mut self, + input_ids: &Tensor, + input_embed: Option<&Tensor>, + offset: u32, + ) -> Result { + let emb = if let Some(input_embed) = input_embed { + let mel_len = input_embed.dim(1)?; + let text_inputs = input_ids.i((.., mel_len))?; + let text_emb = self.mel_embedding.forward(&text_inputs)?.unsqueeze(1)?; + let text_pos_emb = self.mel_pos_embedding.forward(&text_emb)?; + let text_emb = text_emb.add(&text_pos_emb)?; + Tensor::cat(&[input_embed, &text_emb], 1)? + } else { + let emb = self.mel_embedding.forward(input_ids)?; + let emb_pos = self + .mel_pos_embedding + .get_fixed_embedding(offset, input_ids.device())? + .unsqueeze(0)?; + emb.add(&emb_pos)? + }; + let outputs = self.gpt.forward(&emb)?; + let outputs = self.final_norm.forward(&outputs)?; + let logits = self.mel_head.forward(&outputs)?; + let seq_len = logits.dim(1)?; + let logits = logits.narrow(1, seq_len - 1, 1)?; + + Ok(logits) + } + + pub fn generate(&mut self, input_ids: &Tensor, input_embed: &Tensor) -> Result { + let mut logit_processor = get_logit_processor(Some(0.8), Some(0.8), Some(30), 35329); + let mut generate_ids = vec![]; + let mut input_ids = input_ids.clone(); + let mut input_embed = Some(input_embed.clone()); + let mut offset = 0u32; + let mut seq_len = input_ids.dim(1)? as u32; + for _ in 0..self.max_mel_tokens { + let logits = self.gpt_forward(&input_ids, input_embed.as_ref(), offset)?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + if next_token == self.start_mel_token || next_token == self.stop_mel_token { + break; + } + generate_ids.push(next_token); + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + input_embed = None; + offset += seq_len; + seq_len = 1; + } + self.gpt.clear_kv_cache(); + let generate_ids = Tensor::new(generate_ids, &self.device)?; + Ok(generate_ids) + } + + pub fn inference_speech( + &mut self, + speech_condition: &Tensor, + text_inputs: &Tensor, + emo_vec: &Tensor, + ) -> Result<(Tensor, Tensor)> { + let speech_conditioning_latent = + self.get_conditioning(&speech_condition.to_dtype(self.dtype)?)?; + let text_len = text_inputs.dim(0)?; + let tmp = Tensor::zeros(text_len, DType::U32, speech_condition.device())?; + let duration_emb = self.speed_emb.forward(&tmp)?.unsqueeze(1)?; + let duration_emb_half = self + .speed_emb + .forward(&Tensor::ones_like(&tmp)?)? + .unsqueeze(1)?; + let speech_add_emovec = speech_conditioning_latent.broadcast_add(&emo_vec.unsqueeze(1)?)?; + let conds_latent = + Tensor::cat(&[&speech_add_emovec, &duration_emb_half, &duration_emb], 1)?; + let (fake_inputs, inputs_embeds) = self.prepare_inputs(&conds_latent, text_inputs)?; + let output = self.generate(&fake_inputs, &inputs_embeds)?; + Ok((output, speech_conditioning_latent)) + } + + pub fn get_emo_conditioning(&self, speech_conditioning_latent: &Tensor) -> Result { + let (speech_conditioning_input, mask) = self + .emo_conditioning_encoder + .forward(speech_conditioning_latent)?; + let conds = self + .emo_perceiver_encoder + .forward(&speech_conditioning_input, mask.as_ref())?; + let conds = conds.squeeze(1)?; + Ok(conds) + } + + pub fn get_emovec(&self, speech_conditioning_latent: &Tensor) -> Result { + let emo_vec_syn_ori = self.get_emo_conditioning(speech_conditioning_latent)?; + let emo_vec_syn = self.emovec_layer.forward(&emo_vec_syn_ori)?; + let emo_vec = self.emo_layer.forward(&emo_vec_syn)?; + Ok(emo_vec) + } + + pub fn merge_emovec( + &self, + speech_conditioning_latent: &Tensor, + emo_speech_conditioning_latent: &Tensor, + alpha: f64, + ) -> Result { + let emo_vec = self.get_emovec(&emo_speech_conditioning_latent.to_dtype(self.dtype)?)?; + let base_vec = self.get_emovec(&speech_conditioning_latent.to_dtype(self.dtype)?)?; + let out = emo_vec.sub(&base_vec)?.affine(alpha, 0.0)?.add(&base_vec)?; + Ok(out) + } + + pub fn get_logits( + &mut self, + speech_conditioning_inputs: &Tensor, + first_inputs: &Tensor, + second_inputs: &Tensor, + ) -> Result { + let emb = Tensor::cat( + &[speech_conditioning_inputs, first_inputs, second_inputs], + 1, + )?; + let gpt_out = self.gpt.forward(&emb)?; + self.gpt.clear_kv_cache(); + let offset = speech_conditioning_inputs.dim(1)?; + let enc = gpt_out.i((.., offset.., ..))?; + let enc = self.final_norm.forward(&enc)?; + let offset = first_inputs.dim(1)?; + let mel_logits = enc.i((.., offset.., ..))?; + Ok(mel_logits) + } + + pub fn forward( + &mut self, + speech_conditioning_latent: &Tensor, + text_inputs: &Tensor, + mel_codes: &Tensor, + emo_vec: &Tensor, + use_speed: &Tensor, + ) -> Result { + let pad_left = + Tensor::new(vec![self.start_text_token], text_inputs.device())?.unsqueeze(0)?; + let pad_right = + Tensor::new(vec![self.stop_text_token], text_inputs.device())?.unsqueeze(0)?; + let text_inputs = Tensor::cat(&[&pad_left, text_inputs, &pad_right], D::Minus1)?; + let pad_left = + Tensor::new(vec![self.start_mel_token], text_inputs.device())?.unsqueeze(0)?; + let pad_right = + Tensor::new(vec![self.stop_mel_token], text_inputs.device())?.unsqueeze(0)?; + let mel_codes = Tensor::cat(&[&pad_left, mel_codes, &pad_right], D::Minus1)?; + let duration_emb = self + .speed_emb + .forward(&Tensor::zeros_like(use_speed)?)? + .unsqueeze(1)?; + let duration_emb_half = self + .speed_emb + .forward(&Tensor::ones_like(use_speed)?)? + .unsqueeze(1)?; + let speech_add_emovec = speech_conditioning_latent.broadcast_add(&emo_vec.unsqueeze(1)?)?; + let conds = Tensor::cat(&[&speech_add_emovec, &duration_emb_half, &duration_emb], 1)?; + let text_emb = self.text_embedding.forward(&text_inputs)?; + let text_emb_pos = self.text_pos_embedding.forward(&text_inputs)?; + let text_emb = text_emb.add(&text_emb_pos)?; + let mel_emb = self.mel_embedding.forward(&mel_codes)?; + let mel_emb_pos = self.mel_pos_embedding.forward(&mel_codes)?; + let mel_emb = mel_emb.add(&mel_emb_pos)?; + let mel_logits = self.get_logits(&conds, &text_emb, &mel_emb)?; + let len = mel_logits.dim(1)? - 2; + let mel_logits = mel_logits.i((.., 0..len, ..))?; + Ok(mel_logits) + } +} + +pub struct IndexTTS2SpkCache { pub cache_spk_cond: Tensor, pub cache_s2mel_style: Tensor, pub cache_s2mel_prompt: Tensor, pub cache_mel: Tensor, - // pub cache_emo_cond: Tensor, - // pub cache_emo_audio_prompt: Tensor, + pub cache_spk_audio_prompt: String, +} + +pub struct IndexTTS2EmoCache { + cache_emo_audio_prompt: String, + emo_cond_emb: Tensor, } pub struct IndexTTS2Model { - cache: Option, + max_audio_length_seconds: usize, + spk_cache: Option, + emo_cache: Option, feature_extractor: SeamlessM4TFeatureExtractor, semantic_model: W2VBert2_0Model, semantic_mean: Tensor, @@ -880,6 +2004,14 @@ pub struct IndexTTS2Model { mel_energies: Tensor, campplus_model: CAMPPlus, s2mel: MyModel, + emo_num: Vec, + emo_matrix: Vec, + spk_matrix: Vec, + gpt: UnifiedVoice, + bigvgan: BigVGAN, + device: Device, + dtype_f32: DType, + // gpt_dtype: DType, } impl IndexTTS2Model { @@ -888,8 +2020,9 @@ impl IndexTTS2Model { save_dir: &str, config: &IndexTTS2Config, device: &Device, - dtype: DType, ) -> Result { + let dtype_f32 = DType::F32; + let gpt_dtype = get_dtype(None, "bfloat16"); let feature_extractor = SeamlessM4TFeatureExtractor::new( // 80, 80, @@ -900,23 +2033,24 @@ impl IndexTTS2Model { device, )?; let w2vbert2_path = save_dir.to_string() + "/facebook/w2v-bert-2.0"; - let semantic_model = W2VBert2_0Model::init(&w2vbert2_path, device, dtype)?; + let semantic_model = W2VBert2_0Model::init(&w2vbert2_path, device, dtype_f32)?; let semantic_mean_var_path = path.to_string() + "/" + &config.w2v_stat; let dict = read_all_with_key(semantic_mean_var_path, None)?; - let mut semantic_mean = Tensor::new(0.0, device)?.to_dtype(dtype)?; - let mut semantic_std = Tensor::new(1.0, device)?.to_dtype(dtype)?; + let mut semantic_mean = Tensor::new(0.0, device)?.to_dtype(dtype_f32)?; + let mut semantic_std = Tensor::new(1.0, device)?.to_dtype(dtype_f32)?; for (k, v) in dict { if k.eq("mean") { - semantic_mean = v.to_device(device)?.to_dtype(dtype)?; + semantic_mean = v.to_device(device)?.to_dtype(dtype_f32)?; } else if k.eq("var") { - semantic_std = v.to_device(device)?.to_dtype(dtype)?.sqrt()?; + semantic_std = v.to_device(device)?.to_dtype(dtype_f32)?.sqrt()?; } } let semantic_codec_path = save_dir.to_string() + "/amphion/MaskGCT/semantic_codec/model.safetensors"; - let vb = - unsafe { VarBuilder::from_mmaped_safetensors(&[semantic_codec_path], dtype, &device)? }; + let vb = unsafe { + VarBuilder::from_mmaped_safetensors(&[semantic_codec_path], dtype_f32, device)? + }; let semantic_codec = RepCodec::new(vb, &config.semantic_codec)?; let s2mel_filters = mel_filter_bank( config.s2mel.preprocess_params.spect_params.n_fft / 2 + 1, @@ -937,21 +2071,61 @@ impl IndexTTS2Model { .t()?; let s2mel_windows = create_hann_window( config.s2mel.preprocess_params.spect_params.win_length, - dtype, + dtype_f32, device, )?; let (window_shift, window_size, padded_window_size) = get_waveform_and_window_properties(16000, 10.0, 25.0, true)?; let (mel_energies, _) = - kaldi_get_mel_banks(80, padded_window_size, 16000 as f32, 20.0, 0.0, device)?; + kaldi_get_mel_banks(80, padded_window_size, 16000_f32, 20.0, 0.0, device)?; let mel_energies = mel_energies.pad_with_zeros(D::Minus1, 0, 1)?.t()?; let campplus_model_path = save_dir.to_string() + "/iic/speech_campplus_sv_zh-cn_16k-common/campplus_cn_common.bin"; - let campplus_vb = get_vb_model_path(campplus_model_path, dtype, device.clone(), None)?; + let campplus_vb = get_vb_model_path(campplus_model_path, dtype_f32, device.clone(), None)?; let campplus_model = CAMPPlus::new(campplus_vb, 80, 192, 32, 4, 128)?; - let s2mel = MyModel::new(path, &config.s2mel, dtype, device)?; + let s2mel = MyModel::new(path, &config.s2mel, dtype_f32, device)?; + let emo_matrix_path = path.to_string() + "/" + &config.emo_matrix; + let emo_num = config.emo_num.clone(); + let t_emo = load_tensor_from_pt( + &emo_matrix_path, + "feat2/data/0", + Shape::from_dims(&[73, 1280]), + device, + )?; + let emo_matrix = split_tensor(&t_emo, &emo_num, 0)?; + let skp_matrix_path = path.to_string() + "/" + &config.spk_matrix; + let t_spk = load_tensor_from_pt( + &skp_matrix_path, + "feat1/data/0", + Shape::from_dims(&[73, 192]), + device, + )?; + let spk_matrix = split_tensor(&t_spk, &emo_num, 0)?; + let gpt_path = path.to_string() + "/" + &config.gpt_checkpoint; + let gpt_dict = read_all_with_key(gpt_path, None)?; + let mut dict_to_hashmap = HashMap::new(); + for (k, v) in gpt_dict { + dict_to_hashmap.insert(k, v); + } + let gpt_vb = VarBuilder::from_tensors(dict_to_hashmap, gpt_dtype, device); + let gpt = UnifiedVoice::new(gpt_vb, &config.gpt)?; + let bigvgan_model_path = save_dir.to_string() + + "/nv-community/bigvgan_v2_22khz_80band_256x/bigvgan_generator.pt"; + let bigvgan_vb = get_vb_model_path( + bigvgan_model_path, + dtype_f32, + device.clone(), + Some("generator"), + )?; + let bigvgan_config_path = + save_dir.to_string() + "/nv-community/bigvgan_v2_22khz_80band_256x/config.json"; + let bigvgan_cfg: BigVGANConfig = + serde_json::from_slice(&std::fs::read(bigvgan_config_path)?)?; + let bigvgan = BigVGAN::new(bigvgan_vb, &bigvgan_cfg)?; Ok(Self { - cache: None, + max_audio_length_seconds: 15, + spk_cache: None, + emo_cache: None, feature_extractor, semantic_model, semantic_mean, @@ -966,6 +2140,14 @@ impl IndexTTS2Model { mel_energies, campplus_model, s2mel, + emo_num, + emo_matrix, + spk_matrix, + gpt, + device: device.clone(), + dtype_f32, + // gpt_dtype, + bigvgan, }) } @@ -977,6 +2159,7 @@ impl IndexTTS2Model { let output = self.semantic_model .forward(input_features, attention_mask, Some(17), false)?; + let feature = &output.specify_layer_id_hidden_state.unwrap(); let feature = feature .broadcast_sub(&self.semantic_mean)? @@ -988,7 +2171,7 @@ impl IndexTTS2Model { let pad = (self.s2mel_preprocess_params.spect_params.n_fft - self.s2mel_preprocess_params.spect_params.hop_length) / 2; - let pad_audio_22k = pad_reflect_last_dim(&waveform, (pad, pad))?; + let pad_audio_22k = pad_reflect_last_dim(waveform, (pad, pad))?; let spec = torch_stft( &pad_audio_22k, self.s2mel_preprocess_params.spect_params.n_fft, @@ -1001,59 +2184,260 @@ impl IndexTTS2Model { Ok(spec) } + pub fn cut_audio(&self, audio: &Tensor, sr: usize) -> Result<(Tensor, usize)> { + let max_audio_samples = self.max_audio_length_seconds * sr; + let audio_lens = audio.dim(1)?; + let audio = if audio_lens > max_audio_samples { + audio.i((.., 0..max_audio_samples))? + } else { + audio.clone() + }; + Ok((audio, sr)) + } + + pub fn process_spk_info(&self, audio_url: &str) -> Result<(Tensor, Tensor)> { + let (audio, sr) = load_audio(audio_url, &self.device)?; + let (audio, sr) = self.cut_audio(&audio, sr)?; + let audio_22k = resample_simple(&audio, sr as i64, 22050)?; + let audio_16k = resample_simple(&audio, sr as i64, 16000)?; + + Ok((audio_22k, audio_16k)) + } + + pub fn process_emo_info(&self, audio_url: &str) -> Result { + let (audio, sr) = load_audio(audio_url, &self.device)?; + let (audio, sr) = self.cut_audio(&audio, sr)?; + let audio_16k = resample_simple(&audio, sr as i64, 16000)?; + + Ok(audio_16k) + } + + fn is_use_spk_cache(&self, audio_url_vec: &[String]) -> bool { + let mut flag = false; + if audio_url_vec.is_empty() && self.spk_cache.is_some() { + flag = true; + } + if !audio_url_vec.is_empty() + && let Some(cache) = self.spk_cache.as_ref() + && cache.cache_spk_audio_prompt.eq(&audio_url_vec[0]) + { + flag = true; + } + flag + } + fn is_use_emo_cache(&self, audio_url_vec: &[String]) -> bool { + let mut flag = false; + if audio_url_vec.len() < 2 && self.emo_cache.is_some() { + flag = true; + } + if audio_url_vec.len() >= 2 + && let Some(cache) = self.emo_cache.as_ref() + && cache.cache_emo_audio_prompt.eq(&audio_url_vec[1]) + { + flag = true; + } + if audio_url_vec.len() == 1 + && let Some(cache) = self.emo_cache.as_ref() + && cache.cache_emo_audio_prompt.eq(&audio_url_vec[0]) + { + flag = true; + } + flag + } + + fn find_most_similar_cosine(&self, query_vector: &Tensor) -> Result> { + let mut index = vec![]; + for temp in self.spk_matrix.iter() { + let similarities = cosine_similarity(query_vector, temp)?.squeeze(0)?; + let max_idx = similarities.argmax(0)?.to_scalar::()? as usize; + index.push(max_idx); + } + Ok(index) + } + + fn process_emo_vec( + &self, + emo_vec: Vec, + use_random: bool, + style: &Tensor, + emo_weight: f64, + ) -> Result<(Tensor, Tensor)> { + let weight_vector = Tensor::new(emo_vec, &self.device)?.affine(emo_weight, 0.0)?; + let mut rng = rand::rng(); + let random_index = if use_random { + self.emo_num + .iter() + .map(|&x| rng.random_range(0..x.max(1))) + .collect::>() + } else { + self.find_most_similar_cosine(style)? + }; + let mut emo_matrix = vec![]; + for (i, &index) in random_index.iter().enumerate() { + let tmp = self.emo_matrix[i].as_ref().i(index)?.unsqueeze(0)?; + emo_matrix.push(tmp); + } + let emo_matrix = Tensor::cat(&emo_matrix, 0)?; + let emo_vec_mat = weight_vector.unsqueeze(1)?.broadcast_mul(&emo_matrix)?; + let emo_vec_mat = emo_vec_mat.sum(0)?; + let emo_vec_mat = emo_vec_mat.unsqueeze(0)?; + Ok((emo_vec_mat, weight_vector)) + } + pub fn forward( &mut self, input_ids: &Tensor, - audio_22k: Option<&Tensor>, - audio_16k: Option<&Tensor>, - ) -> Result<()> { - if (audio_22k.is_none() || audio_16k.is_none()) && self.cache.is_none() { + mes: &ChatCompletionParameters, + ) -> Result { + let audio_url_vec = extract_audio_url(mes); + if audio_url_vec.is_empty() && self.spk_cache.is_none() && self.emo_cache.is_none() { return Err(anyhow!( - "Missing required audio input: must provide either audio_22k, audio_16k, or have cached prompt data available" + "Missing audio input: please provide an audio prompt URL or initialize the speaker cache first" )); } - let (spk_cond_emb, style, prompt_condition, ref_mel) = if let Some(audio_22k) = audio_22k - && let Some(audio_16k) = audio_16k - { - let (audio_16k_features, audio_16k_mask) = - self.feature_extractor.call(&audio_16k, 16000, true, true)?; - let spk_cond_emb = self.get_emb(&audio_16k_features, audio_16k_mask.as_ref())?; - let (_, s_ref) = self.semantic_codec.quantize(&spk_cond_emb)?; - let ref_mel = self.s2mel_spectrogram(&audio_22k)?; - let ref_target_lengths = Tensor::new(ref_mel.dim(2)? as u32, ref_mel.device())?; - let feat = kaldi_fbank( - &audio_16k, - &self.mel_energies, - self.window_shift, - self.window_size, - self.padded_window_size, - 0.0, - )?; - let feat = feat.broadcast_sub(&feat.mean_keepdim(1)?)?; - let style = self.campplus_model.forward(&feat)?; - let prompt_condition = self - .s2mel - .length_regulator - .forward(&s_ref, &ref_target_lengths)?; - let cache = IndexTTS2Cache { - cache_spk_cond: spk_cond_emb.clone(), - cache_s2mel_style: style.clone(), - cache_s2mel_prompt: prompt_condition.clone(), - cache_mel: ref_mel.clone(), + let (spk_cond_emb, style, prompt_condition, ref_mel) = + if self.is_use_spk_cache(&audio_url_vec) { + let cache = self.spk_cache.as_ref().unwrap(); + ( + cache.cache_spk_cond.clone(), + cache.cache_s2mel_style.clone(), + cache.cache_s2mel_prompt.clone(), + cache.cache_mel.clone(), + ) + } else { + let (audio_22k, audio_16k) = self.process_spk_info(&audio_url_vec[0])?; + let (audio_16k_features, audio_16k_mask) = + self.feature_extractor.call(&audio_16k, 16000, true, true)?; + let spk_cond_emb = self.get_emb(&audio_16k_features, audio_16k_mask.as_ref())?; + let (_, s_ref) = self.semantic_codec.quantize(&spk_cond_emb)?; + let ref_mel = self.s2mel_spectrogram(&audio_22k)?; + let ref_target_lengths = Tensor::new(ref_mel.dim(2)? as u32, ref_mel.device())?; + let feat = kaldi_fbank( + &audio_16k, + &self.mel_energies, + self.window_shift, + self.window_size, + self.padded_window_size, + 0.0, + )? + .squeeze(0)?; + let feat = feat.broadcast_sub(&feat.mean_keepdim(0)?)?.unsqueeze(0)?; + let style = self.campplus_model.forward(&feat)?; + let prompt_condition = self + .s2mel + .length_regulator + .forward(&s_ref, &ref_target_lengths)?; + let cache = IndexTTS2SpkCache { + cache_spk_cond: spk_cond_emb.clone(), + cache_s2mel_style: style.clone(), + cache_s2mel_prompt: prompt_condition.clone(), + cache_mel: ref_mel.clone(), + cache_spk_audio_prompt: audio_url_vec[0].clone(), + }; + self.spk_cache = Some(cache); + (spk_cond_emb, style, prompt_condition, ref_mel) }; - self.cache = Some(cache); - (spk_cond_emb, style, prompt_condition, ref_mel) + + let emo_vector = if let Some(map) = &mes.metadata + && let Some(emo_vector_str) = map.get("emo_vector") + { + serde_json::from_str::>(emo_vector_str).ok() + // match serde_json::from_str::>(emo_vector_str) { + // Ok(emo_vector) => Some(emo_vector), + // Err(_) => None, + // } } else { - let cache = self.cache.as_ref().unwrap(); - ( - cache.cache_spk_cond.clone(), - cache.cache_s2mel_style.clone(), - cache.cache_s2mel_prompt.clone(), - cache.cache_mel.clone(), - ) + None }; - - - Ok(()) + + let use_random = if let Some(map) = &mes.metadata + && let Some(use_random) = map.get("use_random") + { + use_random.parse::().unwrap_or(false) + } else { + false + }; + let emo_weight = if let Some(map) = &mes.metadata + && let Some(emo_weight) = map.get("emo_weight") + { + emo_weight.parse::().unwrap_or(1.0) + } else { + 1.0 + }; + let (emovec_mat, weight_vector) = if let Some(emo_vector) = emo_vector { + let (emovec_mat, weight_vector) = + self.process_emo_vec(emo_vector, use_random, &style, emo_weight)?; + (Some(emovec_mat), Some(weight_vector)) + } else { + (None, None) + }; + + let emo_cond_emb = if self.is_use_emo_cache(&audio_url_vec) { + let cache = self.emo_cache.as_ref().unwrap(); + cache.emo_cond_emb.clone() + } else { + let emo_audio_prompt = if audio_url_vec.len() >= 2 { + audio_url_vec[1].clone() + } else { + audio_url_vec[0].clone() + }; + let emo_audio_16k = self.process_emo_info(&emo_audio_prompt)?; + let (emo_audio_16k_features, emo_audio_16k_mask) = + self.feature_extractor + .call(&emo_audio_16k, 16000, true, true)?; + let emo_cond_emb = + self.get_emb(&emo_audio_16k_features, emo_audio_16k_mask.as_ref())?; + let cache = IndexTTS2EmoCache { + cache_emo_audio_prompt: emo_audio_prompt, + emo_cond_emb: emo_cond_emb.clone(), + }; + self.emo_cache = Some(cache); + emo_cond_emb + }; + let mut emovec = self + .gpt + .merge_emovec(&spk_cond_emb, &emo_cond_emb, emo_weight)?; + if let Some(weight_vector) = weight_vector + && let Some(emovec_mat) = emovec_mat + { + let ratio = 1.0 - weight_vector.sum_all()?.to_scalar::()?; + emovec = emovec + .affine(ratio as f64, 0.0)? + .add(&emovec_mat.to_dtype(emovec.dtype())?)?; + } + let (codes, speech_conditioning_latent) = + self.gpt + .inference_speech(&spk_cond_emb, input_ids, &emovec)?; + let code_len = codes.dim(0)?; + let codes = codes.unsqueeze(0)?; + let use_speed = Tensor::zeros(spk_cond_emb.dim(0)?, DType::U32, &self.device)?; + let latent = self.gpt.forward( + &speech_conditioning_latent, + input_ids, + &codes, + &emovec, + &use_speed, + )?; + let latent = self + .s2mel + .gpt_layer + .forward(&latent.to_dtype(self.dtype_f32)?)?; + let s_infer = self.semantic_codec.quantizer.vq2emb(&codes)?; + let s_infer = s_infer.transpose(1, 2)?; + let s_infer = s_infer.add(&latent)?; + let target_len = (code_len as f32 * 1.72) as u32; + let target_len = Tensor::new(target_len, &self.device)?; + let code = self.s2mel.length_regulator.forward(&s_infer, &target_len)?; + let cat_condition = Tensor::cat(&[&prompt_condition, &code], 1)?; + let x_lens = Tensor::new(vec![cat_condition.dim(1)? as u32], &self.device)?; + let vc_target = + self.s2mel + .cfm + .inference(&cat_condition, &x_lens, &ref_mel, &style, 25, 0.7)?; + let ref_mel_len = ref_mel.dim(D::Minus1)?; + let vc_target_len = vc_target.dim(D::Minus1)?; + let vc_target = vc_target.narrow(D::Minus1, ref_mel_len, vc_target_len - ref_mel_len)?; + let wav = self.bigvgan.forward(&vc_target)?.squeeze(0)?; + Ok(wav) } } diff --git a/src/models/index_tts2/processor.rs b/src/models/index_tts2/processor.rs deleted file mode 100644 index 23839fc..0000000 --- a/src/models/index_tts2/processor.rs +++ /dev/null @@ -1,203 +0,0 @@ -use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; -use anyhow::Result; -use candle_core::{D, DType, Device, IndexOp, Tensor, pickle::read_all_with_key}; -use candle_nn::VarBuilder; - -use crate::{ - models::{ - campplus::CAMPPlus, - feature_extractor::seamless_m4t_feature_extractor::SeamlessM4TFeatureExtractor, - index_tts2::config::{IndexTTS2Config, PreprocessParams}, - mask_gct::model::RepCodec, - w2v_bert_2_0::model::W2VBert2_0Model, - }, - utils::{ - audio_utils::{ - create_hann_window, extract_audio_url, get_waveform_and_window_properties, kaldi_fbank, - kaldi_get_mel_banks, load_audio, mel_filter_bank, resample_simple, torch_stft, - }, - get_vb_model_path, - tensor_utils::pad_reflect_last_dim, - }, -}; - -pub struct IndexTTS2Processor { - device: Device, - max_audio_length_seconds: usize, - // feature_extractor: SeamlessM4TFeatureExtractor, - // semantic_model: W2VBert2_0Model, - // semantic_mean: Tensor, - // semantic_std: Tensor, - // semantic_codec: RepCodec, - // s2mel_filters: Tensor, - // s2mel_windows: Tensor, - // s2mel_preprocess_params: PreprocessParams, - // window_shift: usize, - // window_size: usize, - // padded_window_size: usize, - // mel_energies: Tensor, - // campplus_model: CAMPPlus, -} - -impl IndexTTS2Processor { - pub fn new( - // path: &str, - // save_dir: &str, - // config: &IndexTTS2Config, - device: &Device, - // dtype: DType, - ) -> Result { - // let feature_extractor = SeamlessM4TFeatureExtractor::new( - // // 80, - // 80, - // crate::utils::tensor_utils::PaddingSide::Right, - // 1.0, - // 16000, - // 2, - // device, - // )?; - // let w2vbert2_path = save_dir.to_string() + "/facebook/w2v-bert-2.0"; - // let semantic_model = W2VBert2_0Model::init(&w2vbert2_path, device, dtype)?; - // let semantic_mean_var_path = path.to_string() + "/" + &config.w2v_stat; - // let dict = read_all_with_key(semantic_mean_var_path, None)?; - // let mut semantic_mean = Tensor::new(0.0, device)?.to_dtype(dtype)?; - // let mut semantic_std = Tensor::new(1.0, device)?.to_dtype(dtype)?; - // for (k, v) in dict { - // if k.eq("mean") { - // semantic_mean = v.to_device(device)?.to_dtype(dtype)?; - // } else if k.eq("var") { - // semantic_std = v.to_device(device)?.to_dtype(dtype)?.sqrt()?; - // } - // } - - // let semantic_codec_path = - // save_dir.to_string() + "/amphion/MaskGCT/semantic_codec/model.safetensors"; - // let vb = - // unsafe { VarBuilder::from_mmaped_safetensors(&[semantic_codec_path], dtype, &device)? }; - // let semantic_codec = RepCodec::new(vb, &config.semantic_codec)?; - // let s2mel_filters = mel_filter_bank( - // config.s2mel.preprocess_params.spect_params.n_fft / 2 + 1, - // config.s2mel.preprocess_params.spect_params.n_mels, - // config.s2mel.preprocess_params.spect_params.fmin as f32, - // config - // .s2mel - // .preprocess_params - // .spect_params - // .fmax - // .unwrap_or(config.s2mel.preprocess_params.sr / 2) as f32, - // config.s2mel.preprocess_params.sr as f32, - // Some("slaney"), - // crate::utils::audio_utils::MelScale::Slaney, - // false, - // device, - // )? - // .t()?; - // let s2mel_windows = create_hann_window( - // config.s2mel.preprocess_params.spect_params.win_length, - // dtype, - // device, - // )?; - // let (window_shift, window_size, padded_window_size) = - // get_waveform_and_window_properties(16000, 10.0, 25.0, true)?; - // let (mel_energies, _) = - // kaldi_get_mel_banks(80, padded_window_size, 16000 as f32, 20.0, 0.0, device)?; - // let mel_energies = mel_energies.pad_with_zeros(D::Minus1, 0, 1)?.t()?; - // let campplus_model_path = save_dir.to_string() - // + "/iic/speech_campplus_sv_zh-cn_16k-common/campplus_cn_common.bin"; - // let campplus_vb = get_vb_model_path(campplus_model_path, dtype, device.clone(), None)?; - // let campplus_model = CAMPPlus::new(campplus_vb, 80, 192, 32, 4, 128)?; - Ok(Self { - device: device.clone(), - max_audio_length_seconds: 15, - // feature_extractor, - // semantic_model, - // semantic_mean, - // semantic_std, - // semantic_codec, - // s2mel_filters, - // s2mel_windows, - // s2mel_preprocess_params: config.s2mel.preprocess_params.clone(), - // window_shift, - // window_size, - // padded_window_size, - // mel_energies, - // campplus_model, - }) - } - - pub fn cut_audio(&self, audio: &Tensor, sr: usize) -> Result<(Tensor, usize)> { - let max_audio_samples = self.max_audio_length_seconds * sr; - let audio_lens = audio.dim(1)?; - let audio = if audio_lens > max_audio_samples { - audio.i((.., 0..max_audio_samples))? - } else { - audio.clone() - }; - Ok((audio, sr)) - } - - pub fn extract_audio_and_cut( - &self, - mes: &ChatCompletionParameters, - device: &Device, - ) -> Result<(Tensor, usize, String)> { - let audio_url_vec = extract_audio_url(mes); - let (audio, sr) = load_audio(&audio_url_vec[0], device)?; - let (audio, sr) = self.cut_audio(&audio, sr)?; - Ok((audio, sr, audio_url_vec[0].clone())) - } - - // pub fn get_emb( - // &self, - // input_features: &Tensor, - // attention_mask: Option<&Tensor>, - // ) -> Result { - // let output = - // self.semantic_model - // .forward(input_features, attention_mask, Some(17), false)?; - // let feature = &output.specify_layer_id_hidden_state.unwrap(); - // let feature = feature - // .broadcast_sub(&self.semantic_mean)? - // .broadcast_div(&self.semantic_std)?; - // Ok(feature) - // } - - // pub fn s2mel_spectrogram(&self, waveform: &Tensor) -> Result { - // let pad = (self.s2mel_preprocess_params.spect_params.n_fft - // - self.s2mel_preprocess_params.spect_params.hop_length) - // / 2; - // let pad_audio_22k = pad_reflect_last_dim(&waveform, (pad, pad))?; - // let spec = torch_stft( - // &pad_audio_22k, - // self.s2mel_preprocess_params.spect_params.n_fft, - // self.s2mel_preprocess_params.spect_params.hop_length, - // &self.s2mel_windows, - // )? - // .transpose(1, 2)?; - // let spec = self.s2mel_filters.broadcast_matmul(&spec)?; - // let spec = spec.clamp(1e-5, f64::INFINITY)?.log()?; - // Ok(spec) - // } - pub fn process_info(&self, mes: &ChatCompletionParameters) -> Result<(Tensor, Tensor, String)> { - let (audio, sr, audio_url) = self.extract_audio_and_cut(mes, &self.device)?; - let audio_22k = resample_simple(&audio, sr as i64, 22050)?; - let audio_16k = resample_simple(&audio, sr as i64, 16000)?; - // let (audio_16k_features, audio_16k_mask) = - // self.feature_extractor.call(&audio_16k, 16000, true, true)?; - // let spk_cond_emb = self.get_emb(&audio_16k_features, audio_16k_mask.as_ref())?; - // let (_, s_ref) = self.semantic_codec.quantize(&spk_cond_emb)?; - // let ref_mel = self.s2mel_spectrogram(&audio_22k)?; - // let feat = kaldi_fbank( - // &audio_16k, - // &self.mel_energies, - // self.window_shift, - // self.window_size, - // self.padded_window_size, - // 0.0, - // )?; - // let feat = feat.broadcast_sub(&feat.mean_keepdim(1)?)?; - // let style = self.campplus_model.forward(&feat)?; - // println!("style: {}", style); - Ok((audio_22k, audio_16k, audio_url)) - } -} diff --git a/src/models/index_tts2/utils.rs b/src/models/index_tts2/utils.rs index d80ecfb..5d00928 100644 --- a/src/models/index_tts2/utils.rs +++ b/src/models/index_tts2/utils.rs @@ -1,5 +1,3 @@ -use std::collections::HashMap; - use anyhow::{Result, anyhow}; use regex::Regex; @@ -15,9 +13,12 @@ pub async fn download_index_tts2_need_model(save_dir: Option<&str>) -> Result<() let mask_gct = "amphion/MaskGCT"; // let campplus= "funasr/campplus"; // huggingface let campplus = "iic/speech_campplus_sv_zh-cn_16k-common"; // modelscope + // let bigvgan = "nvidia/bigvgan_v2_22khz_80band_256x"; // huggingface + let bigvgan = "nv-community/bigvgan_v2_22khz_80band_256x"; // modelscope download_model(w2v_bert2_0, &save_dir, 3).await?; download_model(mask_gct, &save_dir, 3).await?; download_model(campplus, &save_dir, 3).await?; + download_model(bigvgan, &save_dir, 3).await?; Ok(()) } diff --git a/src/models/mask_gct/config.rs b/src/models/mask_gct/config.rs index a38930b..455562d 100644 --- a/src/models/mask_gct/config.rs +++ b/src/models/mask_gct/config.rs @@ -18,4 +18,4 @@ fn default_num_quantizers() -> usize { fn default_downsample_scale() -> usize { 1 -} \ No newline at end of file +} diff --git a/src/models/mask_gct/mod.rs b/src/models/mask_gct/mod.rs index 06008be..2852cdb 100644 --- a/src/models/mask_gct/mod.rs +++ b/src/models/mask_gct/mod.rs @@ -1,2 +1,2 @@ +pub mod config; pub mod model; -pub mod config; \ No newline at end of file diff --git a/src/models/mask_gct/model.rs b/src/models/mask_gct/model.rs index a212d52..4745968 100644 --- a/src/models/mask_gct/model.rs +++ b/src/models/mask_gct/model.rs @@ -48,7 +48,7 @@ impl ConvNeXtBlock { let xs = self.pwconv1.forward(&xs)?.gelu()?; let mut xs = self.pwconv2.forward(&xs)?; if let Some(gamma) = &self.gamma { - xs = xs.broadcast_mul(&gamma)?; + xs = xs.broadcast_mul(gamma)?; } let xs = xs.transpose(1, 2)?; let xs = residual.add(&xs)?; @@ -116,10 +116,28 @@ impl FactorizedVectorQuantize { use_l2_normlize: bool, ) -> Result { let (in_project, out_project) = if input_dim != codebook_dim { - let in_project = - WNConv1d::new(vb.pp("in_project"), input_dim, codebook_dim, 1, 1, 0, 1, 1, true)?; - let out_project = - WNConv1d::new(vb.pp("out_project"), codebook_dim, input_dim, 1, 1, 0, 1, 1, true)?; + let in_project = WNConv1d::new( + vb.pp("in_project"), + input_dim, + codebook_dim, + 1, + 1, + 0, + 1, + 1, + true, + )?; + let out_project = WNConv1d::new( + vb.pp("out_project"), + codebook_dim, + input_dim, + 1, + 1, + 0, + 1, + 1, + true, + )?; (Some(in_project), Some(out_project)) } else { (None, None) @@ -166,6 +184,14 @@ impl FactorizedVectorQuantize { } Ok((z_q, indices)) } + + pub fn vq2emb(&self, xs: &Tensor) -> Result { + let mut emb = self.codebook.forward(xs)?.transpose(1, 2)?; + if let Some(out_proj) = &self.out_project { + emb = out_proj.forward(&emb)?; + } + Ok(emb) + } } pub struct ResidualVQ { @@ -210,7 +236,7 @@ impl ResidualVQ { let n_quantizers = n_quantizers.unwrap_or(self.num_quantizers); let mut residual = xs.clone(); let mut quantized_out = Tensor::new(0.0f32, xs.device())?.to_dtype(xs.dtype())?; - for (i, quantizer) in (&self.quantizers).iter().enumerate() { + for (i, quantizer) in self.quantizers.iter().enumerate() { if i >= n_quantizers { break; } @@ -224,8 +250,17 @@ impl ResidualVQ { let all_quantized = Tensor::stack(&all_quantized, 0)?; Ok((quantized_out, all_indices, all_quantized)) } + + pub fn vq2emb(&self, xs: &Tensor) -> Result { + let mut quantized_out = xs.clone(); + for quantizer in &self.quantizers { + quantized_out = quantizer.vq2emb(xs)?; + } + Ok(quantized_out) + } } +#[allow(unused)] pub struct RepCodec { downsample_scale: usize, down: Option, @@ -234,7 +269,7 @@ pub struct RepCodec { encoder_1: Linear, decoder_0: VocosBackbone, decoder_1: Linear, - quantizer: ResidualVQ, + pub quantizer: ResidualVQ, } impl RepCodec { diff --git a/src/models/mod.rs b/src/models/mod.rs index 3ad06af..b0933de 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -1,3 +1,4 @@ +pub mod bigvgan; pub mod campplus; pub mod common; pub mod deepseek_ocr; diff --git a/src/models/qwen3_asr/config.rs b/src/models/qwen3_asr/config.rs index 7373fde..853f186 100644 --- a/src/models/qwen3_asr/config.rs +++ b/src/models/qwen3_asr/config.rs @@ -167,13 +167,13 @@ pub struct Qwen3ASRTextConfig { pub fn qwen3asr_text_config2qwen3_config(cfg: &Qwen3ASRTextConfig) -> Qwen3Config { Qwen3Config { attention_bias: cfg.attention_bias, - attention_dropout: cfg.attention_dropout as f64, - bos_token_id: cfg.bos_token_id.unwrap_or(151643) as u32, - eos_token_id: cfg.eos_token_id.unwrap_or(151645) as u32, + attention_dropout: cfg.attention_dropout, + bos_token_id: cfg.bos_token_id.unwrap_or(151643), + eos_token_id: cfg.eos_token_id.unwrap_or(151645), head_dim: cfg.head_dim, hidden_act: cfg.hidden_act, hidden_size: cfg.hidden_size, - initializer_range: cfg.initializer_range as f64, + initializer_range: cfg.initializer_range, intermediate_size: cfg.intermediate_size, max_position_embeddings: cfg.max_position_embeddings, max_window_layers: 0, diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index 22b70cd..087e9d1 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -1,4 +1,3 @@ - use aha_openai_dive::v1::resources::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; diff --git a/src/models/qwen3_asr/model.rs b/src/models/qwen3_asr/model.rs index 188c863..af4fc87 100644 --- a/src/models/qwen3_asr/model.rs +++ b/src/models/qwen3_asr/model.rs @@ -338,7 +338,7 @@ impl Qwen3ASRThinker { ) -> Result { let mut input_embeds = self.model.embed_tokens.forward(input_ids)?; if let Some(input_features) = input_features { - let audio_feature = self.audio_tower.forward(&input_features)?; + let audio_feature = self.audio_tower.forward(input_features)?; // println!("audio_feature: {}", audio_feature); let audio_mask = get_equal_mask(input_ids, self.audio_token_id)?; let n_audio_tokens = audio_mask.sum_all()?.to_scalar::()?; diff --git a/src/models/qwen3_asr/processor.rs b/src/models/qwen3_asr/processor.rs index e50e9b9..cf894f8 100644 --- a/src/models/qwen3_asr/processor.rs +++ b/src/models/qwen3_asr/processor.rs @@ -82,10 +82,7 @@ impl Qwen3AsrProcessor { pub fn process_audio(&self, mes: &ChatCompletionParameters) -> Result> { let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?; - audio_tensors - .iter() - .map(|audio| float_range_normalize(&audio)) - .collect() + audio_tensors.iter().map(float_range_normalize).collect() } pub fn validate_language(&self, lang: &String) -> bool { @@ -93,10 +90,9 @@ impl Qwen3AsrProcessor { } fn replace_special_tokens(&self, text: &str, token_len: usize) -> String { - let replace = "<|audio_placeholder|>".repeat(token_len as usize); + let replace = "<|audio_placeholder|>".repeat(token_len); let text = text.replacen(&self.audio_token, &replace, 1); - let text = text.replace("<|audio_placeholder|>", &self.audio_token); - text + text.replace("<|audio_placeholder|>", &self.audio_token) } pub fn process_info( @@ -162,11 +158,10 @@ pub struct AudioData { pub fn get_feat_extract_output_lengths(audio_len: usize) -> usize { let input_len_leave = audio_len % 100; - let output_len = if input_len_leave > 0 { + if input_len_leave > 0 { let feat_lengths = (input_len_leave - 1) / 2 + 1; ((feat_lengths - 1) / 2 + 1 - 1) / 2 + 1 + (audio_len / 100) * 13 } else { (audio_len / 100) * 13 - }; - output_len + } } diff --git a/src/models/w2v_bert_2_0/config.rs b/src/models/w2v_bert_2_0/config.rs index 51142ac..4b34833 100644 --- a/src/models/w2v_bert_2_0/config.rs +++ b/src/models/w2v_bert_2_0/config.rs @@ -58,4 +58,4 @@ pub struct W2VBert2_0Config { pub use_weighted_layer_sum: bool, pub vocab_size: Option, pub xvector_output_dim: usize, -} \ No newline at end of file +} diff --git a/src/models/w2v_bert_2_0/mod.rs b/src/models/w2v_bert_2_0/mod.rs index 3621b7a..2852cdb 100644 --- a/src/models/w2v_bert_2_0/mod.rs +++ b/src/models/w2v_bert_2_0/mod.rs @@ -1,2 +1,2 @@ pub mod config; -pub mod model; \ No newline at end of file +pub mod model; diff --git a/src/models/w2v_bert_2_0/model.rs b/src/models/w2v_bert_2_0/model.rs index e9fa3f9..fb29025 100644 --- a/src/models/w2v_bert_2_0/model.rs +++ b/src/models/w2v_bert_2_0/model.rs @@ -45,6 +45,7 @@ impl Wav2Vec2BertFeatureProjection { } } +#[allow(unused)] pub struct Wav2Vec2BertSelfAttention { q_proj: Linear, k_proj: Linear, @@ -203,17 +204,12 @@ impl Wav2Vec2BertSelfAttention { .affine(scale, 0.0)?; if let Some(mask) = attention_mask { // let mask = mask.unsqueeze(1)?.unsqueeze(D::Minus1)?; - Some(relative_position_attn_weights.broadcast_add(&mask)?) + Some(relative_position_attn_weights.broadcast_add(mask)?) } else { Some(relative_position_attn_weights) } } else { - if let Some(mask) = attention_mask { - // let mask = mask.unsqueeze(1)?.unsqueeze(D::Minus1)?; - Some(mask.clone()) - } else { - None - } + attention_mask.cloned() }; let attn_output = eager_attention_forward( @@ -476,7 +472,8 @@ impl Wav2Vec2BertEncoder { let attention_mask = attention_mask .where_cond(&attention_mask_f, &neg_inf_t)? .to_dtype(xs.dtype())? - .affine(1.0, -1.0)?; + .affine(1.0, -1.0)? + .repeat((1, 1, seq_len, 1))?; (xs, Some(attention_mask)) } else { (xs.clone(), None) @@ -490,7 +487,7 @@ impl Wav2Vec2BertEncoder { let mut hidden_states: Vec = vec![]; let mut specify_layer_id_hidden_state = None; - for (i, layer) in (&self.layers).iter().enumerate() { + for (i, layer) in self.layers.iter().enumerate() { if output_hidden_states { hidden_states.push(xs.clone()); } @@ -507,7 +504,7 @@ impl Wav2Vec2BertEncoder { conv_attention_mask, )?; } - let hidden_states = if hidden_states.len() > 0 { + let hidden_states = if !hidden_states.is_empty() { Some(hidden_states) } else { None @@ -521,9 +518,9 @@ impl Wav2Vec2BertEncoder { } pub struct W2VBert2_0Model { - config: W2VBert2_0Config, + // config: W2VBert2_0Config, feature_projection: Wav2Vec2BertFeatureProjection, - masked_spec_embed: Option, + // masked_spec_embed: Option, encoder: Wav2Vec2BertEncoder, // config.add_adapter is false, adapter is None, Wav2Vec2BertAdapter not complish // adapter: Option, @@ -543,21 +540,21 @@ impl W2VBert2_0Model { pub fn new(vb: VarBuilder, config: &W2VBert2_0Config) -> Result { let feature_projection = Wav2Vec2BertFeatureProjection::new(vb.pp("feature_projection"), config)?; - let masked_spec_embed = if config.mask_time_prob > 0.0 || config.mask_time_prob > 0.0 { - Some( - vb.get_with_hints(config.hidden_size, "masked_spec_embed", Init::Uniform { - lo: 0.0, - up: 1.0, - })?, - ) - } else { - None - }; + // let masked_spec_embed = if config.mask_time_prob > 0.0 || config.mask_time_prob > 0.0 { + // Some( + // vb.get_with_hints(config.hidden_size, "masked_spec_embed", Init::Uniform { + // lo: 0.0, + // up: 1.0, + // })?, + // ) + // } else { + // None + // }; let encoder = Wav2Vec2BertEncoder::new(vb.pp("encoder"), config)?; Ok(Self { - config: config.clone(), + // config: config.clone(), feature_projection, - masked_spec_embed, + // masked_spec_embed, encoder, }) } diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index 82ed91d..319dfdb 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -350,7 +350,7 @@ impl Qwen3VLTextRotaryEmbedding { // for dim in 1..3 { for (dim, offset) in (1..3).enumerate() { - let dim = dim +1; + let dim = dim + 1; let length = mrope_section[dim]; let idx = Tensor::arange_step(offset as u32, length as u32, 3, freqs.device())?; let src = freqs.i(dim)?.contiguous()?; // (bs, seq_len, head_dim //2) diff --git a/src/tokenizer/mod.rs b/src/tokenizer/mod.rs index bcd0811..1f1a5e8 100644 --- a/src/tokenizer/mod.rs +++ b/src/tokenizer/mod.rs @@ -58,19 +58,17 @@ impl TokenizerModel { if let Value::Object(tokens_map) = added_tokens_decoder { for (_, token_info) in tokens_map { - if let Value::Object(token_obj) = token_info { - if let Some(content_val) = token_obj.get("content") { - if let Some(content) = content_val.as_str() { - let special = token_obj - .get("special") - .and_then(|v| v.as_bool()) - .unwrap_or(false); + if let Value::Object(token_obj) = token_info + && let Some(content_val) = token_obj.get("content") + && let Some(content) = content_val.as_str() + { + let special = token_obj + .get("special") + .and_then(|v| v.as_bool()) + .unwrap_or(false); - let added_token = - AddedToken::from(content.to_string(), special); - special_tokens.push(added_token); - } - } + let added_token = AddedToken::from(content.to_string(), special); + special_tokens.push(added_token); } } } diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index 1ec8ace..8851319 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -32,7 +32,9 @@ use symphonia::core::meta::MetadataOptions; use symphonia::core::probe::Hint; use crate::utils::get_default_save_dir; -use crate::utils::tensor_utils::{linspace, log10, pad_reflect_last_dim, pad_replicate_last_dim, split_tensor}; +use crate::utils::tensor_utils::{ + linspace, log10, pad_reflect_last_dim, pad_replicate_last_dim, split_tensor, +}; // 重采样方法枚举 #[derive(Debug, Clone, Copy)] @@ -42,12 +44,12 @@ pub enum ResamplingMethod { } // 零阶修正贝塞尔函数 I0 -fn i0(x: f32) -> f32 { +pub 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 { + for k in 1..100 { term = term * half_x_sq / (k * k) as f32; result += term; @@ -1013,6 +1015,40 @@ pub fn crate_hamming_window( Ok(Tensor::from_vec(window, window_size, device)?.to_dtype(dtype)?) } +pub fn crate_kaiser_window( + window_size: usize, + periodic: bool, + beta: f32, + dtype: DType, + device: &Device, +) -> Result { + if window_size < 1 { + return Err(anyhow::anyhow!("window_size must bigger than 0")); + } + if window_size == 1 { + return Ok(Tensor::new(1.0f32, device)?.to_dtype(dtype)?); + } + + let n = if periodic { + window_size as f32 + } else { + (window_size - 1) as f32 + }; + + let n_half = n / 2.0; + let denominator = i0(beta); + + let window = (0..window_size) + .map(|i| { + let x = (i as f32 - n_half) / n_half; + let sqrt_term = (1.0 - x * x).max(0.0).sqrt(); + let numerator = i0(beta * sqrt_term); + numerator / denominator + }) + .collect(); + Ok(Tensor::from_vec(window, window_size, device)?.to_dtype(dtype)?) +} + /// 梅尔频率刻度类型 #[derive(Debug, Clone, Copy)] pub enum MelScale { @@ -1205,7 +1241,7 @@ pub fn torch_stft( ) -> Result { // waveform: already padding // (bs, n_frames, n_fft) - let frames = extract_frames(&waveform, n_fft, hop_length)?; + let frames = extract_frames(waveform, n_fft, hop_length)?; // 应用汉明窗口 let result = frames.broadcast_mul(window)?; // 傅立叶变换 @@ -1575,7 +1611,7 @@ pub fn spectrogram( )?; frames = Tensor::cat(&[buffer_0, buffer_], D::Minus1)?; } - let mut frames = frames.broadcast_mul(&window)?; + let mut frames = frames.broadcast_mul(window)?; let pad_len = fft_length - frame_length; if pad_len > 0 { // (bs, nframes, frame_length) -> (bs, nframes, fft_length) @@ -1625,3 +1661,73 @@ pub fn split_audio_into_chunks(wav: &Tensor, sr: usize, max_chunk_sec: f32) -> R } Ok(wavs) } + +pub fn sinc(x: &Tensor) -> Result { + let pi_x = x.affine(PI, 0.0)?; + + let epsilon = 1e-8; + let mask = x.abs()?.lt(&Tensor::new(epsilon, x.device())?)?; + + let raw_sinc = pi_x.sin()?.div(&pi_x)?; + + // 在接近 0 的位置填充 1.0 + let ones = Tensor::ones_like(x)?; + let res = mask.where_cond(&ones, &raw_sinc)?; + Ok(res) +} + +pub fn kaiser_sinc_filter1d( + cutoff: f32, + half_width: f32, + kernel_size: usize, + device: &Device, + dtype: DType, +) -> Result { + let even = kernel_size.is_multiple_of(2); + let half_size = (kernel_size / 2) as i32; + + // 计算 Kaiser 窗参数 beta + let delta_f = 4.0 * half_width; + let a = 2.285 * (half_size as f32 - 1.0) * std::f32::consts::PI * delta_f + 7.95; + + let beta = if a > 50.0 { + 0.1102 * (a - 8.7) + } else if a >= 21.0 { + 0.5842 * (a - 21.0).powf(0.4) + 0.07886 * (a - 21.0) + } else { + 0.0 + }; + + // 生成 Kaiser 窗 + let window = crate_kaiser_window(kernel_size, false, beta, dtype, device)?; + + // 生成时间序列 + let time: Vec = if even { + ((-half_size)..half_size).map(|i| i as f32 + 0.5).collect() + } else { + (0..kernel_size) + .map(|i| i as f32 - half_size as f32) + .collect() + }; + let time = Tensor::new(time, device)?; + + // 生成滤波器 + let filter_ = if cutoff == 0.0 { + Tensor::zeros((kernel_size,), DType::F32, device)? + } else { + // 2 * cutoff * window * sinc(2 * cutoff * time) + let two_cutoff = (2.0 * cutoff) as f64; + let sinc_input = time.affine(two_cutoff, 0.0)?; + let sinc_vals = sinc(&sinc_input)?; + + let mut filter_val = window.mul(&sinc_vals)?; + filter_val = filter_val.affine(two_cutoff, 0.0)?; + + // 归一化使和为 1 + let sum_val = filter_val.sum_all()?; + filter_val.div(&sum_val)? + }; + + // reshape 为 [1, 1, kernel_size] + Ok(filter_.reshape((1, 1, kernel_size))?) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index e14b690..9a2a4d8 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -3,7 +3,8 @@ pub mod img_utils; pub mod tensor_utils; pub mod video_utils; -use std::io::Read; +use std::fs::File; +use std::io::{Cursor, Read}; use std::{collections::HashMap, fs, path::PathBuf, process::Command, time::Duration}; use aha_openai_dive::v1::resources::{ @@ -27,6 +28,7 @@ use dirs::home_dir; use half::{bf16, f16, slice::HalfFloatSliceExt}; use modelscope::ModelScope; use tokio::time::sleep; +use zip::ZipArchive; pub fn get_device(device: Option<&Device>) -> Device { match device { @@ -274,16 +276,14 @@ pub fn read_pth_tensor_info_cycle>( } } current_obj + } else if let Object::Dict(key_values) = obj { + key_values + .into_iter() + .find(|(k, _)| *k == Object::Unicode(key.to_owned())) + .map(|(_, v)| v) + .ok_or_else(|| anyhow!(format!("key {key} not found")))? } else { - if let Object::Dict(key_values) = obj { - key_values - .into_iter() - .find(|(k, _)| *k == Object::Unicode(key.to_owned())) - .map(|(_, v)| v) - .ok_or_else(|| anyhow!(format!("key {key} not found")))? - } else { - obj - } + obj } } else { obj @@ -309,7 +309,7 @@ pub fn read_pth_tensor_info_cycle>( let tensor_names = tensor_infos.keys(); let mut tensors = Vec::with_capacity(tensor_names.len()); for name in tensor_names { - let _ = match tensor_infos.get(name) { + match tensor_infos.get(name) { None => {} Some(tensor_info) => { let zip_reader = std::io::BufReader::new(std::fs::File::open(&path)?); @@ -806,3 +806,33 @@ pub fn capitalize_first_letter(input: &str) -> String { let remaining = chars.as_str().to_lowercase(); format!("{}{}", first_char, remaining) } + +pub fn load_tensor_from_pt( + path: &str, + zip_name: &str, + shape: Shape, + device: &Device, +) -> Result { + let file = File::open(path)?; + let mut archive = ZipArchive::new(file)?; + // // 列出所有文件(调试用) + // for i in 0..archive.len() { + // let file = archive.by_index(i)?; + // println!("File: {} ({} bytes)", file.name(), file.size()); + // } + // 读取原始字节数据 + let mut data_file = archive.by_name(zip_name)?; + let mut buffer = Vec::new(); + data_file.read_to_end(&mut buffer)?; + // 将字节转换为 f32 (little endian) + let mut cursor = Cursor::new(buffer); + let num_elements = shape.elem_count(); + let mut data = Vec::with_capacity(num_elements); + + for _ in 0..num_elements { + let val = cursor.read_f32::()?; + data.push(val); + } + let t = Tensor::from_vec(data, shape, device)?; + Ok(t) +} diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index 75e35cd..8cbc96b 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -14,7 +14,7 @@ pub fn masked_fill_zeros(hidden_states: &Tensor, mask: &Tensor) -> Result, ) -> Result> { + if input_size == 0 { + return Err(anyhow!("input_size must be > 0")); + } + if output_size == 0 { + return Err(anyhow!("output_size must be > 0")); + } if input_size == 1 { - Ok(vec![0f32; output_size]) - } else if let Some(align_) = align_corner - && align_ - { - Ok((0..output_size) - .map(|i| i as f32 * (input_size - 1) as f32 / (output_size - 1) as f32) - .collect()) + return Ok(vec![0f32; output_size]); + } + let align_corners = align_corner.unwrap_or(false); + if align_corners { + let scale = (input_size - 1) as f32 / (output_size - 1) as f32; + Ok((0..output_size).map(|i| i as f32 * scale).collect()) } else { + let scale = input_size as f32 / output_size as f32; Ok((0..output_size) .map(|i| { - (i as f32 + 0.5) * (input_size as f32 / output_size as f32) - 0.5 - // coord.max(0.0).min((input_size - 1) as f32) + let coord = (i as f32 + 0.5) * scale - 0.5; + coord.clamp(0.0, (input_size - 1) as f32) }) .collect()) } @@ -500,34 +506,27 @@ pub fn interpolate_nearest_1d(t: &Tensor, target_size: usize) -> Result "Input rank must have equal to 3 dimensions" )); } - let shape = t.dims(); - let orig_size = shape[shape.len() - 1]; + + let (bs, channels, orig_size) = t.dims3()?; if orig_size == target_size { return Ok(t.clone()); } - - let (bs, channels, _) = t.dims3()?; - let mut output = Tensor::zeros((bs, channels, target_size), t.dtype(), t.device())?; let coords = compute_1d_coords(orig_size, target_size, None)?; - + let input_data = t.to_vec3::()?; + let mut output_data = vec![vec![vec![0.0f32; target_size]; channels]; bs]; for b in 0..bs { for c in 0..channels { - let input_slice = t.i((b, c))?; - let mut out_i = Vec::new(); - - for &coord in coords.iter().take(target_size) { + for (i, &coord) in coords.iter().enumerate() { // Nearest neighbor: round to nearest integer coordinate - let nearest_idx = coord.floor() as usize; + let nearest_idx = coord.round() as usize; let clamped_idx = nearest_idx.min(orig_size - 1); - let value = input_slice.get(clamped_idx)?; - out_i.push(value); + let value = input_data[b][c][clamped_idx]; + output_data[b][c][i] = value; } - let out_i = Tensor::stack(&out_i, 0)?.unsqueeze(0)?.unsqueeze(0)?; - output = output.slice_assign(&[(b..b + 1), (c..c + 1), (0..target_size)], &out_i)?; } } - output = output.contiguous()?; + let output = Tensor::new(output_data, t.device())?.to_dtype(t.dtype())?; Ok(output) } @@ -1068,3 +1067,14 @@ pub fn sequence_mask(length: &Tensor, max_length: Option) -> Result let mask = x.broadcast_lt(&length)?; Ok(mask) } + +pub fn cosine_similarity(query_vector: &Tensor, matrix: &Tensor) -> Result { + // query_vector: (n, dim) + // matrix: (m, dim) + let query_norm = l2_normalize(query_vector, query_vector.rank() - 1)?; + let matrix_norm = l2_normalize(matrix, matrix.rank() - 1)?; + let similarity = query_norm + .matmul(&matrix_norm.transpose(D::Minus1, D::Minus2)?)? + .squeeze(D::Minus1)?; + Ok(similarity) +} diff --git a/tests/messy_test.rs b/tests/messy_test.rs index f9b0d9c..4349148 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,27 +1,112 @@ // use std::io::Cursor; -use std::time::Instant; - -use aha::utils::tensor_utils::interpolate_nearest_1d; -use anyhow::{Result, anyhow}; -use candle_core::Tensor; -use sentencepiece::SentencePieceProcessor; +use std::fs::File; // use symphonia::core::io::MediaSourceStream; +use std::io::{Read, Seek}; +use std::{io::Cursor, time::Instant}; + +use aha::utils::{load_tensor_from_pt, tensor_utils::interpolate_nearest_1d}; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::{Result, anyhow}; +use byteorder::{LittleEndian, ReadBytesExt}; +use candle_core::{Shape, Tensor}; +use sentencepiece::SentencePieceProcessor; +use zip::ZipArchive; #[test] fn messy_test() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture let device = &candle_core::Device::Cpu; - let save_dir = + let save_dir: String = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; - let model_path = format!("{}/IndexTeam/IndexTTS-2", save_dir); - let bpe_path = model_path.to_string() + "/bpe.model"; - let tokenizer = SentencePieceProcessor::open(bpe_path) - .map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?; - let tokens = tokenizer - .encode("你好啊") - .map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?; - println!("tokens: {:?}", tokens); + let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir); + let emo_matrix_path = model_path.clone() + "/feat2.pt"; + let t_emo = load_tensor_from_pt( + &emo_matrix_path, + "feat2/data/0", + Shape::from_dims(&[73, 1280]), + &device, + )?; + println!("t_emo: {}", t_emo); + let skp_matrix_path = model_path + "/feat1.pt"; + let t_skp = load_tensor_from_pt( + &skp_matrix_path, + "feat1/data/0", + Shape::from_dims(&[73, 192]), + &device, + )?; + println!("t_skp: {}", t_skp); + // let file = File::open(emo_matrix_path)?; + // let mut archive = ZipArchive::new(file)?; + // // 列出所有文件(调试用) + // for i in 0..archive.len() { + // let file = archive.by_index(i)?; + // println!("File: {} ({} bytes)", file.name(), file.size()); + // } + // // 读取原始字节数据 + // let mut data_file = archive.by_name("feat2/data/0")?; + // let mut buffer = Vec::new(); + // data_file.read_to_end(&mut buffer)?; + // // 将字节转换为 f32 (little endian) + // let mut cursor = Cursor::new(buffer); + // let num_elements = 73 * 1280; // 93,440 + // let mut data = Vec::with_capacity(num_elements); + + // for _ in 0..num_elements { + // let val = cursor.read_f32::()?; + // data.push(val); + // } + // let t = Tensor::from_vec(data, (73, 1280), device)?; + // println!("t: {}", t); + // let message = r#" + // { + // "model": "index-tts2", + // "messages": [ + // { + // "role": "user", + // "content": [ + // { + // "type": "audio", + // "audio_url": + // { + // "url": "file:///home/jhq/Videos/voice_01.wav" + // } + // }, + // { + // "type": "text", + // "text": "你好啊" + // } + // ] + // } + // ], + // "metadata": {"emo_vector": "[0, 0, 0, 0, 0, 0, 0.45, 0]"} + // } + // "#; + // let mes: ChatCompletionParameters = serde_json::from_str(message)?; + + // if let Some(map) = &mes.metadata + // && let Some(emo_vector_str) = map.get("emo_vector") + // { + // match serde_json::from_str::>(emo_vector_str) { + // Ok(emo_vector) => { + // println!("Parsed emo_vector: {:?}", emo_vector); + // // 现在 emo_vector 是 Vec: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.45, 0.0] + // } + // Err(e) => { + // eprintln!("Failed to parse emo_vector: {}", e); + // } + // } + // } + // let save_dir = + // aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + // let model_path = format!("{}/IndexTeam/IndexTTS-2", save_dir); + // let bpe_path = model_path.to_string() + "/bpe.model"; + // let tokenizer = SentencePieceProcessor::open(bpe_path) + // .map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?; + // let tokens = tokenizer + // .encode("你好啊") + // .map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?; + // println!("tokens: {:?}", tokens); // let t = Tensor::arange(0.0f32, 40.0, device)?.broadcast_as((1, 40, 40))?; // println!("t: {}", t); // let i_start = Instant::now(); diff --git a/tests/test_index_tts2.rs b/tests/test_index_tts2.rs index 1fd1745..dc33f36 100644 --- a/tests/test_index_tts2.rs +++ b/tests/test_index_tts2.rs @@ -1,8 +1,11 @@ use std::time::Instant; -use anyhow::Result; -use aha::models::index_tts2::{generate::IndexTTS2Generate, utils::download_index_tts2_need_model}; +use aha::{ + models::index_tts2::{generate::IndexTTS2Generate, utils::download_index_tts2_need_model}, + utils::audio_utils::extract_and_save_audio_from_response, +}; use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; #[tokio::test] async fn index_tts2_generate() -> Result<()> { @@ -24,10 +27,10 @@ async fn index_tts2_generate() -> Result<()> { { "url": "file:///home/jhq/Videos/voice_01.wav" } - }, + }, { "type": "text", - "text": "你好啊" + "text": "你好啊,吃饭了吗" } ] } @@ -42,7 +45,11 @@ async fn index_tts2_generate() -> Result<()> { let i_start = Instant::now(); let generate = voxcpm_generate.generate(mes)?; + let save_path = extract_and_save_audio_from_response(&generate, "./")?; + for path in save_path { + println!("save audio: {}", path); + } let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) -} \ No newline at end of file +} diff --git a/tests/test_qwen3_asr.rs b/tests/test_qwen3_asr.rs index ad237c3..dc75cd4 100644 --- a/tests/test_qwen3_asr.rs +++ b/tests/test_qwen3_asr.rs @@ -9,7 +9,7 @@ fn qwen3_asr_generate() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda qwen3_asr_generate -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; - let model_path = format!("{}/Qwen/Qwen3-ASR-0.6B/", save_dir); //Qwen/Qwen3-ASR-1.7B + let model_path = format!("{}/Qwen/Qwen3-ASR-0.6B/", save_dir); //Qwen/Qwen3-ASR-1.7B let message = r#" { "model": "qwen3-asr", diff --git a/tests/weight_test.rs b/tests/weight_test.rs index 0ab8894..432b173 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -2,7 +2,11 @@ use std::collections::HashMap; use aha::utils::{find_type_files, get_device, read_pth_tensor_info_cycle}; use anyhow::Result; -use candle_core::{Device, pickle::read_all_with_key, safetensors}; +use candle_core::{ + Device, + pickle::{read_all_with_key, read_pth_tensor_info}, + safetensors, +}; use candle_nn::VarBuilder; #[test] @@ -203,19 +207,26 @@ fn index_tts2_weight() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda index_tts2_weight -r -- --nocapture let save_dir: String = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; - let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir); - let s2mel_path = model_path+ "/s2mel.pth"; + let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir); + let bigvgan_path = format!( + "{}/nv-community/bigvgan_v2_22khz_80band_256x/bigvgan_generator.pt", + save_dir + ); + // let gpt_path = model_path+ "/gpt.pth"; + + // let spk_matrix_path = model_path+ "/feat1.pt"; + // let s2mel_path = model_path+ "/s2mel.pth"; // let wac2vec2_path = model_path+ "/wav2vec2bert_stats.pt"; // let model_path = format!("{}/iic/speech_campplus_sv_zh-cn_16k-common/", save_dir); // let campplus_path = model_path+ "/campplus_cn_common.bin"; // let model_list = find_type_files(&model_path, "safetensors")?; - let model_list = vec![s2mel_path]; - // let mut dict_to_hashmap = HashMap::new(); - // let mut dtype = candle_core::DType::F32; + let model_list = vec![bigvgan_path]; + // // 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"))?; - // let dict = read_all_with_key(m, Some("net"))?; - let dict = read_pth_tensor_info_cycle(m, Some("net.cfm"))?; + let dict = read_all_with_key(m, Some("generator"))?; + // let dict = read_pth_tensor_info_cycle(m, Some("net.cfm"))?; // dtype = dict[0].1.dtype(); for (k, v) in dict { // if k.contains("model") { @@ -230,9 +241,9 @@ fn index_tts2_weight() -> Result<()> { // let model_list = vec![semantic_codec_path]; // for m in model_list { // let weights = safetensors::load(m, &device)?; - // for (key, tensor) in weights.iter() { + // for (key, tensor) in weights.iter() { // println!("=== {} === {:?}", key, tensor.shape()); // } // } Ok(()) -} \ No newline at end of file +}