updata index tts

This commit is contained in:
jhqxxx
2026-02-14 15:52:30 +08:00
parent 7c832e0ce8
commit 2c34fc2d79
40 changed files with 2602 additions and 540 deletions
Generated
+1
View File
@@ -50,6 +50,7 @@ dependencies = [
"minijinja",
"modelscope",
"num",
"rand 0.9.2",
"rayon",
"realfft",
"regex",
+1
View File
@@ -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"]
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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
```
## 功能特性
+1 -2
View File
@@ -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();
+46
View File
@@ -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<usize>,
pub upsample_kernel_sizes: Vec<usize>,
pub upsample_initial_channel: usize,
pub resblock_kernel_sizes: Vec<usize>,
pub resblock_dilation_sizes: Vec<Vec<usize>>,
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<usize>,
pub cqtd_hop_lengths: Vec<usize>,
pub cqtd_n_octaves: Vec<usize>,
pub cqtd_bins_per_octaves: Vec<usize>,
pub mpd_reshapes: Vec<usize>,
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<usize>,
pub fmax_for_loss: Option<usize>,
pub normalize_volume: bool,
pub num_workers: usize,
}
+333
View File
@@ -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<usize>) -> Result<Self> {
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<Tensor> {
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<usize>) -> Result<Self> {
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<Tensor> {
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<Self> {
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<Tensor> {
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<Self> {
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<Tensor> {
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<WNConv1d>,
convs2: Vec<WNConv1d>,
activations: Vec<TorchActivation1d>,
}
impl AMPBlock1 {
pub fn new(
vb: VarBuilder,
channels: usize,
kernel_size: usize,
dilation: Vec<usize>,
) -> Result<Self> {
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<Tensor> {
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<WNConvTranspose1d>,
resblocks: Vec<AMPBlock1>,
activation_post: TorchActivation1d,
conv_post: WNConv1d,
use_tanh_at_final: bool,
}
impl BigVGAN {
pub fn new(vb: VarBuilder, cfg: &BigVGANConfig) -> Result<Self> {
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<Tensor> {
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)
}
}
+10 -6
View File
@@ -25,7 +25,11 @@ impl Shortcut {
) -> Result<Self> {
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<Tensor> {
@@ -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,
+257 -14
View File
@@ -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<LayerNorm> {
pub fn get_layer_norm_without_weight(vb: VarBuilder, eps: f64, dim: usize) -> Result<LayerNorm> {
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<Tensor> {
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<Self> {
Ok(Self { dim })
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
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<Tensor> {
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<Self> {
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<Tensor> {
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<Self> {
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<Tensor> {
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<usize>,
) -> Result<Self> {
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<Tensor> {
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<GPT2Block>,
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<Self> {
// 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<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
b_size,
seq_len,
0,
xs.device(),
)?)
}
};
for 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()
}
}
}
+1 -1
View File
@@ -18,4 +18,4 @@ pub struct FeatureExtractor {
fn default_sampling_rate() -> usize {
16000
}
}
+2 -2
View File
@@ -1,3 +1,3 @@
pub mod seamless_m4t_feature_extractor;
pub mod config;
pub mod feature_extraction_whisper;
pub mod config;
pub mod seamless_m4t_feature_extractor;
@@ -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)?;
-2
View File
@@ -11,8 +11,6 @@ pub struct GlmAsrNanoProcessorConfig {
pub max_audio_len: usize,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrNanoConfig {
pub audio_config: GlmAsrAudioConfig,
+1 -1
View File
@@ -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<Self> {
let path = path.to_string();
+3 -3
View File
@@ -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,
}
}
}
+20 -39
View File
@@ -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<String>,
// 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<DType>) -> Result<Self> {
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<ChatCompletionResponse> {
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)
}
}
+2 -2
View File
@@ -1,5 +1,5 @@
pub mod config;
pub mod generate;
pub mod model;
pub mod processor;
pub mod utils;
// pub mod processor;
pub mod utils;
File diff suppressed because it is too large Load Diff
-203
View File
@@ -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<Self> {
// 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<Tensor> {
// 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<Tensor> {
// 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))
}
}
+3 -2
View File
@@ -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(())
}
+1 -1
View File
@@ -18,4 +18,4 @@ fn default_num_quantizers() -> usize {
fn default_downsample_scale() -> usize {
1
}
}
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod config;
pub mod model;
pub mod config;
+42 -7
View File
@@ -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<Self> {
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<Tensor> {
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<Tensor> {
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<Conv1d>,
@@ -234,7 +269,7 @@ pub struct RepCodec {
encoder_1: Linear,
decoder_0: VocosBackbone,
decoder_1: Linear,
quantizer: ResidualVQ,
pub quantizer: ResidualVQ,
}
impl RepCodec {
+1
View File
@@ -1,3 +1,4 @@
pub mod bigvgan;
pub mod campplus;
pub mod common;
pub mod deepseek_ocr;
+4 -4
View File
@@ -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,
-1
View File
@@ -1,4 +1,3 @@
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
+1 -1
View File
@@ -338,7 +338,7 @@ impl Qwen3ASRThinker {
) -> Result<Tensor> {
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::<u32>()?;
+5 -10
View File
@@ -82,10 +82,7 @@ impl Qwen3AsrProcessor {
pub fn process_audio(&self, mes: &ChatCompletionParameters) -> Result<Vec<Tensor>> {
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
}
}
+1 -1
View File
@@ -58,4 +58,4 @@ pub struct W2VBert2_0Config {
pub use_weighted_layer_sum: bool,
pub vocab_size: Option<usize>,
pub xvector_output_dim: usize,
}
}
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod config;
pub mod model;
pub mod model;
+21 -24
View File
@@ -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<Tensor> = 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<Tensor>,
// masked_spec_embed: Option<Tensor>,
encoder: Wav2Vec2BertEncoder,
// config.add_adapter is false, adapter is None, Wav2Vec2BertAdapter not complish
// adapter: Option<Wav2Vec2BertAdapter>,
@@ -543,21 +540,21 @@ impl W2VBert2_0Model {
pub fn new(vb: VarBuilder, config: &W2VBert2_0Config) -> Result<Self> {
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,
})
}
+1 -1
View File
@@ -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)
+10 -12
View File
@@ -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);
}
}
}
+111 -5
View File
@@ -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<Tensor> {
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<Tensor> {
// 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<Tensor> {
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<Tensor> {
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<f32> = 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))?)
}
+41 -11
View File
@@ -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<P: AsRef<std::path::Path>>(
}
}
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<P: AsRef<std::path::Path>>(
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<Tensor> {
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::<LittleEndian>()?;
data.push(val);
}
let t = Tensor::from_vec(data, shape, device)?;
Ok(t)
}
+36 -26
View File
@@ -14,7 +14,7 @@ pub fn masked_fill_zeros(hidden_states: &Tensor, mask: &Tensor) -> Result<Tensor
let mask = mask
.unsqueeze(D::Minus1)?
.broadcast_as(hidden_states.shape())?;
let hidden_states = mask.where_cond(&hidden_states, &on_false)?;
let hidden_states = mask.where_cond(hidden_states, &on_false)?;
Ok(hidden_states)
}
@@ -431,19 +431,25 @@ pub fn compute_1d_coords(
output_size: usize,
align_corner: Option<bool>,
) -> Result<Vec<f32>> {
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<Tensor>
"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::<f32>()?;
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<u32>) -> Result<Tensor>
let mask = x.broadcast_lt(&length)?;
Ok(mask)
}
pub fn cosine_similarity(query_vector: &Tensor, matrix: &Tensor) -> Result<Tensor> {
// 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)
}
+100 -15
View File
@@ -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::<LittleEndian>()?;
// 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::<Vec<f32>>(emo_vector_str) {
// Ok(emo_vector) => {
// println!("Parsed emo_vector: {:?}", emo_vector);
// // 现在 emo_vector 是 Vec<f32>: [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();
+12 -5
View File
@@ -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(())
}
}
+1 -1
View File
@@ -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",
+21 -10
View File
@@ -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(())
}
}