updata index tts

This commit is contained in:
jhqxxx
2026-02-14 15:52:30 +08:00
parent 7c832e0ce8
commit 2c34fc2d79
40 changed files with 2602 additions and 540 deletions
+46
View File
@@ -0,0 +1,46 @@
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct BigVGANConfig {
pub resblock: String,
pub num_gpus: usize,
pub batch_size: usize,
pub learning_rate: f64,
pub adam_b1: f64,
pub adam_b2: f64,
pub lr_decay: f64,
pub seed: u32,
pub upsample_rates: Vec<usize>,
pub upsample_kernel_sizes: Vec<usize>,
pub upsample_initial_channel: usize,
pub resblock_kernel_sizes: Vec<usize>,
pub resblock_dilation_sizes: Vec<Vec<usize>>,
pub use_tanh_at_final: bool,
pub use_bias_at_final: bool,
pub activation: String,
pub snake_logscale: bool,
pub use_cqtd_instead_of_mrd: bool,
pub cqtd_filters: usize,
pub cqtd_max_filters: usize,
pub cqtd_filters_scale: usize,
pub cqtd_dilations: Vec<usize>,
pub cqtd_hop_lengths: Vec<usize>,
pub cqtd_n_octaves: Vec<usize>,
pub cqtd_bins_per_octaves: Vec<usize>,
pub mpd_reshapes: Vec<usize>,
pub use_spectral_norm: bool,
pub discriminator_channel_mult: usize,
pub use_multiscale_melloss: bool,
pub lambda_melloss: f64,
pub clip_grad_norm: f64,
pub segment_size: usize,
pub num_mels: usize,
pub num_freq: usize,
pub n_fft: usize,
pub hop_size: usize,
pub win_size: usize,
pub sampling_rate: usize,
pub fmin: usize,
pub fmax: Option<usize>,
pub fmax_for_loss: Option<usize>,
pub normalize_volume: bool,
pub num_workers: usize,
}
+333
View File
@@ -0,0 +1,333 @@
use anyhow::Result;
use candle_core::{D, Tensor};
use candle_nn::{Init, VarBuilder};
use crate::{
models::{
bigvgan::config::BigVGANConfig,
common::{WNConv1d, WNConvTranspose1d},
},
utils::tensor_utils::pad_replicate_last_dim,
};
pub mod config;
pub struct UpSample1d {
// ratio: usize,
stride: usize,
pad: usize,
pad_left: usize,
pad_right: usize,
filter: Tensor,
}
impl UpSample1d {
pub fn new(vb: VarBuilder, ratio: usize, kernel_size: Option<usize>) -> Result<Self> {
let stride = ratio;
let kernel_size = kernel_size.unwrap_or(6 * ratio / 2 * 2);
let pad = kernel_size / ratio - 1;
let pad_left = pad * stride + (kernel_size - stride) / 2;
let pad_right = pad * stride + (kernel_size - stride + 1) / 2;
let filter = vb.get_with_hints((1, 1, kernel_size), "filter", Init::Const(0.0))?;
Ok(Self {
// ratio,
stride,
pad,
pad_left,
pad_right,
filter,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let c = xs.dim(1)?;
let xs = pad_replicate_last_dim(xs, (self.pad, self.pad))?;
let xs = xs.conv_transpose1d(&self.filter.repeat((c, 1, 1))?, 0, 0, self.stride, 1, c)?;
let xs_last_dim = xs.dim(D::Minus1)?;
let xs_len = xs_last_dim - self.pad_left - self.pad_right;
let xs = xs.narrow(D::Minus1, self.pad_left, xs_len)?;
Ok(xs)
}
}
pub struct DownSample1d {
stride: usize,
// kernel_size: usize,
pad_left: usize,
pad_right: usize,
filter: Tensor,
}
impl DownSample1d {
pub fn new(vb: VarBuilder, ratio: usize, kernel_size: Option<usize>) -> Result<Self> {
let stride = ratio;
let kernel_size = kernel_size.unwrap_or(6 * ratio / 2 * 2);
let even = if kernel_size.is_multiple_of(2) { 1 } else { 0 };
let pad_left = kernel_size / 2 - even;
let pad_right = kernel_size / 2;
let filter = vb.get_with_hints((1, 1, kernel_size), "lowpass.filter", Init::Const(0.0))?;
Ok(Self {
stride,
// kernel_size,
pad_left,
pad_right,
filter,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let c = xs.dim(1)?;
let xs = pad_replicate_last_dim(xs, (self.pad_left, self.pad_right))?;
let xs = xs.conv1d(&self.filter.repeat((c, 1, 1))?, 0, self.stride, 1, c)?;
Ok(xs)
}
}
pub struct SnakeBeta {
alpha: Tensor,
beta: Tensor,
no_div_by_zero: f64,
}
impl SnakeBeta {
pub fn new(vb: VarBuilder, in_features: usize) -> Result<Self> {
let no_div_by_zero = 0.000000001f64;
let alpha = vb
.get_with_hints(in_features, "alpha", Init::Const(0.0))?
.unsqueeze(0)?
.unsqueeze(D::Minus1)?
.contiguous()?
.exp()?;
let beta = vb
.get_with_hints(in_features, "beta", Init::Const(0.0))?
.unsqueeze(0)?
.unsqueeze(D::Minus1)?
.contiguous()?
.exp()?;
Ok(Self {
alpha,
beta,
no_div_by_zero,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let beta = (1.0 / self.beta.affine(1.0, self.no_div_by_zero)?)?;
let xs = xs
.broadcast_mul(&self.alpha)?
.sin()?
.powf(2.0)?
.broadcast_mul(&beta)?
.add(xs)?;
Ok(xs)
}
}
pub struct TorchActivation1d {
upsample: UpSample1d,
downsample: DownSample1d,
act: SnakeBeta,
}
impl TorchActivation1d {
pub fn new(
vb: VarBuilder,
up_ratio: usize,
down_ratio: usize,
up_kernel_size: usize,
down_kernel_size: usize,
channels: usize,
) -> Result<Self> {
let upsample = UpSample1d::new(vb.pp("upsample"), up_ratio, Some(up_kernel_size))?;
let downsample =
DownSample1d::new(vb.pp("downsample"), down_ratio, Some(down_kernel_size))?;
let act = SnakeBeta::new(vb.pp("act"), channels)?;
Ok(Self {
upsample,
downsample,
act,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let xs = self.upsample.forward(xs)?;
let xs = self.act.forward(&xs)?;
let xs = self.downsample.forward(&xs)?;
Ok(xs)
}
}
pub struct AMPBlock1 {
convs1: Vec<WNConv1d>,
convs2: Vec<WNConv1d>,
activations: Vec<TorchActivation1d>,
}
impl AMPBlock1 {
pub fn new(
vb: VarBuilder,
channels: usize,
kernel_size: usize,
dilation: Vec<usize>,
) -> Result<Self> {
let vb_convs1 = vb.pp("convs1");
let mut convs1 = vec![];
for (i, &d) in dilation.iter().enumerate() {
let pad = ((kernel_size * d - d) as f32 / 2.0).round() as usize;
let layer = WNConv1d::new(
vb_convs1.pp(i),
channels,
channels,
kernel_size,
d,
pad,
1,
1,
true,
)?;
convs1.push(layer);
}
let vb_convs2 = vb.pp("convs2");
let mut convs2 = vec![];
for (i, _) in dilation.iter().enumerate() {
let pad = ((kernel_size - 1) as f32 / 2.0).round() as usize;
let layer = WNConv1d::new(
vb_convs2.pp(i),
channels,
channels,
kernel_size,
1,
pad,
1,
1,
true,
)?;
convs2.push(layer);
}
let num_layer = convs1.len() + convs2.len();
let act_vb = vb.pp("activations");
let mut activations = vec![];
for i in 0..num_layer {
let layer = TorchActivation1d::new(act_vb.pp(i), 2, 2, 12, 12, channels)?;
activations.push(layer);
}
Ok(Self {
convs1,
convs2,
activations,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let lens = self.convs1.len();
let mut xs = xs.clone();
for i in 0..lens {
let xt = self.activations[i * 2].forward(&xs)?;
let xt = self.convs1[i].forward(&xt)?;
let xt = self.activations[i * 2 + 1].forward(&xt)?;
let xt = self.convs2[i].forward(&xt)?;
xs = xs.add(&xt)?;
}
Ok(xs)
}
}
pub struct BigVGAN {
num_kernels: usize,
num_upsamples: usize,
conv_pre: WNConv1d,
ups: Vec<WNConvTranspose1d>,
resblocks: Vec<AMPBlock1>,
activation_post: TorchActivation1d,
conv_post: WNConv1d,
use_tanh_at_final: bool,
}
impl BigVGAN {
pub fn new(vb: VarBuilder, cfg: &BigVGANConfig) -> Result<Self> {
let num_kernels = cfg.resblock_kernel_sizes.len();
let num_upsamples = cfg.upsample_rates.len();
let conv_pre = WNConv1d::new(
vb.pp("conv_pre"),
cfg.num_mels,
cfg.upsample_initial_channel,
7,
1,
3,
1,
1,
true,
)?;
let vb_ups = vb.pp("ups");
let mut ups = vec![];
for (i, (&u, &k)) in cfg
.upsample_rates
.iter()
.zip(cfg.upsample_kernel_sizes.iter())
.enumerate()
{
let in_c = cfg.upsample_initial_channel / (2_i32.pow(i as u32) as usize);
let out_c = cfg.upsample_initial_channel / (2_i32.pow(i as u32 + 1) as usize);
let pad = (k - u) / 2;
let layer =
WNConvTranspose1d::new(vb_ups.pp(i).pp("0"), in_c, out_c, 1, k, pad, 0, 1, u)?;
ups.push(layer);
}
let vb_resblocks = vb.pp("resblocks");
let mut resblocks = vec![];
let ups_len = ups.len();
let res_len = cfg.resblock_kernel_sizes.len();
let mut ch = 0;
for i in 0..ups_len {
ch = cfg.upsample_initial_channel / (2_i32.pow(i as u32 + 1) as usize);
for (j, (&k, d)) in cfg
.resblock_kernel_sizes
.iter()
.zip(cfg.resblock_dilation_sizes.iter())
.enumerate()
{
let layer = AMPBlock1::new(vb_resblocks.pp(i * res_len + j), ch, k, d.clone())?;
resblocks.push(layer);
}
}
let activation_post = TorchActivation1d::new(vb.pp("activation_post"), 2, 2, 12, 12, ch)?;
let conv_post = WNConv1d::new(vb.pp("conv_post"), ch, 1, 7, 1, 3, 1, 1, false)?;
Ok(Self {
num_kernels,
num_upsamples,
conv_pre,
ups,
resblocks,
activation_post,
conv_post,
use_tanh_at_final: cfg.use_tanh_at_final,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let mut xs = self.conv_pre.forward(xs)?;
for i in 0..self.num_upsamples {
xs = self.ups[i].forward(&xs)?;
let mut xs_j = xs.zeros_like()?;
for j in 0..self.num_kernels {
let j_ = self.resblocks[i * self.num_kernels + j].forward(&xs)?;
xs_j = xs_j.add(&j_)?;
}
xs = xs_j.affine(1.0 / (self.num_kernels as f64), 0.0)?;
}
xs = self.activation_post.forward(&xs)?;
xs = self.conv_post.forward(&xs)?;
xs = if self.use_tanh_at_final {
xs.tanh()?
} else {
xs.clamp(-1.0, 1.0)?
};
Ok(xs)
}
}
+10 -6
View File
@@ -25,7 +25,11 @@ impl Shortcut {
) -> Result<Self> {
let conv_0 = get_conv2d(vb.pp("0"), in_c, out_c, ks, padding, 1, 1, 1, bias)?;
let bn_1 = get_batch_norm(vb.pp("1"), 1e-5, out_c, true)?;
Ok(Self { conv_0, bn_1, stride })
Ok(Self {
conv_0,
bn_1,
stride,
})
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
@@ -36,7 +40,7 @@ impl Shortcut {
let indices = Tensor::arange(0u32, half_h as u32, x.device())?.affine(2.0, 0.0)?;
x = x.index_select(&indices, 2)?;
}
x = self.bn_1.forward_t(&x, false)?;
x = self.bn_1.forward_t(&x, false)?;
Ok(x)
}
}
@@ -106,7 +110,7 @@ impl BasicResBlock {
} else {
xs = xs.add(&residual)?;
}
xs = xs.relu()?;
xs = xs.relu()?;
Ok(xs)
}
}
@@ -435,7 +439,7 @@ impl DenseLayer {
.forward(&xs.unsqueeze(D::Minus1)?)?
.squeeze(D::Minus1)?
} else {
self.linear.forward(&xs)?
self.linear.forward(xs)?
};
let xs = self.nonlinear.forward_t(&xs, false)?;
Ok(xs)
@@ -463,7 +467,7 @@ impl XVector {
let mut channels = init_channels;
let mut blocks = vec![];
let mut transits = vec![];
let params = vec![(12, 3, 1), (24, 3, 2), (16, 3, 2)];
let params = [(12, 3, 1), (24, 3, 2), (16, 3, 2)];
for (i, (num_layers, ks, dilation)) in params.iter().enumerate() {
let block = CAMDenseTDNNBlock::new(
vb.pp(format!("block{}", i + 1)),
@@ -477,7 +481,7 @@ impl XVector {
false,
)?;
blocks.push(block);
channels = channels + num_layers * growth_rate;
channels += num_layers * growth_rate;
let transit = TransitLayer::new(
vb.pp(format!("transit{}", i + 1)),
channels,
+257 -14
View File
@@ -2,8 +2,8 @@ use anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig,
ConvTranspose1d, ConvTranspose1dConfig, Embedding, GroupNorm, Init, LayerNorm, LayerNormConfig,
Linear, Module, ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d,
ConvTranspose1d, ConvTranspose1dConfig, Embedding, Init, LayerNorm, LayerNormConfig, Linear,
Module, ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d,
conv2d_no_bias, embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm,
};
use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator};
@@ -700,11 +700,7 @@ pub fn get_layer_norm(vb: VarBuilder, eps: f64, dim: usize, affine: bool) -> Res
Ok(norm)
}
pub fn get_layer_norm_without_weight(
vb: VarBuilder,
eps: f64,
dim: usize,
) -> Result<LayerNorm> {
pub fn get_layer_norm_without_weight(vb: VarBuilder, eps: f64, dim: usize) -> Result<LayerNorm> {
let weight = Tensor::ones(dim, vb.dtype(), vb.device())?;
let bias = Tensor::zeros(dim, vb.dtype(), vb.device())?;
Ok(LayerNorm::new(weight, bias, eps))
@@ -1027,11 +1023,25 @@ impl GLU {
Ok(Self { dim })
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let half_dim = xs.dim(self.dim)? / 2;
let a = xs.narrow(self.dim, 0, half_dim)?;
let b = xs.narrow(self.dim, half_dim, half_dim)?;
let b = sigmoid(&b)?;
let xs = a.mul(&b)?;
let x_ = xs.chunk(2, self.dim)?;
let x_1 = sigmoid(x_[1].as_ref())?;
let xs = x_1.mul(x_[0].as_ref())?;
Ok(xs)
}
}
pub struct GEGLU {
dim: usize,
}
impl GEGLU {
pub fn new(dim: usize) -> Result<Self> {
Ok(Self { dim })
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let x_ = xs.chunk(2, self.dim)?;
let x_1 = x_[1].as_ref().gelu()?;
let xs = x_1.mul(x_[0].as_ref())?;
Ok(xs)
}
}
@@ -1103,8 +1113,8 @@ impl WNConvTranspose1d {
let normalized_weight = weight_v.broadcast_div(&weight_norm)?;
let scaled_weight = normalized_weight.broadcast_mul(&weight_g)?;
let config = ConvTranspose1dConfig {
padding: padding,
output_padding: output_padding,
padding,
output_padding,
stride,
dilation,
groups,
@@ -1184,3 +1194,236 @@ pub fn mish(xs: &Tensor) -> Result<Tensor> {
let xs = xs.mul(&tanh)?;
Ok(xs)
}
pub struct GPT2Attention {
num_heads: usize,
head_dim: usize,
c_attn: Linear,
c_proj: Linear,
kv_cache: Option<(Tensor, Tensor)>,
}
impl GPT2Attention {
pub fn new(vb: VarBuilder, hidden_size: usize, num_heads: usize) -> Result<Self> {
let c_attn_weight = vb
.get_with_hints(
(hidden_size, 3 * hidden_size),
"c_attn.weight",
Init::Const(1.0),
)?
.t()?;
let c_attn_bias = vb.get_with_hints(3 * hidden_size, "c_attn.bias", Init::Const(0.0))?;
let c_attn = Linear::new(c_attn_weight, Some(c_attn_bias));
// let c_attn = linear_b(3 * hidden_size, hidden_size, true, vb.pp("c_attn"))?;
let c_proj_weight = vb
.get_with_hints(
(hidden_size, hidden_size),
"c_proj.weight",
Init::Const(1.0),
)?
.t()?;
let c_proj_bias = vb.get_with_hints(hidden_size, "c_proj.bias", Init::Const(0.0))?;
let c_proj = Linear::new(c_proj_weight, Some(c_proj_bias));
// let c_proj = linear_b(hidden_size, hidden_size, true, vb.pp("c_proj"))?;
let head_dim = hidden_size / num_heads;
Ok(Self {
num_heads,
head_dim,
c_attn,
c_proj,
kv_cache: None,
})
}
pub fn forward(&mut self, xs: &Tensor, attention_mask: Option<&Tensor>) -> Result<Tensor> {
let (b, seq_len, _) = xs.dims3()?;
let xs = self.c_attn.forward(xs)?;
let xs_splits = xs.chunk(3, 2)?;
let query_states = xs_splits[0]
.as_ref()
.reshape((b, seq_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let key_states = xs_splits[1]
.as_ref()
.reshape((b, seq_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let value_states = xs_splits[2]
.as_ref()
.reshape((b, seq_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let (key_states, value_states) = match &self.kv_cache {
None => (key_states, value_states),
Some((prev_k, prev_v)) => {
let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
(key_states, value_states)
}
};
self.kv_cache = Some((key_states.clone(), value_states.clone()));
let scale = 1f64 / f64::sqrt(self.head_dim as f64);
let attn_output = eager_attention_forward(
&query_states,
&key_states,
&value_states,
None,
attention_mask,
scale,
)?;
let attn_output = attn_output.reshape((b, seq_len, self.num_heads * self.head_dim))?;
let attn_output = attn_output.apply(&self.c_proj)?;
Ok(attn_output)
}
pub fn clear_kv_cache(&mut self) {
self.kv_cache = None
}
}
pub struct GPT2MLP {
linear1: Linear,
linear2: Linear,
act: Activation,
}
impl GPT2MLP {
pub fn new(
vb: VarBuilder,
in_dim: usize,
middle_dim: usize,
out_dim: usize,
act: Activation,
) -> Result<Self> {
let c_fc_weight = vb
.get_with_hints((in_dim, middle_dim), "c_fc.weight", Init::Const(1.0))?
.t()?;
let c_fc_bias = vb.get_with_hints(middle_dim, "c_fc.bias", Init::Const(0.0))?;
let c_fc = Linear::new(c_fc_weight, Some(c_fc_bias));
let c_proj_weight = vb
.get_with_hints((middle_dim, out_dim), "c_proj.weight", Init::Const(1.0))?
.t()?;
let c_proj_bias = vb.get_with_hints(out_dim, "c_proj.bias", Init::Const(0.0))?;
let c_proj = Linear::new(c_proj_weight, Some(c_proj_bias));
Ok(Self {
linear1: c_fc,
linear2: c_proj,
act,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let xs = xs
.apply(&self.linear1)?
.apply(&self.act)?
.apply(&self.linear2)?;
Ok(xs)
}
}
pub struct GPT2Block {
ln_1: LayerNorm,
attn: GPT2Attention,
ln_2: LayerNorm,
mlp: GPT2MLP,
}
impl GPT2Block {
pub fn new(
vb: VarBuilder,
hidden_size: usize,
num_heads: usize,
inner_dim: Option<usize>,
) -> Result<Self> {
let inner_dim = inner_dim.unwrap_or(4 * hidden_size);
let ln_1 = get_layer_norm(vb.pp("ln_1"), 1e-5, hidden_size, true)?;
let attn = GPT2Attention::new(vb.pp("attn"), hidden_size, num_heads)?;
let ln_2 = get_layer_norm(vb.pp("ln_2"), 1e-5, hidden_size, true)?;
let mlp = GPT2MLP::new(
vb.pp("mlp"),
hidden_size,
inner_dim,
hidden_size,
Activation::NewGelu,
)?;
Ok(Self {
ln_1,
attn,
ln_2,
mlp,
})
}
pub fn forward(&mut self, xs: &Tensor, attention_mask: Option<&Tensor>) -> Result<Tensor> {
let residual = xs.clone();
let xs = self.ln_1.forward(xs)?;
let xs = self.attn.forward(&xs, attention_mask)?;
let residual = xs.add(&residual)?;
let xs = self.ln_2.forward(&residual)?;
let xs = self.mlp.forward(&xs)?;
let xs = xs.add(&residual)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
self.attn.clear_kv_cache()
}
}
#[allow(unused)]
pub struct GPT2Model {
wte: Embedding,
// wpe: Embedding,
h: Vec<GPT2Block>,
ln_f: LayerNorm,
}
#[allow(unused)]
impl GPT2Model {
pub fn new(
vb: VarBuilder,
hidden_size: usize,
num_heads: usize,
num_hidden_layers: usize,
wte_embeddings: &Tensor,
) -> Result<Self> {
// let wte = embedding(vocab_size, hidden_size, vb.pp("wte"))?;
let wte = Embedding::new(wte_embeddings.clone(), hidden_size);
// let wpe = embedding(max_position_embeddings, hidden_size, vb.pp("wpe"))?;
let vb_layers = vb.pp("h");
let mut h = vec![];
for i in 0..num_hidden_layers {
let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?;
h.push(block);
}
let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?;
Ok(Self { wte, h, ln_f })
}
pub fn forward(&mut self, inputs_embeds: &Tensor) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
b_size,
seq_len,
0,
xs.device(),
)?)
}
};
for block in &mut self.h {
xs = block.forward(&xs, attention_mask.as_ref())?;
}
xs = self.ln_f.forward(&xs)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.h.iter_mut() {
layer.clear_kv_cache()
}
}
}
+1 -1
View File
@@ -18,4 +18,4 @@ pub struct FeatureExtractor {
fn default_sampling_rate() -> usize {
16000
}
}
+2 -2
View File
@@ -1,3 +1,3 @@
pub mod seamless_m4t_feature_extractor;
pub mod config;
pub mod feature_extraction_whisper;
pub mod config;
pub mod seamless_m4t_feature_extractor;
@@ -82,7 +82,7 @@ impl SeamlessM4TFeatureExtractor {
0.97,
Some(&self.mel_filters),
Some("log"),
1.192092955078125e-07,
1.192_092_9e-7,
true,
)?
.transpose(D::Minus1, D::Minus2)?;
-2
View File
@@ -11,8 +11,6 @@ pub struct GlmAsrNanoProcessorConfig {
pub max_audio_len: usize,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct GlmAsrNanoConfig {
pub audio_config: GlmAsrAudioConfig,
+1 -1
View File
@@ -1,4 +1,3 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor};
@@ -27,6 +26,7 @@ pub struct GlmAsrNanoProcessor {
whisper_feature_extrator: WhisperFeatureExtractor,
}
#[allow(unused)]
impl GlmAsrNanoProcessor {
pub fn new(path: &str, device: &Device, dtype: DType) -> Result<Self> {
let path = path.to_string();
+3 -3
View File
@@ -5,7 +5,7 @@ use crate::models::mask_gct::config::SemanticCodec;
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct IndexTTS2Config {
pub dataset: Dataset,
pub gpt: Gpt,
pub gpt: GptConfig,
pub semantic_codec: SemanticCodec,
pub s2mel: S2MelConfig,
pub gpt_checkpoint: String,
@@ -39,7 +39,7 @@ pub struct Mel {
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct Gpt {
pub struct GptConfig {
pub model_dim: usize,
pub max_mel_tokens: usize,
pub max_text_tokens: usize,
@@ -206,7 +206,7 @@ impl DiTModelArgs {
has_cross_attention: false,
context_dim: 0,
uvit_skip_connection: config.uvit_skip_connection,
time_as_token: config.uvit_skip_connection,
time_as_token: config.time_as_token,
}
}
}
+20 -39
View File
@@ -1,78 +1,59 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use aha_openai_dive::v1::resources::chat::{ChatCompletionParameters, ChatCompletionResponse};
use anyhow::{Result, anyhow};
use base64::{Engine, prelude::BASE64_STANDARD};
use candle_core::{DType, Device};
use sentencepiece::SentencePieceProcessor;
use crate::{
models::index_tts2::{
config::IndexTTS2Config,
model::IndexTTS2Model,
processor::IndexTTS2Processor,
utils::{TextNormalizer, tokenize_by_cjk_char},
config::IndexTTS2Config, model::IndexTTS2Model, utils::tokenize_by_cjk_char,
},
tokenizer::sentencepiece_encode,
utils::{
audio_utils::extract_audio_url, extract_user_text, get_default_save_dir, get_device,
get_dtype,
audio_utils::get_audio_wav_u8, build_audio_completion_response, extract_user_text,
get_default_save_dir, get_device,
},
};
pub struct IndexTTS2Generate {
processor: IndexTTS2Processor,
tokenizer: SentencePieceProcessor,
config: IndexTTS2Config,
cache_spk_audio_prompt: Option<String>,
// config: IndexTTS2Config,
model: IndexTTS2Model,
device: Device,
sample_rate: u32,
model_name: String,
}
#[allow(unused)]
impl IndexTTS2Generate {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let config_path = path.to_string() + "/config.yaml";
let save_dir = get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let config: IndexTTS2Config = serde_yaml::from_slice(&std::fs::read(config_path)?)?;
let device = get_device(device);
let dtype = get_dtype(dtype, "bf16");
let processor = IndexTTS2Processor::new(&device)?;
let bpe_path = path.to_string() + "/bpe.model";
let tokenizer = SentencePieceProcessor::open(bpe_path)
.map_err(|e| anyhow!(format!("load bpe,model file error:{}", e)))?;
let model = IndexTTS2Model::new(path, &save_dir, &config, &device, dtype)?;
let model = IndexTTS2Model::new(path, &save_dir, &config, &device)?;
Ok(Self {
processor,
tokenizer,
config,
cache_spk_audio_prompt: None,
// config,
model,
device,
sample_rate: 22050,
model_name: "index-tts2".to_string(),
})
}
pub fn use_prompt(&self, mes: &ChatCompletionParameters) -> bool {
if let Some(cache) = &self.cache_spk_audio_prompt {
let audio_vec = extract_audio_url(mes);
if audio_vec.len() == 0 {
true
} else {
if cache.eq(&audio_vec[0]) { true } else { false }
}
} else {
false
}
}
pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<()> {
pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let text = extract_user_text(&mes)?;
let text = tokenize_by_cjk_char(&text, true);
let input_ids = sentencepiece_encode(&text, &self.tokenizer, &self.device)?;
let (audio_22k, audio_16k) = if self.use_prompt(&mes) {
(None, None)
} else {
let (audio_22k, audio_16k, prompt) = self.processor.process_info(&mes)?;
self.cache_spk_audio_prompt = Some(prompt);
(Some(audio_22k), Some(audio_16k))
};
let _ = self.model.forward(&input_ids, audio_22k.as_ref(), audio_16k.as_ref())?;
Ok(())
let audio = self.model.forward(&input_ids, &mes)?;
let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate)?;
let base64_audio = BASE64_STANDARD.encode(wav_u8);
let response = build_audio_completion_response(&base64_audio, &self.model_name);
Ok(response)
}
}
+2 -2
View File
@@ -1,5 +1,5 @@
pub mod config;
pub mod generate;
pub mod model;
pub mod processor;
pub mod utils;
// pub mod processor;
pub mod utils;
File diff suppressed because it is too large Load Diff
-203
View File
@@ -1,203 +0,0 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
use crate::{
models::{
campplus::CAMPPlus,
feature_extractor::seamless_m4t_feature_extractor::SeamlessM4TFeatureExtractor,
index_tts2::config::{IndexTTS2Config, PreprocessParams},
mask_gct::model::RepCodec,
w2v_bert_2_0::model::W2VBert2_0Model,
},
utils::{
audio_utils::{
create_hann_window, extract_audio_url, get_waveform_and_window_properties, kaldi_fbank,
kaldi_get_mel_banks, load_audio, mel_filter_bank, resample_simple, torch_stft,
},
get_vb_model_path,
tensor_utils::pad_reflect_last_dim,
},
};
pub struct IndexTTS2Processor {
device: Device,
max_audio_length_seconds: usize,
// feature_extractor: SeamlessM4TFeatureExtractor,
// semantic_model: W2VBert2_0Model,
// semantic_mean: Tensor,
// semantic_std: Tensor,
// semantic_codec: RepCodec,
// s2mel_filters: Tensor,
// s2mel_windows: Tensor,
// s2mel_preprocess_params: PreprocessParams,
// window_shift: usize,
// window_size: usize,
// padded_window_size: usize,
// mel_energies: Tensor,
// campplus_model: CAMPPlus,
}
impl IndexTTS2Processor {
pub fn new(
// path: &str,
// save_dir: &str,
// config: &IndexTTS2Config,
device: &Device,
// dtype: DType,
) -> Result<Self> {
// let feature_extractor = SeamlessM4TFeatureExtractor::new(
// // 80,
// 80,
// crate::utils::tensor_utils::PaddingSide::Right,
// 1.0,
// 16000,
// 2,
// device,
// )?;
// let w2vbert2_path = save_dir.to_string() + "/facebook/w2v-bert-2.0";
// let semantic_model = W2VBert2_0Model::init(&w2vbert2_path, device, dtype)?;
// let semantic_mean_var_path = path.to_string() + "/" + &config.w2v_stat;
// let dict = read_all_with_key(semantic_mean_var_path, None)?;
// let mut semantic_mean = Tensor::new(0.0, device)?.to_dtype(dtype)?;
// let mut semantic_std = Tensor::new(1.0, device)?.to_dtype(dtype)?;
// for (k, v) in dict {
// if k.eq("mean") {
// semantic_mean = v.to_device(device)?.to_dtype(dtype)?;
// } else if k.eq("var") {
// semantic_std = v.to_device(device)?.to_dtype(dtype)?.sqrt()?;
// }
// }
// let semantic_codec_path =
// save_dir.to_string() + "/amphion/MaskGCT/semantic_codec/model.safetensors";
// let vb =
// unsafe { VarBuilder::from_mmaped_safetensors(&[semantic_codec_path], dtype, &device)? };
// let semantic_codec = RepCodec::new(vb, &config.semantic_codec)?;
// let s2mel_filters = mel_filter_bank(
// config.s2mel.preprocess_params.spect_params.n_fft / 2 + 1,
// config.s2mel.preprocess_params.spect_params.n_mels,
// config.s2mel.preprocess_params.spect_params.fmin as f32,
// config
// .s2mel
// .preprocess_params
// .spect_params
// .fmax
// .unwrap_or(config.s2mel.preprocess_params.sr / 2) as f32,
// config.s2mel.preprocess_params.sr as f32,
// Some("slaney"),
// crate::utils::audio_utils::MelScale::Slaney,
// false,
// device,
// )?
// .t()?;
// let s2mel_windows = create_hann_window(
// config.s2mel.preprocess_params.spect_params.win_length,
// dtype,
// device,
// )?;
// let (window_shift, window_size, padded_window_size) =
// get_waveform_and_window_properties(16000, 10.0, 25.0, true)?;
// let (mel_energies, _) =
// kaldi_get_mel_banks(80, padded_window_size, 16000 as f32, 20.0, 0.0, device)?;
// let mel_energies = mel_energies.pad_with_zeros(D::Minus1, 0, 1)?.t()?;
// let campplus_model_path = save_dir.to_string()
// + "/iic/speech_campplus_sv_zh-cn_16k-common/campplus_cn_common.bin";
// let campplus_vb = get_vb_model_path(campplus_model_path, dtype, device.clone(), None)?;
// let campplus_model = CAMPPlus::new(campplus_vb, 80, 192, 32, 4, 128)?;
Ok(Self {
device: device.clone(),
max_audio_length_seconds: 15,
// feature_extractor,
// semantic_model,
// semantic_mean,
// semantic_std,
// semantic_codec,
// s2mel_filters,
// s2mel_windows,
// s2mel_preprocess_params: config.s2mel.preprocess_params.clone(),
// window_shift,
// window_size,
// padded_window_size,
// mel_energies,
// campplus_model,
})
}
pub fn cut_audio(&self, audio: &Tensor, sr: usize) -> Result<(Tensor, usize)> {
let max_audio_samples = self.max_audio_length_seconds * sr;
let audio_lens = audio.dim(1)?;
let audio = if audio_lens > max_audio_samples {
audio.i((.., 0..max_audio_samples))?
} else {
audio.clone()
};
Ok((audio, sr))
}
pub fn extract_audio_and_cut(
&self,
mes: &ChatCompletionParameters,
device: &Device,
) -> Result<(Tensor, usize, String)> {
let audio_url_vec = extract_audio_url(mes);
let (audio, sr) = load_audio(&audio_url_vec[0], device)?;
let (audio, sr) = self.cut_audio(&audio, sr)?;
Ok((audio, sr, audio_url_vec[0].clone()))
}
// pub fn get_emb(
// &self,
// input_features: &Tensor,
// attention_mask: Option<&Tensor>,
// ) -> Result<Tensor> {
// let output =
// self.semantic_model
// .forward(input_features, attention_mask, Some(17), false)?;
// let feature = &output.specify_layer_id_hidden_state.unwrap();
// let feature = feature
// .broadcast_sub(&self.semantic_mean)?
// .broadcast_div(&self.semantic_std)?;
// Ok(feature)
// }
// pub fn s2mel_spectrogram(&self, waveform: &Tensor) -> Result<Tensor> {
// let pad = (self.s2mel_preprocess_params.spect_params.n_fft
// - self.s2mel_preprocess_params.spect_params.hop_length)
// / 2;
// let pad_audio_22k = pad_reflect_last_dim(&waveform, (pad, pad))?;
// let spec = torch_stft(
// &pad_audio_22k,
// self.s2mel_preprocess_params.spect_params.n_fft,
// self.s2mel_preprocess_params.spect_params.hop_length,
// &self.s2mel_windows,
// )?
// .transpose(1, 2)?;
// let spec = self.s2mel_filters.broadcast_matmul(&spec)?;
// let spec = spec.clamp(1e-5, f64::INFINITY)?.log()?;
// Ok(spec)
// }
pub fn process_info(&self, mes: &ChatCompletionParameters) -> Result<(Tensor, Tensor, String)> {
let (audio, sr, audio_url) = self.extract_audio_and_cut(mes, &self.device)?;
let audio_22k = resample_simple(&audio, sr as i64, 22050)?;
let audio_16k = resample_simple(&audio, sr as i64, 16000)?;
// let (audio_16k_features, audio_16k_mask) =
// self.feature_extractor.call(&audio_16k, 16000, true, true)?;
// let spk_cond_emb = self.get_emb(&audio_16k_features, audio_16k_mask.as_ref())?;
// let (_, s_ref) = self.semantic_codec.quantize(&spk_cond_emb)?;
// let ref_mel = self.s2mel_spectrogram(&audio_22k)?;
// let feat = kaldi_fbank(
// &audio_16k,
// &self.mel_energies,
// self.window_shift,
// self.window_size,
// self.padded_window_size,
// 0.0,
// )?;
// let feat = feat.broadcast_sub(&feat.mean_keepdim(1)?)?;
// let style = self.campplus_model.forward(&feat)?;
// println!("style: {}", style);
Ok((audio_22k, audio_16k, audio_url))
}
}
+3 -2
View File
@@ -1,5 +1,3 @@
use std::collections::HashMap;
use anyhow::{Result, anyhow};
use regex::Regex;
@@ -15,9 +13,12 @@ pub async fn download_index_tts2_need_model(save_dir: Option<&str>) -> Result<()
let mask_gct = "amphion/MaskGCT";
// let campplus= "funasr/campplus"; // huggingface
let campplus = "iic/speech_campplus_sv_zh-cn_16k-common"; // modelscope
// let bigvgan = "nvidia/bigvgan_v2_22khz_80band_256x"; // huggingface
let bigvgan = "nv-community/bigvgan_v2_22khz_80band_256x"; // modelscope
download_model(w2v_bert2_0, &save_dir, 3).await?;
download_model(mask_gct, &save_dir, 3).await?;
download_model(campplus, &save_dir, 3).await?;
download_model(bigvgan, &save_dir, 3).await?;
Ok(())
}
+1 -1
View File
@@ -18,4 +18,4 @@ fn default_num_quantizers() -> usize {
fn default_downsample_scale() -> usize {
1
}
}
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod config;
pub mod model;
pub mod config;
+42 -7
View File
@@ -48,7 +48,7 @@ impl ConvNeXtBlock {
let xs = self.pwconv1.forward(&xs)?.gelu()?;
let mut xs = self.pwconv2.forward(&xs)?;
if let Some(gamma) = &self.gamma {
xs = xs.broadcast_mul(&gamma)?;
xs = xs.broadcast_mul(gamma)?;
}
let xs = xs.transpose(1, 2)?;
let xs = residual.add(&xs)?;
@@ -116,10 +116,28 @@ impl FactorizedVectorQuantize {
use_l2_normlize: bool,
) -> Result<Self> {
let (in_project, out_project) = if input_dim != codebook_dim {
let in_project =
WNConv1d::new(vb.pp("in_project"), input_dim, codebook_dim, 1, 1, 0, 1, 1, true)?;
let out_project =
WNConv1d::new(vb.pp("out_project"), codebook_dim, input_dim, 1, 1, 0, 1, 1, true)?;
let in_project = WNConv1d::new(
vb.pp("in_project"),
input_dim,
codebook_dim,
1,
1,
0,
1,
1,
true,
)?;
let out_project = WNConv1d::new(
vb.pp("out_project"),
codebook_dim,
input_dim,
1,
1,
0,
1,
1,
true,
)?;
(Some(in_project), Some(out_project))
} else {
(None, None)
@@ -166,6 +184,14 @@ impl FactorizedVectorQuantize {
}
Ok((z_q, indices))
}
pub fn vq2emb(&self, xs: &Tensor) -> Result<Tensor> {
let mut emb = self.codebook.forward(xs)?.transpose(1, 2)?;
if let Some(out_proj) = &self.out_project {
emb = out_proj.forward(&emb)?;
}
Ok(emb)
}
}
pub struct ResidualVQ {
@@ -210,7 +236,7 @@ impl ResidualVQ {
let n_quantizers = n_quantizers.unwrap_or(self.num_quantizers);
let mut residual = xs.clone();
let mut quantized_out = Tensor::new(0.0f32, xs.device())?.to_dtype(xs.dtype())?;
for (i, quantizer) in (&self.quantizers).iter().enumerate() {
for (i, quantizer) in self.quantizers.iter().enumerate() {
if i >= n_quantizers {
break;
}
@@ -224,8 +250,17 @@ impl ResidualVQ {
let all_quantized = Tensor::stack(&all_quantized, 0)?;
Ok((quantized_out, all_indices, all_quantized))
}
pub fn vq2emb(&self, xs: &Tensor) -> Result<Tensor> {
let mut quantized_out = xs.clone();
for quantizer in &self.quantizers {
quantized_out = quantizer.vq2emb(xs)?;
}
Ok(quantized_out)
}
}
#[allow(unused)]
pub struct RepCodec {
downsample_scale: usize,
down: Option<Conv1d>,
@@ -234,7 +269,7 @@ pub struct RepCodec {
encoder_1: Linear,
decoder_0: VocosBackbone,
decoder_1: Linear,
quantizer: ResidualVQ,
pub quantizer: ResidualVQ,
}
impl RepCodec {
+1
View File
@@ -1,3 +1,4 @@
pub mod bigvgan;
pub mod campplus;
pub mod common;
pub mod deepseek_ocr;
+4 -4
View File
@@ -167,13 +167,13 @@ pub struct Qwen3ASRTextConfig {
pub fn qwen3asr_text_config2qwen3_config(cfg: &Qwen3ASRTextConfig) -> Qwen3Config {
Qwen3Config {
attention_bias: cfg.attention_bias,
attention_dropout: cfg.attention_dropout as f64,
bos_token_id: cfg.bos_token_id.unwrap_or(151643) as u32,
eos_token_id: cfg.eos_token_id.unwrap_or(151645) as u32,
attention_dropout: cfg.attention_dropout,
bos_token_id: cfg.bos_token_id.unwrap_or(151643),
eos_token_id: cfg.eos_token_id.unwrap_or(151645),
head_dim: cfg.head_dim,
hidden_act: cfg.hidden_act,
hidden_size: cfg.hidden_size,
initializer_range: cfg.initializer_range as f64,
initializer_range: cfg.initializer_range,
intermediate_size: cfg.intermediate_size,
max_position_embeddings: cfg.max_position_embeddings,
max_window_layers: 0,
-1
View File
@@ -1,4 +1,3 @@
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
+1 -1
View File
@@ -338,7 +338,7 @@ impl Qwen3ASRThinker {
) -> Result<Tensor> {
let mut input_embeds = self.model.embed_tokens.forward(input_ids)?;
if let Some(input_features) = input_features {
let audio_feature = self.audio_tower.forward(&input_features)?;
let audio_feature = self.audio_tower.forward(input_features)?;
// println!("audio_feature: {}", audio_feature);
let audio_mask = get_equal_mask(input_ids, self.audio_token_id)?;
let n_audio_tokens = audio_mask.sum_all()?.to_scalar::<u32>()?;
+5 -10
View File
@@ -82,10 +82,7 @@ impl Qwen3AsrProcessor {
pub fn process_audio(&self, mes: &ChatCompletionParameters) -> Result<Vec<Tensor>> {
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?;
audio_tensors
.iter()
.map(|audio| float_range_normalize(&audio))
.collect()
audio_tensors.iter().map(float_range_normalize).collect()
}
pub fn validate_language(&self, lang: &String) -> bool {
@@ -93,10 +90,9 @@ impl Qwen3AsrProcessor {
}
fn replace_special_tokens(&self, text: &str, token_len: usize) -> String {
let replace = "<|audio_placeholder|>".repeat(token_len as usize);
let replace = "<|audio_placeholder|>".repeat(token_len);
let text = text.replacen(&self.audio_token, &replace, 1);
let text = text.replace("<|audio_placeholder|>", &self.audio_token);
text
text.replace("<|audio_placeholder|>", &self.audio_token)
}
pub fn process_info(
@@ -162,11 +158,10 @@ pub struct AudioData {
pub fn get_feat_extract_output_lengths(audio_len: usize) -> usize {
let input_len_leave = audio_len % 100;
let output_len = if input_len_leave > 0 {
if input_len_leave > 0 {
let feat_lengths = (input_len_leave - 1) / 2 + 1;
((feat_lengths - 1) / 2 + 1 - 1) / 2 + 1 + (audio_len / 100) * 13
} else {
(audio_len / 100) * 13
};
output_len
}
}
+1 -1
View File
@@ -58,4 +58,4 @@ pub struct W2VBert2_0Config {
pub use_weighted_layer_sum: bool,
pub vocab_size: Option<usize>,
pub xvector_output_dim: usize,
}
}
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod config;
pub mod model;
pub mod model;
+21 -24
View File
@@ -45,6 +45,7 @@ impl Wav2Vec2BertFeatureProjection {
}
}
#[allow(unused)]
pub struct Wav2Vec2BertSelfAttention {
q_proj: Linear,
k_proj: Linear,
@@ -203,17 +204,12 @@ impl Wav2Vec2BertSelfAttention {
.affine(scale, 0.0)?;
if let Some(mask) = attention_mask {
// let mask = mask.unsqueeze(1)?.unsqueeze(D::Minus1)?;
Some(relative_position_attn_weights.broadcast_add(&mask)?)
Some(relative_position_attn_weights.broadcast_add(mask)?)
} else {
Some(relative_position_attn_weights)
}
} else {
if let Some(mask) = attention_mask {
// let mask = mask.unsqueeze(1)?.unsqueeze(D::Minus1)?;
Some(mask.clone())
} else {
None
}
attention_mask.cloned()
};
let attn_output = eager_attention_forward(
@@ -476,7 +472,8 @@ impl Wav2Vec2BertEncoder {
let attention_mask = attention_mask
.where_cond(&attention_mask_f, &neg_inf_t)?
.to_dtype(xs.dtype())?
.affine(1.0, -1.0)?;
.affine(1.0, -1.0)?
.repeat((1, 1, seq_len, 1))?;
(xs, Some(attention_mask))
} else {
(xs.clone(), None)
@@ -490,7 +487,7 @@ impl Wav2Vec2BertEncoder {
let mut hidden_states: Vec<Tensor> = vec![];
let mut specify_layer_id_hidden_state = None;
for (i, layer) in (&self.layers).iter().enumerate() {
for (i, layer) in self.layers.iter().enumerate() {
if output_hidden_states {
hidden_states.push(xs.clone());
}
@@ -507,7 +504,7 @@ impl Wav2Vec2BertEncoder {
conv_attention_mask,
)?;
}
let hidden_states = if hidden_states.len() > 0 {
let hidden_states = if !hidden_states.is_empty() {
Some(hidden_states)
} else {
None
@@ -521,9 +518,9 @@ impl Wav2Vec2BertEncoder {
}
pub struct W2VBert2_0Model {
config: W2VBert2_0Config,
// config: W2VBert2_0Config,
feature_projection: Wav2Vec2BertFeatureProjection,
masked_spec_embed: Option<Tensor>,
// masked_spec_embed: Option<Tensor>,
encoder: Wav2Vec2BertEncoder,
// config.add_adapter is false, adapter is None, Wav2Vec2BertAdapter not complish
// adapter: Option<Wav2Vec2BertAdapter>,
@@ -543,21 +540,21 @@ impl W2VBert2_0Model {
pub fn new(vb: VarBuilder, config: &W2VBert2_0Config) -> Result<Self> {
let feature_projection =
Wav2Vec2BertFeatureProjection::new(vb.pp("feature_projection"), config)?;
let masked_spec_embed = if config.mask_time_prob > 0.0 || config.mask_time_prob > 0.0 {
Some(
vb.get_with_hints(config.hidden_size, "masked_spec_embed", Init::Uniform {
lo: 0.0,
up: 1.0,
})?,
)
} else {
None
};
// let masked_spec_embed = if config.mask_time_prob > 0.0 || config.mask_time_prob > 0.0 {
// Some(
// vb.get_with_hints(config.hidden_size, "masked_spec_embed", Init::Uniform {
// lo: 0.0,
// up: 1.0,
// })?,
// )
// } else {
// None
// };
let encoder = Wav2Vec2BertEncoder::new(vb.pp("encoder"), config)?;
Ok(Self {
config: config.clone(),
// config: config.clone(),
feature_projection,
masked_spec_embed,
// masked_spec_embed,
encoder,
})
}