updata index tts
This commit is contained in:
Generated
+1
@@ -50,6 +50,7 @@ dependencies = [
|
||||
"minijinja",
|
||||
"modelscope",
|
||||
"num",
|
||||
"rand 0.9.2",
|
||||
"rayon",
|
||||
"realfft",
|
||||
"regex",
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
## 功能特性
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,4 +18,4 @@ pub struct FeatureExtractor {
|
||||
|
||||
fn default_sampling_rate() -> usize {
|
||||
16000
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
@@ -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,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();
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
+1505
-121
File diff suppressed because it is too large
Load Diff
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -18,4 +18,4 @@ fn default_num_quantizers() -> usize {
|
||||
|
||||
fn default_downsample_scale() -> usize {
|
||||
1
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
pub mod config;
|
||||
pub mod model;
|
||||
pub mod config;
|
||||
@@ -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,3 +1,4 @@
|
||||
pub mod bigvgan;
|
||||
pub mod campplus;
|
||||
pub mod common;
|
||||
pub mod deepseek_ocr;
|
||||
|
||||
@@ -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,4 +1,3 @@
|
||||
|
||||
use aha_openai_dive::v1::resources::chat::{
|
||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||
};
|
||||
|
||||
@@ -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>()?;
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,2 +1,2 @@
|
||||
pub mod config;
|
||||
pub mod model;
|
||||
pub mod model;
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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();
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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(())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user