updata index tts
This commit is contained in:
Generated
+1
@@ -50,6 +50,7 @@ dependencies = [
|
|||||||
"minijinja",
|
"minijinja",
|
||||||
"modelscope",
|
"modelscope",
|
||||||
"num",
|
"num",
|
||||||
|
"rand 0.9.2",
|
||||||
"rayon",
|
"rayon",
|
||||||
"realfft",
|
"realfft",
|
||||||
"regex",
|
"regex",
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ half = "2.7.1"
|
|||||||
byteorder = "1.5.0"
|
byteorder = "1.5.0"
|
||||||
sentencepiece = "0.13.1"
|
sentencepiece = "0.13.1"
|
||||||
regex = "1.12.3"
|
regex = "1.12.3"
|
||||||
|
rand = "0.9.2"
|
||||||
|
|
||||||
[features]
|
[features]
|
||||||
flash-attn = ["candle-flash-attn"]
|
flash-attn = ["candle-flash-attn"]
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ cargo build --release --features ffmpeg
|
|||||||
```bash
|
```bash
|
||||||
# Install build dependencies
|
# Install build dependencies
|
||||||
sudo apt-get update
|
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
|
# For FFmpeg feature
|
||||||
sudo apt-get install -y ffmpeg libavutil-dev libavcodec-dev \
|
sudo apt-get install -y ffmpeg libavutil-dev libavcodec-dev \
|
||||||
@@ -166,7 +166,7 @@ cargo build --release
|
|||||||
# Follow Linux instructions inside WSL2
|
# Follow Linux instructions inside WSL2
|
||||||
wsl
|
wsl
|
||||||
sudo apt-get update
|
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
|
## Feature Flags
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ cargo build --release --features ffmpeg
|
|||||||
```bash
|
```bash
|
||||||
# 安装构建依赖
|
# 安装构建依赖
|
||||||
sudo apt-get update
|
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 功能所需
|
# FFmpeg 功能所需
|
||||||
sudo apt-get install -y ffmpeg libavutil-dev libavcodec-dev \
|
sudo apt-get install -y ffmpeg libavutil-dev libavcodec-dev \
|
||||||
@@ -165,7 +165,7 @@ cargo build --release
|
|||||||
# 在 WSL2 中按照 Linux 说明操作
|
# 在 WSL2 中按照 Linux 说明操作
|
||||||
wsl
|
wsl
|
||||||
sudo apt-get update
|
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 anyhow::{Ok, Result};
|
||||||
|
|
||||||
use crate::exec::ExecModel;
|
use crate::exec::ExecModel;
|
||||||
|
use crate::models::GenerateModel;
|
||||||
use crate::models::qwen3_asr::generate::Qwen3AsrGenerateModel;
|
use crate::models::qwen3_asr::generate::Qwen3AsrGenerateModel;
|
||||||
use crate::models::{GenerateModel};
|
|
||||||
|
|
||||||
pub struct Qwen3ASRExec;
|
pub struct Qwen3ASRExec;
|
||||||
|
|
||||||
impl ExecModel for Qwen3ASRExec {
|
impl ExecModel for Qwen3ASRExec {
|
||||||
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
|
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let mut model = Qwen3AsrGenerateModel::init(weight_path, None, None)?;
|
let mut model = Qwen3AsrGenerateModel::init(weight_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
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> {
|
) -> Result<Self> {
|
||||||
let conv_0 = get_conv2d(vb.pp("0"), in_c, out_c, ks, padding, 1, 1, 1, bias)?;
|
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)?;
|
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> {
|
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)?;
|
let indices = Tensor::arange(0u32, half_h as u32, x.device())?.affine(2.0, 0.0)?;
|
||||||
x = x.index_select(&indices, 2)?;
|
x = x.index_select(&indices, 2)?;
|
||||||
}
|
}
|
||||||
x = self.bn_1.forward_t(&x, false)?;
|
x = self.bn_1.forward_t(&x, false)?;
|
||||||
Ok(x)
|
Ok(x)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -106,7 +110,7 @@ impl BasicResBlock {
|
|||||||
} else {
|
} else {
|
||||||
xs = xs.add(&residual)?;
|
xs = xs.add(&residual)?;
|
||||||
}
|
}
|
||||||
xs = xs.relu()?;
|
xs = xs.relu()?;
|
||||||
Ok(xs)
|
Ok(xs)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -435,7 +439,7 @@ impl DenseLayer {
|
|||||||
.forward(&xs.unsqueeze(D::Minus1)?)?
|
.forward(&xs.unsqueeze(D::Minus1)?)?
|
||||||
.squeeze(D::Minus1)?
|
.squeeze(D::Minus1)?
|
||||||
} else {
|
} else {
|
||||||
self.linear.forward(&xs)?
|
self.linear.forward(xs)?
|
||||||
};
|
};
|
||||||
let xs = self.nonlinear.forward_t(&xs, false)?;
|
let xs = self.nonlinear.forward_t(&xs, false)?;
|
||||||
Ok(xs)
|
Ok(xs)
|
||||||
@@ -463,7 +467,7 @@ impl XVector {
|
|||||||
let mut channels = init_channels;
|
let mut channels = init_channels;
|
||||||
let mut blocks = vec![];
|
let mut blocks = vec![];
|
||||||
let mut transits = 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() {
|
for (i, (num_layers, ks, dilation)) in params.iter().enumerate() {
|
||||||
let block = CAMDenseTDNNBlock::new(
|
let block = CAMDenseTDNNBlock::new(
|
||||||
vb.pp(format!("block{}", i + 1)),
|
vb.pp(format!("block{}", i + 1)),
|
||||||
@@ -477,7 +481,7 @@ impl XVector {
|
|||||||
false,
|
false,
|
||||||
)?;
|
)?;
|
||||||
blocks.push(block);
|
blocks.push(block);
|
||||||
channels = channels + num_layers * growth_rate;
|
channels += num_layers * growth_rate;
|
||||||
let transit = TransitLayer::new(
|
let transit = TransitLayer::new(
|
||||||
vb.pp(format!("transit{}", i + 1)),
|
vb.pp(format!("transit{}", i + 1)),
|
||||||
channels,
|
channels,
|
||||||
|
|||||||
+257
-14
@@ -2,8 +2,8 @@ use anyhow::{Result, anyhow};
|
|||||||
use candle_core::{D, IndexOp, Tensor};
|
use candle_core::{D, IndexOp, Tensor};
|
||||||
use candle_nn::{
|
use candle_nn::{
|
||||||
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig,
|
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig,
|
||||||
ConvTranspose1d, ConvTranspose1dConfig, Embedding, GroupNorm, Init, LayerNorm, LayerNormConfig,
|
ConvTranspose1d, ConvTranspose1dConfig, Embedding, Init, LayerNorm, LayerNormConfig, Linear,
|
||||||
Linear, Module, ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d,
|
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,
|
conv2d_no_bias, embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm,
|
||||||
};
|
};
|
||||||
use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator};
|
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)
|
Ok(norm)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get_layer_norm_without_weight(
|
pub fn get_layer_norm_without_weight(vb: VarBuilder, eps: f64, dim: usize) -> Result<LayerNorm> {
|
||||||
vb: VarBuilder,
|
|
||||||
eps: f64,
|
|
||||||
dim: usize,
|
|
||||||
) -> Result<LayerNorm> {
|
|
||||||
let weight = Tensor::ones(dim, vb.dtype(), vb.device())?;
|
let weight = Tensor::ones(dim, vb.dtype(), vb.device())?;
|
||||||
let bias = Tensor::zeros(dim, vb.dtype(), vb.device())?;
|
let bias = Tensor::zeros(dim, vb.dtype(), vb.device())?;
|
||||||
Ok(LayerNorm::new(weight, bias, eps))
|
Ok(LayerNorm::new(weight, bias, eps))
|
||||||
@@ -1027,11 +1023,25 @@ impl GLU {
|
|||||||
Ok(Self { dim })
|
Ok(Self { dim })
|
||||||
}
|
}
|
||||||
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
||||||
let half_dim = xs.dim(self.dim)? / 2;
|
let x_ = xs.chunk(2, self.dim)?;
|
||||||
let a = xs.narrow(self.dim, 0, half_dim)?;
|
let x_1 = sigmoid(x_[1].as_ref())?;
|
||||||
let b = xs.narrow(self.dim, half_dim, half_dim)?;
|
let xs = x_1.mul(x_[0].as_ref())?;
|
||||||
let b = sigmoid(&b)?;
|
Ok(xs)
|
||||||
let xs = a.mul(&b)?;
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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)
|
Ok(xs)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1103,8 +1113,8 @@ impl WNConvTranspose1d {
|
|||||||
let normalized_weight = weight_v.broadcast_div(&weight_norm)?;
|
let normalized_weight = weight_v.broadcast_div(&weight_norm)?;
|
||||||
let scaled_weight = normalized_weight.broadcast_mul(&weight_g)?;
|
let scaled_weight = normalized_weight.broadcast_mul(&weight_g)?;
|
||||||
let config = ConvTranspose1dConfig {
|
let config = ConvTranspose1dConfig {
|
||||||
padding: padding,
|
padding,
|
||||||
output_padding: output_padding,
|
output_padding,
|
||||||
stride,
|
stride,
|
||||||
dilation,
|
dilation,
|
||||||
groups,
|
groups,
|
||||||
@@ -1184,3 +1194,236 @@ pub fn mish(xs: &Tensor) -> Result<Tensor> {
|
|||||||
let xs = xs.mul(&tanh)?;
|
let xs = xs.mul(&tanh)?;
|
||||||
Ok(xs)
|
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 {
|
fn default_sampling_rate() -> usize {
|
||||||
16000
|
16000
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,3 +1,3 @@
|
|||||||
pub mod seamless_m4t_feature_extractor;
|
pub mod config;
|
||||||
pub mod feature_extraction_whisper;
|
pub mod feature_extraction_whisper;
|
||||||
pub mod config;
|
pub mod seamless_m4t_feature_extractor;
|
||||||
|
|||||||
@@ -82,7 +82,7 @@ impl SeamlessM4TFeatureExtractor {
|
|||||||
0.97,
|
0.97,
|
||||||
Some(&self.mel_filters),
|
Some(&self.mel_filters),
|
||||||
Some("log"),
|
Some("log"),
|
||||||
1.192092955078125e-07,
|
1.192_092_9e-7,
|
||||||
true,
|
true,
|
||||||
)?
|
)?
|
||||||
.transpose(D::Minus1, D::Minus2)?;
|
.transpose(D::Minus1, D::Minus2)?;
|
||||||
|
|||||||
@@ -11,8 +11,6 @@ pub struct GlmAsrNanoProcessorConfig {
|
|||||||
pub max_audio_len: usize,
|
pub max_audio_len: usize,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
||||||
pub struct GlmAsrNanoConfig {
|
pub struct GlmAsrNanoConfig {
|
||||||
pub audio_config: GlmAsrAudioConfig,
|
pub audio_config: GlmAsrAudioConfig,
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
|
|
||||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||||
@@ -27,6 +26,7 @@ pub struct GlmAsrNanoProcessor {
|
|||||||
whisper_feature_extrator: WhisperFeatureExtractor,
|
whisper_feature_extrator: WhisperFeatureExtractor,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
impl GlmAsrNanoProcessor {
|
impl GlmAsrNanoProcessor {
|
||||||
pub fn new(path: &str, device: &Device, dtype: DType) -> Result<Self> {
|
pub fn new(path: &str, device: &Device, dtype: DType) -> Result<Self> {
|
||||||
let path = path.to_string();
|
let path = path.to_string();
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ use crate::models::mask_gct::config::SemanticCodec;
|
|||||||
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
||||||
pub struct IndexTTS2Config {
|
pub struct IndexTTS2Config {
|
||||||
pub dataset: Dataset,
|
pub dataset: Dataset,
|
||||||
pub gpt: Gpt,
|
pub gpt: GptConfig,
|
||||||
pub semantic_codec: SemanticCodec,
|
pub semantic_codec: SemanticCodec,
|
||||||
pub s2mel: S2MelConfig,
|
pub s2mel: S2MelConfig,
|
||||||
pub gpt_checkpoint: String,
|
pub gpt_checkpoint: String,
|
||||||
@@ -39,7 +39,7 @@ pub struct Mel {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
||||||
pub struct Gpt {
|
pub struct GptConfig {
|
||||||
pub model_dim: usize,
|
pub model_dim: usize,
|
||||||
pub max_mel_tokens: usize,
|
pub max_mel_tokens: usize,
|
||||||
pub max_text_tokens: usize,
|
pub max_text_tokens: usize,
|
||||||
@@ -206,7 +206,7 @@ impl DiTModelArgs {
|
|||||||
has_cross_attention: false,
|
has_cross_attention: false,
|
||||||
context_dim: 0,
|
context_dim: 0,
|
||||||
uvit_skip_connection: config.uvit_skip_connection,
|
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 anyhow::{Result, anyhow};
|
||||||
|
use base64::{Engine, prelude::BASE64_STANDARD};
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use sentencepiece::SentencePieceProcessor;
|
use sentencepiece::SentencePieceProcessor;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::index_tts2::{
|
models::index_tts2::{
|
||||||
config::IndexTTS2Config,
|
config::IndexTTS2Config, model::IndexTTS2Model, utils::tokenize_by_cjk_char,
|
||||||
model::IndexTTS2Model,
|
|
||||||
processor::IndexTTS2Processor,
|
|
||||||
utils::{TextNormalizer, tokenize_by_cjk_char},
|
|
||||||
},
|
},
|
||||||
tokenizer::sentencepiece_encode,
|
tokenizer::sentencepiece_encode,
|
||||||
utils::{
|
utils::{
|
||||||
audio_utils::extract_audio_url, extract_user_text, get_default_save_dir, get_device,
|
audio_utils::get_audio_wav_u8, build_audio_completion_response, extract_user_text,
|
||||||
get_dtype,
|
get_default_save_dir, get_device,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct IndexTTS2Generate {
|
pub struct IndexTTS2Generate {
|
||||||
processor: IndexTTS2Processor,
|
|
||||||
tokenizer: SentencePieceProcessor,
|
tokenizer: SentencePieceProcessor,
|
||||||
config: IndexTTS2Config,
|
// config: IndexTTS2Config,
|
||||||
cache_spk_audio_prompt: Option<String>,
|
|
||||||
model: IndexTTS2Model,
|
model: IndexTTS2Model,
|
||||||
device: Device,
|
device: Device,
|
||||||
|
sample_rate: u32,
|
||||||
|
model_name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
impl IndexTTS2Generate {
|
impl IndexTTS2Generate {
|
||||||
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||||
let config_path = path.to_string() + "/config.yaml";
|
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 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 config: IndexTTS2Config = serde_yaml::from_slice(&std::fs::read(config_path)?)?;
|
||||||
let device = get_device(device);
|
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 bpe_path = path.to_string() + "/bpe.model";
|
||||||
let tokenizer = SentencePieceProcessor::open(bpe_path)
|
let tokenizer = SentencePieceProcessor::open(bpe_path)
|
||||||
.map_err(|e| anyhow!(format!("load bpe,model file error:{}", e)))?;
|
.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 {
|
Ok(Self {
|
||||||
processor,
|
|
||||||
tokenizer,
|
tokenizer,
|
||||||
config,
|
// config,
|
||||||
cache_spk_audio_prompt: None,
|
|
||||||
model,
|
model,
|
||||||
device,
|
device,
|
||||||
|
sample_rate: 22050,
|
||||||
|
model_name: "index-tts2".to_string(),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn use_prompt(&self, mes: &ChatCompletionParameters) -> bool {
|
pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
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<()> {
|
|
||||||
let text = extract_user_text(&mes)?;
|
let text = extract_user_text(&mes)?;
|
||||||
let text = tokenize_by_cjk_char(&text, true);
|
let text = tokenize_by_cjk_char(&text, true);
|
||||||
let input_ids = sentencepiece_encode(&text, &self.tokenizer, &self.device)?;
|
let input_ids = sentencepiece_encode(&text, &self.tokenizer, &self.device)?;
|
||||||
let (audio_22k, audio_16k) = if self.use_prompt(&mes) {
|
|
||||||
(None, None)
|
let audio = self.model.forward(&input_ids, &mes)?;
|
||||||
} else {
|
let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate)?;
|
||||||
let (audio_22k, audio_16k, prompt) = self.processor.process_info(&mes)?;
|
let base64_audio = BASE64_STANDARD.encode(wav_u8);
|
||||||
self.cache_spk_audio_prompt = Some(prompt);
|
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
||||||
(Some(audio_22k), Some(audio_16k))
|
Ok(response)
|
||||||
};
|
|
||||||
let _ = self.model.forward(&input_ids, audio_22k.as_ref(), audio_16k.as_ref())?;
|
|
||||||
Ok(())
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
pub mod config;
|
pub mod config;
|
||||||
pub mod generate;
|
pub mod generate;
|
||||||
pub mod model;
|
pub mod model;
|
||||||
pub mod processor;
|
// pub mod processor;
|
||||||
pub mod utils;
|
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 anyhow::{Result, anyhow};
|
||||||
use regex::Regex;
|
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 mask_gct = "amphion/MaskGCT";
|
||||||
// let campplus= "funasr/campplus"; // huggingface
|
// let campplus= "funasr/campplus"; // huggingface
|
||||||
let campplus = "iic/speech_campplus_sv_zh-cn_16k-common"; // modelscope
|
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(w2v_bert2_0, &save_dir, 3).await?;
|
||||||
download_model(mask_gct, &save_dir, 3).await?;
|
download_model(mask_gct, &save_dir, 3).await?;
|
||||||
download_model(campplus, &save_dir, 3).await?;
|
download_model(campplus, &save_dir, 3).await?;
|
||||||
|
download_model(bigvgan, &save_dir, 3).await?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -18,4 +18,4 @@ fn default_num_quantizers() -> usize {
|
|||||||
|
|
||||||
fn default_downsample_scale() -> usize {
|
fn default_downsample_scale() -> usize {
|
||||||
1
|
1
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,2 +1,2 @@
|
|||||||
|
pub mod config;
|
||||||
pub mod model;
|
pub mod model;
|
||||||
pub mod config;
|
|
||||||
@@ -48,7 +48,7 @@ impl ConvNeXtBlock {
|
|||||||
let xs = self.pwconv1.forward(&xs)?.gelu()?;
|
let xs = self.pwconv1.forward(&xs)?.gelu()?;
|
||||||
let mut xs = self.pwconv2.forward(&xs)?;
|
let mut xs = self.pwconv2.forward(&xs)?;
|
||||||
if let Some(gamma) = &self.gamma {
|
if let Some(gamma) = &self.gamma {
|
||||||
xs = xs.broadcast_mul(&gamma)?;
|
xs = xs.broadcast_mul(gamma)?;
|
||||||
}
|
}
|
||||||
let xs = xs.transpose(1, 2)?;
|
let xs = xs.transpose(1, 2)?;
|
||||||
let xs = residual.add(&xs)?;
|
let xs = residual.add(&xs)?;
|
||||||
@@ -116,10 +116,28 @@ impl FactorizedVectorQuantize {
|
|||||||
use_l2_normlize: bool,
|
use_l2_normlize: bool,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
let (in_project, out_project) = if input_dim != codebook_dim {
|
let (in_project, out_project) = if input_dim != codebook_dim {
|
||||||
let in_project =
|
let in_project = WNConv1d::new(
|
||||||
WNConv1d::new(vb.pp("in_project"), input_dim, codebook_dim, 1, 1, 0, 1, 1, true)?;
|
vb.pp("in_project"),
|
||||||
let out_project =
|
input_dim,
|
||||||
WNConv1d::new(vb.pp("out_project"), codebook_dim, input_dim, 1, 1, 0, 1, 1, true)?;
|
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))
|
(Some(in_project), Some(out_project))
|
||||||
} else {
|
} else {
|
||||||
(None, None)
|
(None, None)
|
||||||
@@ -166,6 +184,14 @@ impl FactorizedVectorQuantize {
|
|||||||
}
|
}
|
||||||
Ok((z_q, indices))
|
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 {
|
pub struct ResidualVQ {
|
||||||
@@ -210,7 +236,7 @@ impl ResidualVQ {
|
|||||||
let n_quantizers = n_quantizers.unwrap_or(self.num_quantizers);
|
let n_quantizers = n_quantizers.unwrap_or(self.num_quantizers);
|
||||||
let mut residual = xs.clone();
|
let mut residual = xs.clone();
|
||||||
let mut quantized_out = Tensor::new(0.0f32, xs.device())?.to_dtype(xs.dtype())?;
|
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 {
|
if i >= n_quantizers {
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
@@ -224,8 +250,17 @@ impl ResidualVQ {
|
|||||||
let all_quantized = Tensor::stack(&all_quantized, 0)?;
|
let all_quantized = Tensor::stack(&all_quantized, 0)?;
|
||||||
Ok((quantized_out, all_indices, all_quantized))
|
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 {
|
pub struct RepCodec {
|
||||||
downsample_scale: usize,
|
downsample_scale: usize,
|
||||||
down: Option<Conv1d>,
|
down: Option<Conv1d>,
|
||||||
@@ -234,7 +269,7 @@ pub struct RepCodec {
|
|||||||
encoder_1: Linear,
|
encoder_1: Linear,
|
||||||
decoder_0: VocosBackbone,
|
decoder_0: VocosBackbone,
|
||||||
decoder_1: Linear,
|
decoder_1: Linear,
|
||||||
quantizer: ResidualVQ,
|
pub quantizer: ResidualVQ,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl RepCodec {
|
impl RepCodec {
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
pub mod bigvgan;
|
||||||
pub mod campplus;
|
pub mod campplus;
|
||||||
pub mod common;
|
pub mod common;
|
||||||
pub mod deepseek_ocr;
|
pub mod deepseek_ocr;
|
||||||
|
|||||||
@@ -167,13 +167,13 @@ pub struct Qwen3ASRTextConfig {
|
|||||||
pub fn qwen3asr_text_config2qwen3_config(cfg: &Qwen3ASRTextConfig) -> Qwen3Config {
|
pub fn qwen3asr_text_config2qwen3_config(cfg: &Qwen3ASRTextConfig) -> Qwen3Config {
|
||||||
Qwen3Config {
|
Qwen3Config {
|
||||||
attention_bias: cfg.attention_bias,
|
attention_bias: cfg.attention_bias,
|
||||||
attention_dropout: cfg.attention_dropout as f64,
|
attention_dropout: cfg.attention_dropout,
|
||||||
bos_token_id: cfg.bos_token_id.unwrap_or(151643) as u32,
|
bos_token_id: cfg.bos_token_id.unwrap_or(151643),
|
||||||
eos_token_id: cfg.eos_token_id.unwrap_or(151645) as u32,
|
eos_token_id: cfg.eos_token_id.unwrap_or(151645),
|
||||||
head_dim: cfg.head_dim,
|
head_dim: cfg.head_dim,
|
||||||
hidden_act: cfg.hidden_act,
|
hidden_act: cfg.hidden_act,
|
||||||
hidden_size: cfg.hidden_size,
|
hidden_size: cfg.hidden_size,
|
||||||
initializer_range: cfg.initializer_range as f64,
|
initializer_range: cfg.initializer_range,
|
||||||
intermediate_size: cfg.intermediate_size,
|
intermediate_size: cfg.intermediate_size,
|
||||||
max_position_embeddings: cfg.max_position_embeddings,
|
max_position_embeddings: cfg.max_position_embeddings,
|
||||||
max_window_layers: 0,
|
max_window_layers: 0,
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
|
|
||||||
use aha_openai_dive::v1::resources::chat::{
|
use aha_openai_dive::v1::resources::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -338,7 +338,7 @@ impl Qwen3ASRThinker {
|
|||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let mut input_embeds = self.model.embed_tokens.forward(input_ids)?;
|
let mut input_embeds = self.model.embed_tokens.forward(input_ids)?;
|
||||||
if let Some(input_features) = input_features {
|
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);
|
// println!("audio_feature: {}", audio_feature);
|
||||||
let audio_mask = get_equal_mask(input_ids, self.audio_token_id)?;
|
let audio_mask = get_equal_mask(input_ids, self.audio_token_id)?;
|
||||||
let n_audio_tokens = audio_mask.sum_all()?.to_scalar::<u32>()?;
|
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>> {
|
pub fn process_audio(&self, mes: &ChatCompletionParameters) -> Result<Vec<Tensor>> {
|
||||||
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?;
|
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?;
|
||||||
audio_tensors
|
audio_tensors.iter().map(float_range_normalize).collect()
|
||||||
.iter()
|
|
||||||
.map(|audio| float_range_normalize(&audio))
|
|
||||||
.collect()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn validate_language(&self, lang: &String) -> bool {
|
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 {
|
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.replacen(&self.audio_token, &replace, 1);
|
||||||
let text = text.replace("<|audio_placeholder|>", &self.audio_token);
|
text.replace("<|audio_placeholder|>", &self.audio_token)
|
||||||
text
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn process_info(
|
pub fn process_info(
|
||||||
@@ -162,11 +158,10 @@ pub struct AudioData {
|
|||||||
|
|
||||||
pub fn get_feat_extract_output_lengths(audio_len: usize) -> usize {
|
pub fn get_feat_extract_output_lengths(audio_len: usize) -> usize {
|
||||||
let input_len_leave = audio_len % 100;
|
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;
|
let feat_lengths = (input_len_leave - 1) / 2 + 1;
|
||||||
((feat_lengths - 1) / 2 + 1 - 1) / 2 + 1 + (audio_len / 100) * 13
|
((feat_lengths - 1) / 2 + 1 - 1) / 2 + 1 + (audio_len / 100) * 13
|
||||||
} else {
|
} else {
|
||||||
(audio_len / 100) * 13
|
(audio_len / 100) * 13
|
||||||
};
|
}
|
||||||
output_len
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -58,4 +58,4 @@ pub struct W2VBert2_0Config {
|
|||||||
pub use_weighted_layer_sum: bool,
|
pub use_weighted_layer_sum: bool,
|
||||||
pub vocab_size: Option<usize>,
|
pub vocab_size: Option<usize>,
|
||||||
pub xvector_output_dim: usize,
|
pub xvector_output_dim: usize,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,2 +1,2 @@
|
|||||||
pub mod config;
|
pub mod config;
|
||||||
pub mod model;
|
pub mod model;
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ impl Wav2Vec2BertFeatureProjection {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
pub struct Wav2Vec2BertSelfAttention {
|
pub struct Wav2Vec2BertSelfAttention {
|
||||||
q_proj: Linear,
|
q_proj: Linear,
|
||||||
k_proj: Linear,
|
k_proj: Linear,
|
||||||
@@ -203,17 +204,12 @@ impl Wav2Vec2BertSelfAttention {
|
|||||||
.affine(scale, 0.0)?;
|
.affine(scale, 0.0)?;
|
||||||
if let Some(mask) = attention_mask {
|
if let Some(mask) = attention_mask {
|
||||||
// let mask = mask.unsqueeze(1)?.unsqueeze(D::Minus1)?;
|
// 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 {
|
} else {
|
||||||
Some(relative_position_attn_weights)
|
Some(relative_position_attn_weights)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
if let Some(mask) = attention_mask {
|
attention_mask.cloned()
|
||||||
// let mask = mask.unsqueeze(1)?.unsqueeze(D::Minus1)?;
|
|
||||||
Some(mask.clone())
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
}
|
|
||||||
};
|
};
|
||||||
|
|
||||||
let attn_output = eager_attention_forward(
|
let attn_output = eager_attention_forward(
|
||||||
@@ -476,7 +472,8 @@ impl Wav2Vec2BertEncoder {
|
|||||||
let attention_mask = attention_mask
|
let attention_mask = attention_mask
|
||||||
.where_cond(&attention_mask_f, &neg_inf_t)?
|
.where_cond(&attention_mask_f, &neg_inf_t)?
|
||||||
.to_dtype(xs.dtype())?
|
.to_dtype(xs.dtype())?
|
||||||
.affine(1.0, -1.0)?;
|
.affine(1.0, -1.0)?
|
||||||
|
.repeat((1, 1, seq_len, 1))?;
|
||||||
(xs, Some(attention_mask))
|
(xs, Some(attention_mask))
|
||||||
} else {
|
} else {
|
||||||
(xs.clone(), None)
|
(xs.clone(), None)
|
||||||
@@ -490,7 +487,7 @@ impl Wav2Vec2BertEncoder {
|
|||||||
let mut hidden_states: Vec<Tensor> = vec![];
|
let mut hidden_states: Vec<Tensor> = vec![];
|
||||||
let mut specify_layer_id_hidden_state = None;
|
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 {
|
if output_hidden_states {
|
||||||
hidden_states.push(xs.clone());
|
hidden_states.push(xs.clone());
|
||||||
}
|
}
|
||||||
@@ -507,7 +504,7 @@ impl Wav2Vec2BertEncoder {
|
|||||||
conv_attention_mask,
|
conv_attention_mask,
|
||||||
)?;
|
)?;
|
||||||
}
|
}
|
||||||
let hidden_states = if hidden_states.len() > 0 {
|
let hidden_states = if !hidden_states.is_empty() {
|
||||||
Some(hidden_states)
|
Some(hidden_states)
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
@@ -521,9 +518,9 @@ impl Wav2Vec2BertEncoder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct W2VBert2_0Model {
|
pub struct W2VBert2_0Model {
|
||||||
config: W2VBert2_0Config,
|
// config: W2VBert2_0Config,
|
||||||
feature_projection: Wav2Vec2BertFeatureProjection,
|
feature_projection: Wav2Vec2BertFeatureProjection,
|
||||||
masked_spec_embed: Option<Tensor>,
|
// masked_spec_embed: Option<Tensor>,
|
||||||
encoder: Wav2Vec2BertEncoder,
|
encoder: Wav2Vec2BertEncoder,
|
||||||
// config.add_adapter is false, adapter is None, Wav2Vec2BertAdapter not complish
|
// config.add_adapter is false, adapter is None, Wav2Vec2BertAdapter not complish
|
||||||
// adapter: Option<Wav2Vec2BertAdapter>,
|
// adapter: Option<Wav2Vec2BertAdapter>,
|
||||||
@@ -543,21 +540,21 @@ impl W2VBert2_0Model {
|
|||||||
pub fn new(vb: VarBuilder, config: &W2VBert2_0Config) -> Result<Self> {
|
pub fn new(vb: VarBuilder, config: &W2VBert2_0Config) -> Result<Self> {
|
||||||
let feature_projection =
|
let feature_projection =
|
||||||
Wav2Vec2BertFeatureProjection::new(vb.pp("feature_projection"), config)?;
|
Wav2Vec2BertFeatureProjection::new(vb.pp("feature_projection"), config)?;
|
||||||
let masked_spec_embed = if config.mask_time_prob > 0.0 || config.mask_time_prob > 0.0 {
|
// let masked_spec_embed = if config.mask_time_prob > 0.0 || config.mask_time_prob > 0.0 {
|
||||||
Some(
|
// Some(
|
||||||
vb.get_with_hints(config.hidden_size, "masked_spec_embed", Init::Uniform {
|
// vb.get_with_hints(config.hidden_size, "masked_spec_embed", Init::Uniform {
|
||||||
lo: 0.0,
|
// lo: 0.0,
|
||||||
up: 1.0,
|
// up: 1.0,
|
||||||
})?,
|
// })?,
|
||||||
)
|
// )
|
||||||
} else {
|
// } else {
|
||||||
None
|
// None
|
||||||
};
|
// };
|
||||||
let encoder = Wav2Vec2BertEncoder::new(vb.pp("encoder"), config)?;
|
let encoder = Wav2Vec2BertEncoder::new(vb.pp("encoder"), config)?;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
config: config.clone(),
|
// config: config.clone(),
|
||||||
feature_projection,
|
feature_projection,
|
||||||
masked_spec_embed,
|
// masked_spec_embed,
|
||||||
encoder,
|
encoder,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -350,7 +350,7 @@ impl Qwen3VLTextRotaryEmbedding {
|
|||||||
|
|
||||||
// for dim in 1..3 {
|
// for dim in 1..3 {
|
||||||
for (dim, offset) in (1..3).enumerate() {
|
for (dim, offset) in (1..3).enumerate() {
|
||||||
let dim = dim +1;
|
let dim = dim + 1;
|
||||||
let length = mrope_section[dim];
|
let length = mrope_section[dim];
|
||||||
let idx = Tensor::arange_step(offset as u32, length as u32, 3, freqs.device())?;
|
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)
|
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 {
|
if let Value::Object(tokens_map) = added_tokens_decoder {
|
||||||
for (_, token_info) in tokens_map {
|
for (_, token_info) in tokens_map {
|
||||||
if let Value::Object(token_obj) = token_info {
|
if let Value::Object(token_obj) = token_info
|
||||||
if let Some(content_val) = token_obj.get("content") {
|
&& let Some(content_val) = token_obj.get("content")
|
||||||
if let Some(content) = content_val.as_str() {
|
&& let Some(content) = content_val.as_str()
|
||||||
let special = token_obj
|
{
|
||||||
.get("special")
|
let special = token_obj
|
||||||
.and_then(|v| v.as_bool())
|
.get("special")
|
||||||
.unwrap_or(false);
|
.and_then(|v| v.as_bool())
|
||||||
|
.unwrap_or(false);
|
||||||
|
|
||||||
let added_token =
|
let added_token = AddedToken::from(content.to_string(), special);
|
||||||
AddedToken::from(content.to_string(), special);
|
special_tokens.push(added_token);
|
||||||
special_tokens.push(added_token);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+111
-5
@@ -32,7 +32,9 @@ use symphonia::core::meta::MetadataOptions;
|
|||||||
use symphonia::core::probe::Hint;
|
use symphonia::core::probe::Hint;
|
||||||
|
|
||||||
use crate::utils::get_default_save_dir;
|
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)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
@@ -42,12 +44,12 @@ pub enum ResamplingMethod {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 零阶修正贝塞尔函数 I0
|
// 零阶修正贝塞尔函数 I0
|
||||||
fn i0(x: f32) -> f32 {
|
pub fn i0(x: f32) -> f32 {
|
||||||
let mut result = 1.0;
|
let mut result = 1.0;
|
||||||
let mut term = 1.0;
|
let mut term = 1.0;
|
||||||
let half_x_sq = x * x / 4.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;
|
term = term * half_x_sq / (k * k) as f32;
|
||||||
result += term;
|
result += term;
|
||||||
|
|
||||||
@@ -1013,6 +1015,40 @@ pub fn crate_hamming_window(
|
|||||||
Ok(Tensor::from_vec(window, window_size, device)?.to_dtype(dtype)?)
|
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)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
pub enum MelScale {
|
pub enum MelScale {
|
||||||
@@ -1205,7 +1241,7 @@ pub fn torch_stft(
|
|||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
// waveform: already padding
|
// waveform: already padding
|
||||||
// (bs, n_frames, n_fft)
|
// (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)?;
|
let result = frames.broadcast_mul(window)?;
|
||||||
// 傅立叶变换
|
// 傅立叶变换
|
||||||
@@ -1575,7 +1611,7 @@ pub fn spectrogram(
|
|||||||
)?;
|
)?;
|
||||||
frames = Tensor::cat(&[buffer_0, buffer_], D::Minus1)?;
|
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;
|
let pad_len = fft_length - frame_length;
|
||||||
if pad_len > 0 {
|
if pad_len > 0 {
|
||||||
// (bs, nframes, frame_length) -> (bs, nframes, fft_length)
|
// (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)
|
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 tensor_utils;
|
||||||
pub mod video_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 std::{collections::HashMap, fs, path::PathBuf, process::Command, time::Duration};
|
||||||
|
|
||||||
use aha_openai_dive::v1::resources::{
|
use aha_openai_dive::v1::resources::{
|
||||||
@@ -27,6 +28,7 @@ use dirs::home_dir;
|
|||||||
use half::{bf16, f16, slice::HalfFloatSliceExt};
|
use half::{bf16, f16, slice::HalfFloatSliceExt};
|
||||||
use modelscope::ModelScope;
|
use modelscope::ModelScope;
|
||||||
use tokio::time::sleep;
|
use tokio::time::sleep;
|
||||||
|
use zip::ZipArchive;
|
||||||
|
|
||||||
pub fn get_device(device: Option<&Device>) -> Device {
|
pub fn get_device(device: Option<&Device>) -> Device {
|
||||||
match device {
|
match device {
|
||||||
@@ -274,16 +276,14 @@ pub fn read_pth_tensor_info_cycle<P: AsRef<std::path::Path>>(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
current_obj
|
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 {
|
} else {
|
||||||
if let Object::Dict(key_values) = obj {
|
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
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
obj
|
obj
|
||||||
@@ -309,7 +309,7 @@ pub fn read_pth_tensor_info_cycle<P: AsRef<std::path::Path>>(
|
|||||||
let tensor_names = tensor_infos.keys();
|
let tensor_names = tensor_infos.keys();
|
||||||
let mut tensors = Vec::with_capacity(tensor_names.len());
|
let mut tensors = Vec::with_capacity(tensor_names.len());
|
||||||
for name in tensor_names {
|
for name in tensor_names {
|
||||||
let _ = match tensor_infos.get(name) {
|
match tensor_infos.get(name) {
|
||||||
None => {}
|
None => {}
|
||||||
Some(tensor_info) => {
|
Some(tensor_info) => {
|
||||||
let zip_reader = std::io::BufReader::new(std::fs::File::open(&path)?);
|
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();
|
let remaining = chars.as_str().to_lowercase();
|
||||||
format!("{}{}", first_char, remaining)
|
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
|
let mask = mask
|
||||||
.unsqueeze(D::Minus1)?
|
.unsqueeze(D::Minus1)?
|
||||||
.broadcast_as(hidden_states.shape())?;
|
.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)
|
Ok(hidden_states)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -431,19 +431,25 @@ pub fn compute_1d_coords(
|
|||||||
output_size: usize,
|
output_size: usize,
|
||||||
align_corner: Option<bool>,
|
align_corner: Option<bool>,
|
||||||
) -> Result<Vec<f32>> {
|
) -> 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 {
|
if input_size == 1 {
|
||||||
Ok(vec![0f32; output_size])
|
return Ok(vec![0f32; output_size]);
|
||||||
} else if let Some(align_) = align_corner
|
}
|
||||||
&& align_
|
let align_corners = align_corner.unwrap_or(false);
|
||||||
{
|
if align_corners {
|
||||||
Ok((0..output_size)
|
let scale = (input_size - 1) as f32 / (output_size - 1) as f32;
|
||||||
.map(|i| i as f32 * (input_size - 1) as f32 / (output_size - 1) as f32)
|
Ok((0..output_size).map(|i| i as f32 * scale).collect())
|
||||||
.collect())
|
|
||||||
} else {
|
} else {
|
||||||
|
let scale = input_size as f32 / output_size as f32;
|
||||||
Ok((0..output_size)
|
Ok((0..output_size)
|
||||||
.map(|i| {
|
.map(|i| {
|
||||||
(i as f32 + 0.5) * (input_size as f32 / output_size as f32) - 0.5
|
let coord = (i as f32 + 0.5) * scale - 0.5;
|
||||||
// coord.max(0.0).min((input_size - 1) as f32)
|
coord.clamp(0.0, (input_size - 1) as f32)
|
||||||
})
|
})
|
||||||
.collect())
|
.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"
|
"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 {
|
if orig_size == target_size {
|
||||||
return Ok(t.clone());
|
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 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 b in 0..bs {
|
||||||
for c in 0..channels {
|
for c in 0..channels {
|
||||||
let input_slice = t.i((b, c))?;
|
for (i, &coord) in coords.iter().enumerate() {
|
||||||
let mut out_i = Vec::new();
|
|
||||||
|
|
||||||
for &coord in coords.iter().take(target_size) {
|
|
||||||
// Nearest neighbor: round to nearest integer coordinate
|
// 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 clamped_idx = nearest_idx.min(orig_size - 1);
|
||||||
|
|
||||||
let value = input_slice.get(clamped_idx)?;
|
let value = input_data[b][c][clamped_idx];
|
||||||
out_i.push(value);
|
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)
|
Ok(output)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1068,3 +1067,14 @@ pub fn sequence_mask(length: &Tensor, max_length: Option<u32>) -> Result<Tensor>
|
|||||||
let mask = x.broadcast_lt(&length)?;
|
let mask = x.broadcast_lt(&length)?;
|
||||||
Ok(mask)
|
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::io::Cursor;
|
||||||
|
|
||||||
use std::time::Instant;
|
use std::fs::File;
|
||||||
|
|
||||||
use aha::utils::tensor_utils::interpolate_nearest_1d;
|
|
||||||
use anyhow::{Result, anyhow};
|
|
||||||
use candle_core::Tensor;
|
|
||||||
use sentencepiece::SentencePieceProcessor;
|
|
||||||
// use symphonia::core::io::MediaSourceStream;
|
// 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]
|
#[test]
|
||||||
fn messy_test() -> Result<()> {
|
fn messy_test() -> Result<()> {
|
||||||
// RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture
|
||||||
let device = &candle_core::Device::Cpu;
|
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"))?;
|
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 model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
|
||||||
let bpe_path = model_path.to_string() + "/bpe.model";
|
let emo_matrix_path = model_path.clone() + "/feat2.pt";
|
||||||
let tokenizer = SentencePieceProcessor::open(bpe_path)
|
let t_emo = load_tensor_from_pt(
|
||||||
.map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?;
|
&emo_matrix_path,
|
||||||
let tokens = tokenizer
|
"feat2/data/0",
|
||||||
.encode("你好啊")
|
Shape::from_dims(&[73, 1280]),
|
||||||
.map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?;
|
&device,
|
||||||
println!("tokens: {:?}", tokens);
|
)?;
|
||||||
|
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))?;
|
// let t = Tensor::arange(0.0f32, 40.0, device)?.broadcast_as((1, 40, 40))?;
|
||||||
// println!("t: {}", t);
|
// println!("t: {}", t);
|
||||||
// let i_start = Instant::now();
|
// let i_start = Instant::now();
|
||||||
|
|||||||
@@ -1,8 +1,11 @@
|
|||||||
use std::time::Instant;
|
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 aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
|
use anyhow::Result;
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn index_tts2_generate() -> Result<()> {
|
async fn index_tts2_generate() -> Result<()> {
|
||||||
@@ -24,10 +27,10 @@ async fn index_tts2_generate() -> Result<()> {
|
|||||||
{
|
{
|
||||||
"url": "file:///home/jhq/Videos/voice_01.wav"
|
"url": "file:///home/jhq/Videos/voice_01.wav"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"type": "text",
|
"type": "text",
|
||||||
"text": "你好啊"
|
"text": "你好啊,吃饭了吗"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
@@ -42,7 +45,11 @@ async fn index_tts2_generate() -> Result<()> {
|
|||||||
|
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let generate = voxcpm_generate.generate(mes)?;
|
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();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ fn qwen3_asr_generate() -> Result<()> {
|
|||||||
// RUST_BACKTRACE=1 cargo test -F cuda qwen3_asr_generate -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda qwen3_asr_generate -r -- --nocapture
|
||||||
let save_dir =
|
let save_dir =
|
||||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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#"
|
let message = r#"
|
||||||
{
|
{
|
||||||
"model": "qwen3-asr",
|
"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 aha::utils::{find_type_files, get_device, read_pth_tensor_info_cycle};
|
||||||
use anyhow::Result;
|
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;
|
use candle_nn::VarBuilder;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -203,19 +207,26 @@ fn index_tts2_weight() -> Result<()> {
|
|||||||
// RUST_BACKTRACE=1 cargo test -F cuda index_tts2_weight -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda index_tts2_weight -r -- --nocapture
|
||||||
let save_dir: String =
|
let save_dir: String =
|
||||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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 model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
|
||||||
let s2mel_path = model_path+ "/s2mel.pth";
|
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 wac2vec2_path = model_path+ "/wav2vec2bert_stats.pt";
|
||||||
// let model_path = format!("{}/iic/speech_campplus_sv_zh-cn_16k-common/", save_dir);
|
// let model_path = format!("{}/iic/speech_campplus_sv_zh-cn_16k-common/", save_dir);
|
||||||
// let campplus_path = model_path+ "/campplus_cn_common.bin";
|
// let campplus_path = model_path+ "/campplus_cn_common.bin";
|
||||||
// let model_list = find_type_files(&model_path, "safetensors")?;
|
// let model_list = find_type_files(&model_path, "safetensors")?;
|
||||||
let model_list = vec![s2mel_path];
|
let model_list = vec![bigvgan_path];
|
||||||
// let mut dict_to_hashmap = HashMap::new();
|
// // let mut dict_to_hashmap = HashMap::new();
|
||||||
// let mut dtype = candle_core::DType::F32;
|
// // let mut dtype = candle_core::DType::F32;
|
||||||
for m in model_list {
|
for m in model_list {
|
||||||
// let dict = read_all_with_key(m, Some("state_dict"))?;
|
// let dict = read_all_with_key(m, Some("state_dict"))?;
|
||||||
// let dict = read_all_with_key(m, Some("net"))?;
|
let dict = read_all_with_key(m, Some("generator"))?;
|
||||||
let dict = read_pth_tensor_info_cycle(m, Some("net.cfm"))?;
|
// let dict = read_pth_tensor_info_cycle(m, Some("net.cfm"))?;
|
||||||
// dtype = dict[0].1.dtype();
|
// dtype = dict[0].1.dtype();
|
||||||
for (k, v) in dict {
|
for (k, v) in dict {
|
||||||
// if k.contains("model") {
|
// if k.contains("model") {
|
||||||
@@ -230,9 +241,9 @@ fn index_tts2_weight() -> Result<()> {
|
|||||||
// let model_list = vec![semantic_codec_path];
|
// let model_list = vec![semantic_codec_path];
|
||||||
// for m in model_list {
|
// for m in model_list {
|
||||||
// let weights = safetensors::load(m, &device)?;
|
// let weights = safetensors::load(m, &device)?;
|
||||||
// for (key, tensor) in weights.iter() {
|
// for (key, tensor) in weights.iter() {
|
||||||
// println!("=== {} === {:?}", key, tensor.shape());
|
// println!("=== {} === {:?}", key, tensor.shape());
|
||||||
// }
|
// }
|
||||||
// }
|
// }
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user