style: cargo fmt --workspace (whitespace/wrapping only, no semantic change)
Whole-workspace rustfmt pass picked up while iterating on Mamba GPU backward work. Verified formatting-only via diff sampling; no logic changed. Co-Authored-By: Claude Sonnet 5 <[email protected]>
This commit is contained in:
@@ -7,15 +7,21 @@
|
||||
// rtx-tts uses the concrete-tensor API and does not yet require backend dispatch.
|
||||
#![allow(deprecated)]
|
||||
|
||||
use crate::acoustic::{AcousticModel, MelSpectrogramConfig, VarianceAdaptor, VarianceAdaptorConfig};
|
||||
use crate::Result;
|
||||
use rtx_nn::layers::{Module, linear::Linear, attention::{MultiHeadAttention, AttentionConfig}};
|
||||
use crate::acoustic::{
|
||||
AcousticModel, MelSpectrogramConfig, VarianceAdaptor, VarianceAdaptorConfig,
|
||||
};
|
||||
use rtx_nn::layers::activation::ReLU;
|
||||
use rtx_nn::layers::norm::layer_norm::LayerNorm;
|
||||
use rtx_nn::layers::conv::{Conv1d, Conv1dConfig, Conv1dPadding};
|
||||
use rtx_nn::layers::dropout::Dropout;
|
||||
use rtx_nn::layers::embedding::{Embedding, EmbeddingConfig};
|
||||
use rtx_nn::layers::conv::{Conv1d, Conv1dConfig, Conv1dPadding};
|
||||
use rtx_tensor::{Tensor, Device, DType};
|
||||
use rtx_nn::layers::norm::layer_norm::LayerNorm;
|
||||
use rtx_nn::layers::{
|
||||
Module,
|
||||
attention::{AttentionConfig, MultiHeadAttention},
|
||||
linear::Linear,
|
||||
};
|
||||
use rtx_tensor::{DType, Device, Tensor};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Configuration for FastSpeech2
|
||||
@@ -77,20 +83,28 @@ impl FastSpeech2Config {
|
||||
/// Validate configuration
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
if self.encoder_layers == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("encoder_layers must be > 0".into()));
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"encoder_layers must be > 0".into(),
|
||||
));
|
||||
}
|
||||
if self.decoder_layers == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("decoder_layers must be > 0".into()));
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"decoder_layers must be > 0".into(),
|
||||
));
|
||||
}
|
||||
if self.hidden_size == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("hidden_size must be > 0".into()));
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"hidden_size must be > 0".into(),
|
||||
));
|
||||
}
|
||||
if self.num_heads == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("num_heads must be > 0".into()));
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"num_heads must be > 0".into(),
|
||||
));
|
||||
}
|
||||
if self.hidden_size % self.num_heads != 0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"hidden_size must be divisible by num_heads".into()
|
||||
"hidden_size must be divisible by num_heads".into(),
|
||||
));
|
||||
}
|
||||
self.mel_config.validate()?;
|
||||
@@ -170,12 +184,14 @@ impl FFTBlock {
|
||||
|
||||
fn forward(&self, x: &Tensor, mask: Option<&Tensor>) -> Result<Tensor> {
|
||||
// Self-attention with residual
|
||||
let (attn_out, _, _) = self.self_attn.forward_with_cache(
|
||||
x, None, None, None, mask, None, false
|
||||
).map_err(|e| crate::TtsError::ModelError(format!("Attention forward failed: {e}")))?;
|
||||
let (attn_out, _, _) = self
|
||||
.self_attn
|
||||
.forward_with_cache(x, None, None, None, mask, None, false)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Attention forward failed: {e}")))?;
|
||||
|
||||
let attn_out = if self.training {
|
||||
self.dropout.forward(&attn_out)
|
||||
self.dropout
|
||||
.forward(&attn_out)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Dropout failed: {e}")))?
|
||||
} else {
|
||||
attn_out
|
||||
@@ -184,27 +200,37 @@ impl FFTBlock {
|
||||
let residual = (x + &attn_out)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Add failed: {e}")))?;
|
||||
|
||||
let normed = self.attn_norm.forward(&residual)
|
||||
let normed = self
|
||||
.attn_norm
|
||||
.forward(&residual)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("LayerNorm failed: {e}")))?;
|
||||
|
||||
// Feed-forward with residual
|
||||
let ff_input = normed.transpose(1, 2)
|
||||
let ff_input = normed
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let conv1_out = self.conv1.forward(&ff_input)
|
||||
let conv1_out = self
|
||||
.conv1
|
||||
.forward(&ff_input)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1 failed: {e}")))?;
|
||||
|
||||
let activated = conv1_out.relu()
|
||||
let activated = conv1_out
|
||||
.relu()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("ReLU failed: {e}")))?;
|
||||
|
||||
let conv2_out = self.conv2.forward(&activated)
|
||||
let conv2_out = self
|
||||
.conv2
|
||||
.forward(&activated)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv2 failed: {e}")))?;
|
||||
|
||||
let ff_out = conv2_out.transpose(1, 2)
|
||||
let ff_out = conv2_out
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let ff_out = if self.training {
|
||||
self.dropout.forward(&ff_out)
|
||||
self.dropout
|
||||
.forward(&ff_out)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Dropout failed: {e}")))?
|
||||
} else {
|
||||
ff_out
|
||||
@@ -213,7 +239,8 @@ impl FFTBlock {
|
||||
let residual = (&normed + &ff_out)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Add failed: {e}")))?;
|
||||
|
||||
self.conv_norm.forward(&residual)
|
||||
self.conv_norm
|
||||
.forward(&residual)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("LayerNorm failed: {e}")))
|
||||
}
|
||||
|
||||
@@ -320,17 +347,12 @@ impl FastSpeech2 {
|
||||
let mut variance_config = config.variance_config.clone();
|
||||
variance_config.hidden_dim = config.hidden_size;
|
||||
|
||||
let phoneme_emb_config = EmbeddingConfig::new(
|
||||
config.phoneme_vocab_size,
|
||||
config.hidden_size,
|
||||
);
|
||||
let phoneme_emb_config =
|
||||
EmbeddingConfig::new(config.phoneme_vocab_size, config.hidden_size);
|
||||
let phoneme_emb = Embedding::new(phoneme_emb_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Phoneme embedding failed: {e}")))?;
|
||||
|
||||
let position_emb_config = EmbeddingConfig::new(
|
||||
config.max_seq_len,
|
||||
config.hidden_size,
|
||||
);
|
||||
let position_emb_config = EmbeddingConfig::new(config.max_seq_len, config.hidden_size);
|
||||
let position_emb = Embedding::new(position_emb_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Position embedding failed: {e}")))?;
|
||||
|
||||
@@ -367,19 +389,28 @@ impl AcousticModel for FastSpeech2 {
|
||||
let seq_len = shape.dims()[1];
|
||||
|
||||
// Embed phonemes
|
||||
let phoneme_embedded = self.phoneme_emb.forward(phonemes)
|
||||
let phoneme_embedded = self
|
||||
.phoneme_emb
|
||||
.forward(phonemes)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Phoneme embedding failed: {e}")))?;
|
||||
|
||||
// Create position indices (cast to f32; Embedding interprets them as integer indices)
|
||||
let positions: Vec<f32> = (0..seq_len as u32).map(|i| i as f32).collect();
|
||||
let position_ids = Tensor::from_vec(
|
||||
positions.into_iter().cycle().take(batch_size * seq_len).collect(),
|
||||
positions
|
||||
.into_iter()
|
||||
.cycle()
|
||||
.take(batch_size * seq_len)
|
||||
.collect(),
|
||||
&[batch_size, seq_len],
|
||||
&self.device
|
||||
).map_err(|e| crate::TtsError::TensorError(format!("Position tensor failed: {e}")))?;
|
||||
&self.device,
|
||||
)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Position tensor failed: {e}")))?;
|
||||
|
||||
// Embed positions
|
||||
let position_embedded = self.position_emb.forward(&position_ids)
|
||||
let position_embedded = self
|
||||
.position_emb
|
||||
.forward(&position_ids)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Position embedding failed: {e}")))?;
|
||||
|
||||
// Add positional embeddings
|
||||
@@ -395,7 +426,9 @@ impl AcousticModel for FastSpeech2 {
|
||||
let decoded = self.decoder.forward(hidden, None)?;
|
||||
|
||||
// Project to mel dimension
|
||||
let mel = self.mel_linear.forward(&decoded)
|
||||
let mel = self
|
||||
.mel_linear
|
||||
.forward(&decoded)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Mel projection failed: {e}")))?;
|
||||
|
||||
// Transpose to [batch, mel_dim, time] format
|
||||
@@ -485,7 +518,13 @@ mod tests {
|
||||
|
||||
let batch_size = 2;
|
||||
let seq_len = 10;
|
||||
let phonemes = Tensor::randint(0, config.phoneme_vocab_size as i32, &[batch_size, seq_len], &device).unwrap();
|
||||
let phonemes = Tensor::randint(
|
||||
0,
|
||||
config.phoneme_vocab_size as i32,
|
||||
&[batch_size, seq_len],
|
||||
&device,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let encoded = model.encode(&phonemes, None);
|
||||
assert!(encoded.is_ok());
|
||||
@@ -505,7 +544,13 @@ mod tests {
|
||||
|
||||
let batch_size = 2;
|
||||
let seq_len = 10;
|
||||
let phonemes = Tensor::randint(0, config.phoneme_vocab_size as i32, &[batch_size, seq_len], &device).unwrap();
|
||||
let phonemes = Tensor::randint(
|
||||
0,
|
||||
config.phoneme_vocab_size as i32,
|
||||
&[batch_size, seq_len],
|
||||
&device,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let mel = model.forward(&phonemes, None);
|
||||
assert!(mel.is_ok());
|
||||
@@ -642,7 +687,8 @@ mod tests {
|
||||
config.decoder_layers = 2;
|
||||
let model = FastSpeech2::new(config.clone(), &device).unwrap();
|
||||
|
||||
let phonemes = Tensor::randint(0, config.phoneme_vocab_size as i32, &[1, 5], &device).unwrap();
|
||||
let phonemes =
|
||||
Tensor::randint(0, config.phoneme_vocab_size as i32, &[1, 5], &device).unwrap();
|
||||
let mel = model.forward(&phonemes, None);
|
||||
assert!(mel.is_ok());
|
||||
}
|
||||
|
||||
@@ -34,9 +34,8 @@ pub mod variance_adaptor;
|
||||
pub use fastspeech2::{FastSpeech2, FastSpeech2Config};
|
||||
pub use tacotron::{Tacotron2, Tacotron2Config};
|
||||
pub use variance_adaptor::{
|
||||
VarianceAdaptor, VarianceAdaptorConfig,
|
||||
DurationPredictor, PitchPredictor, EnergyPredictor,
|
||||
LengthRegulator,
|
||||
DurationPredictor, EnergyPredictor, LengthRegulator, PitchPredictor, VarianceAdaptor,
|
||||
VarianceAdaptorConfig,
|
||||
};
|
||||
|
||||
/// Configuration for mel spectrogram generation
|
||||
@@ -100,42 +99,42 @@ impl MelSpectrogramConfig {
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
if self.n_mels == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"n_mels must be greater than 0".into()
|
||||
"n_mels must be greater than 0".into(),
|
||||
));
|
||||
}
|
||||
if self.sample_rate == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"sample_rate must be greater than 0".into()
|
||||
"sample_rate must be greater than 0".into(),
|
||||
));
|
||||
}
|
||||
if self.n_fft == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"n_fft must be greater than 0".into()
|
||||
"n_fft must be greater than 0".into(),
|
||||
));
|
||||
}
|
||||
if self.hop_length == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"hop_length must be greater than 0".into()
|
||||
"hop_length must be greater than 0".into(),
|
||||
));
|
||||
}
|
||||
if self.win_length == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"win_length must be greater than 0".into()
|
||||
"win_length must be greater than 0".into(),
|
||||
));
|
||||
}
|
||||
if self.f_min < 0.0 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"f_min must be non-negative".into()
|
||||
"f_min must be non-negative".into(),
|
||||
));
|
||||
}
|
||||
if self.f_max <= self.f_min {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"f_max must be greater than f_min".into()
|
||||
"f_max must be greater than f_min".into(),
|
||||
));
|
||||
}
|
||||
if self.f_max > (self.sample_rate / 2) as f32 {
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"f_max must be <= sample_rate/2 (Nyquist frequency)".into()
|
||||
"f_max must be <= sample_rate/2 (Nyquist frequency)".into(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
|
||||
@@ -7,13 +7,17 @@
|
||||
// rtx-tts uses the concrete-tensor API and does not yet require backend dispatch.
|
||||
#![allow(deprecated)]
|
||||
|
||||
use crate::acoustic::{AcousticModel, MelSpectrogramConfig};
|
||||
use crate::Result;
|
||||
use rtx_nn::layers::{Module, linear::Linear, conv::{Conv1d, Conv1dConfig, Conv1dPadding}};
|
||||
use rtx_nn::layers::activation::{ReLU, Tanh, Sigmoid};
|
||||
use crate::acoustic::{AcousticModel, MelSpectrogramConfig};
|
||||
use rtx_nn::layers::activation::{ReLU, Sigmoid, Tanh};
|
||||
use rtx_nn::layers::dropout::Dropout;
|
||||
use rtx_nn::layers::embedding::{Embedding, EmbeddingConfig};
|
||||
use rtx_tensor::{Tensor, Device};
|
||||
use rtx_nn::layers::{
|
||||
Module,
|
||||
conv::{Conv1d, Conv1dConfig, Conv1dPadding},
|
||||
linear::Linear,
|
||||
};
|
||||
use rtx_tensor::{Device, Tensor};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Configuration for Tacotron2
|
||||
@@ -79,10 +83,14 @@ impl Tacotron2Config {
|
||||
/// Validate configuration
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
if self.encoder_dim == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("encoder_dim must be > 0".into()));
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"encoder_dim must be > 0".into(),
|
||||
));
|
||||
}
|
||||
if self.decoder_dim == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("decoder_dim must be > 0".into()));
|
||||
return Err(crate::TtsError::InvalidConfig(
|
||||
"decoder_dim must be > 0".into(),
|
||||
));
|
||||
}
|
||||
if self.mel_dim == 0 {
|
||||
return Err(crate::TtsError::InvalidConfig("mel_dim must be > 0".into()));
|
||||
@@ -119,19 +127,30 @@ impl PreNet {
|
||||
}
|
||||
|
||||
fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
let h1 = self.fc1.forward(x)
|
||||
let h1 = self
|
||||
.fc1
|
||||
.forward(x)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("FC1 forward failed: {e}")))?;
|
||||
let h1 = self.activation.forward(&h1)
|
||||
let h1 = self
|
||||
.activation
|
||||
.forward(&h1)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("ReLU failed: {e}")))?;
|
||||
let h1 = self.dropout.forward(&h1)
|
||||
let h1 = self
|
||||
.dropout
|
||||
.forward(&h1)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Dropout failed: {e}")))?;
|
||||
|
||||
let h2 = self.fc2.forward(&h1)
|
||||
let h2 = self
|
||||
.fc2
|
||||
.forward(&h1)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("FC2 forward failed: {e}")))?;
|
||||
let h2 = self.activation.forward(&h2)
|
||||
let h2 = self
|
||||
.activation
|
||||
.forward(&h2)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("ReLU failed: {e}")))?;
|
||||
|
||||
self.dropout.forward(&h2)
|
||||
self.dropout
|
||||
.forward(&h2)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Dropout failed: {e}")))
|
||||
}
|
||||
|
||||
@@ -156,7 +175,11 @@ impl PostNet {
|
||||
let mut batch_norm_layers = Vec::new();
|
||||
|
||||
for i in 0..config.postnet_n_convs {
|
||||
let in_channels = if i == 0 { config.mel_dim } else { config.postnet_channels };
|
||||
let in_channels = if i == 0 {
|
||||
config.mel_dim
|
||||
} else {
|
||||
config.postnet_channels
|
||||
};
|
||||
let out_channels = if i == config.postnet_n_convs - 1 {
|
||||
config.mel_dim
|
||||
} else {
|
||||
@@ -175,12 +198,16 @@ impl PostNet {
|
||||
bias: true,
|
||||
};
|
||||
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d creation failed: {e}")))?);
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Conv1d creation failed: {e}"))
|
||||
})?);
|
||||
|
||||
// Simplified batch norm
|
||||
batch_norm_layers.push(Linear::new(out_channels, out_channels, true, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Batch norm creation failed: {e}")))?);
|
||||
batch_norm_layers.push(
|
||||
Linear::new(out_channels, out_channels, true, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Batch norm creation failed: {e}"))
|
||||
})?,
|
||||
);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
@@ -195,29 +222,42 @@ impl PostNet {
|
||||
fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
let mut hidden = x.clone();
|
||||
|
||||
for (i, (conv, bn)) in self.conv_layers.iter().zip(self.batch_norm_layers.iter()).enumerate() {
|
||||
let conv_out = conv.forward(&hidden)
|
||||
for (i, (conv, bn)) in self
|
||||
.conv_layers
|
||||
.iter()
|
||||
.zip(self.batch_norm_layers.iter())
|
||||
.enumerate()
|
||||
{
|
||||
let conv_out = conv
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv forward failed: {e}")))?;
|
||||
|
||||
// Transpose for batch norm
|
||||
let transposed = conv_out.transpose(1, 2)
|
||||
let transposed = conv_out
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let bn_out = bn.forward(&transposed)
|
||||
let bn_out = bn
|
||||
.forward(&transposed)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Batch norm failed: {e}")))?;
|
||||
|
||||
// Transpose back
|
||||
hidden = bn_out.transpose(1, 2)
|
||||
hidden = bn_out
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
// Tanh activation except for last layer
|
||||
if i < self.conv_layers.len() - 1 {
|
||||
hidden = self.tanh.forward(&hidden)
|
||||
hidden = self
|
||||
.tanh
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Tanh failed: {e}")))?;
|
||||
}
|
||||
|
||||
if self.training {
|
||||
hidden = self.dropout.forward(&hidden)
|
||||
hidden = self
|
||||
.dropout
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Dropout failed: {e}")))?;
|
||||
}
|
||||
}
|
||||
@@ -252,11 +292,13 @@ impl LocationAttention {
|
||||
attention_dim: usize,
|
||||
device: &Device,
|
||||
) -> Result<Self> {
|
||||
let query_layer = Linear::new(query_dim, attention_dim, false, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Query layer creation failed: {e}")))?;
|
||||
let query_layer = Linear::new(query_dim, attention_dim, false, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Query layer creation failed: {e}"))
|
||||
})?;
|
||||
|
||||
let memory_layer = Linear::new(memory_dim, attention_dim, false, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Memory layer creation failed: {e}")))?;
|
||||
let memory_layer = Linear::new(memory_dim, attention_dim, false, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Memory layer creation failed: {e}"))
|
||||
})?;
|
||||
|
||||
let location_conv_config = Conv1dConfig {
|
||||
in_channels: 2,
|
||||
@@ -270,11 +312,13 @@ impl LocationAttention {
|
||||
bias: true,
|
||||
};
|
||||
|
||||
let location_conv = Conv1d::from_config(location_conv_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Location conv creation failed: {e}")))?;
|
||||
let location_conv = Conv1d::from_config(location_conv_config, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Location conv creation failed: {e}"))
|
||||
})?;
|
||||
|
||||
let location_layer = Linear::new(32, attention_dim, false, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Location layer creation failed: {e}")))?;
|
||||
let location_layer = Linear::new(32, attention_dim, false, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Location layer creation failed: {e}"))
|
||||
})?;
|
||||
|
||||
let v = Linear::new(attention_dim, 1, false, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("V layer creation failed: {e}")))?;
|
||||
@@ -296,24 +340,34 @@ impl LocationAttention {
|
||||
memory: &Tensor,
|
||||
attention_weights_cat: &Tensor,
|
||||
) -> Result<Tensor> {
|
||||
let processed_query = self.query_layer.forward(query)
|
||||
let processed_query = self
|
||||
.query_layer
|
||||
.forward(query)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Query layer failed: {e}")))?;
|
||||
|
||||
let processed_memory = self.memory_layer.forward(memory)
|
||||
let processed_memory = self
|
||||
.memory_layer
|
||||
.forward(memory)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Memory layer failed: {e}")))?;
|
||||
|
||||
// Process location features
|
||||
let processed_location = self.location_conv.forward(attention_weights_cat)
|
||||
let processed_location = self
|
||||
.location_conv
|
||||
.forward(attention_weights_cat)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Location conv failed: {e}")))?;
|
||||
|
||||
let processed_location = processed_location.transpose(1, 2)
|
||||
let processed_location = processed_location
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let processed_location = self.location_layer.forward(&processed_location)
|
||||
let processed_location = self
|
||||
.location_layer
|
||||
.forward(&processed_location)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Location layer failed: {e}")))?;
|
||||
|
||||
// Expand query to match memory sequence length
|
||||
let query_expanded = processed_query.unsqueeze(1)
|
||||
let query_expanded = processed_query
|
||||
.unsqueeze(1)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Unsqueeze failed: {e}")))?;
|
||||
|
||||
// Compute alignment energies
|
||||
@@ -322,13 +376,17 @@ impl LocationAttention {
|
||||
let energies = (&energies + &processed_location)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Add failed: {e}")))?;
|
||||
|
||||
let energies = energies.tanh()
|
||||
let energies = energies
|
||||
.tanh()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Tanh failed: {e}")))?;
|
||||
|
||||
let alignment = self.v.forward(&energies)
|
||||
let alignment = self
|
||||
.v
|
||||
.forward(&energies)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("V layer failed: {e}")))?;
|
||||
|
||||
let alignment = alignment.squeeze(Some(2))
|
||||
let alignment = alignment
|
||||
.squeeze(Some(2))
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Squeeze failed: {e}")))?;
|
||||
|
||||
// Softmax
|
||||
@@ -367,8 +425,9 @@ impl Encoder {
|
||||
bias: true,
|
||||
};
|
||||
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d creation failed: {e}")))?);
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Conv1d creation failed: {e}"))
|
||||
})?);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
@@ -381,28 +440,36 @@ impl Encoder {
|
||||
}
|
||||
|
||||
fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
let embedded = self.embedding.forward(x)
|
||||
let embedded = self
|
||||
.embedding
|
||||
.forward(x)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Embedding failed: {e}")))?;
|
||||
|
||||
// Transpose for conv: [batch, seq, dim] -> [batch, dim, seq]
|
||||
let mut hidden = embedded.transpose(1, 2)
|
||||
let mut hidden = embedded
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
for conv in &self.conv_layers {
|
||||
let conv_out = conv.forward(&hidden)
|
||||
let conv_out = conv
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv forward failed: {e}")))?;
|
||||
|
||||
hidden = conv_out.relu()
|
||||
hidden = conv_out
|
||||
.relu()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("ReLU failed: {e}")))?;
|
||||
|
||||
if self.training {
|
||||
hidden = self.dropout.forward(&hidden)
|
||||
hidden = self
|
||||
.dropout
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Dropout failed: {e}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
// Transpose back: [batch, dim, seq] -> [batch, seq, dim]
|
||||
hidden.transpose(1, 2)
|
||||
hidden
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))
|
||||
}
|
||||
|
||||
@@ -452,8 +519,13 @@ impl Tacotron2 {
|
||||
let decoder_rnn_hidden = Linear::new(config.decoder_dim, config.decoder_dim, true, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Decoder RNN hidden failed: {e}")))?;
|
||||
|
||||
let mel_projection = Linear::new(config.decoder_dim + config.encoder_dim, config.mel_dim, true, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Mel projection failed: {e}")))?;
|
||||
let mel_projection = Linear::new(
|
||||
config.decoder_dim + config.encoder_dim,
|
||||
config.mel_dim,
|
||||
true,
|
||||
device,
|
||||
)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Mel projection failed: {e}")))?;
|
||||
|
||||
let gate_projection = Linear::new(config.decoder_dim + config.encoder_dim, 1, true, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Gate projection failed: {e}")))?;
|
||||
@@ -492,23 +564,25 @@ impl Tacotron2 {
|
||||
let prenet_out = self.prenet.forward(decoder_input)?;
|
||||
|
||||
// Attention
|
||||
let attention_weights = self.attention.forward(
|
||||
decoder_hidden,
|
||||
encoder_outputs,
|
||||
attention_weights_cat,
|
||||
)?;
|
||||
let attention_weights =
|
||||
self.attention
|
||||
.forward(decoder_hidden, encoder_outputs, attention_weights_cat)?;
|
||||
|
||||
// Apply attention to encoder outputs
|
||||
let attention_weights_expanded = attention_weights.unsqueeze(1)
|
||||
let attention_weights_expanded = attention_weights
|
||||
.unsqueeze(1)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Unsqueeze failed: {e}")))?;
|
||||
|
||||
let encoder_outputs_transposed = encoder_outputs.transpose(1, 2)
|
||||
let encoder_outputs_transposed = encoder_outputs
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let attention_context = rtx_tensor::ops::matmul(&attention_weights_expanded, &encoder_outputs_transposed)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Matmul failed: {e}")))?;
|
||||
let attention_context =
|
||||
rtx_tensor::ops::matmul(&attention_weights_expanded, &encoder_outputs_transposed)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Matmul failed: {e}")))?;
|
||||
|
||||
let attention_context = attention_context.squeeze(Some(1))
|
||||
let attention_context = attention_context
|
||||
.squeeze(Some(1))
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Squeeze failed: {e}")))?;
|
||||
|
||||
// Concatenate prenet output and attention context
|
||||
@@ -516,29 +590,48 @@ impl Tacotron2 {
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Cat failed: {e}")))?;
|
||||
|
||||
// Decoder RNN step
|
||||
let rnn_input_proj = self.decoder_rnn_input.forward(&decoder_rnn_input)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("RNN input projection failed: {e}")))?;
|
||||
let rnn_input_proj = self
|
||||
.decoder_rnn_input
|
||||
.forward(&decoder_rnn_input)
|
||||
.map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("RNN input projection failed: {e}"))
|
||||
})?;
|
||||
|
||||
let rnn_hidden_proj = self.decoder_rnn_hidden.forward(decoder_hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("RNN hidden projection failed: {e}")))?;
|
||||
let rnn_hidden_proj = self
|
||||
.decoder_rnn_hidden
|
||||
.forward(decoder_hidden)
|
||||
.map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("RNN hidden projection failed: {e}"))
|
||||
})?;
|
||||
|
||||
let decoder_hidden_new = (&rnn_input_proj + &rnn_hidden_proj)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Add failed: {e}")))?;
|
||||
|
||||
let decoder_hidden_new = decoder_hidden_new.tanh()
|
||||
let decoder_hidden_new = decoder_hidden_new
|
||||
.tanh()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Tanh failed: {e}")))?;
|
||||
|
||||
// Projection to mel
|
||||
let projection_input = Tensor::cat(&[decoder_hidden_new.clone(), attention_context.clone()], 1)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Cat failed: {e}")))?;
|
||||
let projection_input =
|
||||
Tensor::cat(&[decoder_hidden_new.clone(), attention_context.clone()], 1)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Cat failed: {e}")))?;
|
||||
|
||||
let mel_output = self.mel_projection.forward(&projection_input)
|
||||
let mel_output = self
|
||||
.mel_projection
|
||||
.forward(&projection_input)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Mel projection failed: {e}")))?;
|
||||
|
||||
let gate_output = self.gate_projection.forward(&projection_input)
|
||||
let gate_output = self
|
||||
.gate_projection
|
||||
.forward(&projection_input)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Gate projection failed: {e}")))?;
|
||||
|
||||
Ok((mel_output, gate_output, decoder_hidden_new, attention_weights))
|
||||
Ok((
|
||||
mel_output,
|
||||
gate_output,
|
||||
decoder_hidden_new,
|
||||
attention_weights,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -577,20 +670,28 @@ impl AcousticModel for Tacotron2 {
|
||||
current_input = mel_out;
|
||||
|
||||
// Update attention weights cat
|
||||
let attn_expanded = attn_weights.unsqueeze(1)
|
||||
let attn_expanded = attn_weights
|
||||
.unsqueeze(1)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Unsqueeze failed: {e}")))?;
|
||||
|
||||
current_attn = Tensor::cat(&[
|
||||
current_attn.slice(1, 1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Slice failed: {e}")))?,
|
||||
attn_expanded
|
||||
], 1).map_err(|e| crate::TtsError::TensorError(format!("Cat failed: {e}")))?;
|
||||
current_attn = Tensor::cat(
|
||||
&[
|
||||
current_attn
|
||||
.slice(1, 1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Slice failed: {e}")))?,
|
||||
attn_expanded,
|
||||
],
|
||||
1,
|
||||
)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Cat failed: {e}")))?;
|
||||
|
||||
// Check gate (stop condition)
|
||||
let gate_sigmoid = gate_out.sigmoid()
|
||||
let gate_sigmoid = gate_out
|
||||
.sigmoid()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Sigmoid failed: {e}")))?;
|
||||
|
||||
let gate_val = gate_sigmoid.to_vec()
|
||||
let gate_val = gate_sigmoid
|
||||
.to_vec()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("To vec failed: {e}")))?;
|
||||
|
||||
if gate_val.iter().all(|&v| v > 0.5) {
|
||||
@@ -603,7 +704,8 @@ impl AcousticModel for Tacotron2 {
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Stack failed: {e}")))?;
|
||||
|
||||
// Transpose to [batch, mel_dim, time]
|
||||
let mel_transposed = mel_stacked.transpose(1, 2)
|
||||
let mel_transposed = mel_stacked
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
// Apply postnet
|
||||
@@ -699,7 +801,8 @@ mod tests {
|
||||
|
||||
let batch_size = 2;
|
||||
let seq_len = 10;
|
||||
let input = Tensor::randint(0, config.vocab_size as i32, &[batch_size, seq_len], &device).unwrap();
|
||||
let input =
|
||||
Tensor::randint(0, config.vocab_size as i32, &[batch_size, seq_len], &device).unwrap();
|
||||
|
||||
let output = encoder.forward(&input);
|
||||
assert!(output.is_ok());
|
||||
@@ -727,7 +830,8 @@ mod tests {
|
||||
|
||||
let batch_size = 2;
|
||||
let seq_len = 10;
|
||||
let phonemes = Tensor::randint(0, config.vocab_size as i32, &[batch_size, seq_len], &device).unwrap();
|
||||
let phonemes =
|
||||
Tensor::randint(0, config.vocab_size as i32, &[batch_size, seq_len], &device).unwrap();
|
||||
|
||||
let encoded = model.encode(&phonemes, None);
|
||||
assert!(encoded.is_ok());
|
||||
|
||||
@@ -8,10 +8,10 @@
|
||||
#![allow(deprecated)]
|
||||
|
||||
use crate::Result;
|
||||
use rtx_nn::layers::{Module, linear::Linear};
|
||||
use rtx_nn::layers::activation::ReLU;
|
||||
use rtx_nn::layers::conv::{Conv1d, Conv1dConfig, Conv1dPadding};
|
||||
use rtx_tensor::{Tensor, Device};
|
||||
use rtx_nn::layers::{Module, linear::Linear};
|
||||
use rtx_tensor::{Device, Tensor};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Configuration for variance adaptor
|
||||
@@ -57,18 +57,24 @@ impl SoftplusApprox {
|
||||
// For numerical stability, use: max(0, x) + log(1 + exp(-|x|))
|
||||
let zeros = Tensor::zeros_like(x)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Zeros failed: {e}")))?;
|
||||
let relu_part = x.maximum(&zeros)
|
||||
let relu_part = x
|
||||
.maximum(&zeros)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Maximum failed: {e}")))?;
|
||||
|
||||
let abs_x = x.abs()
|
||||
let abs_x = x
|
||||
.abs()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Abs failed: {e}")))?;
|
||||
let neg_abs = abs_x.mul_scalar(-1.0)
|
||||
let neg_abs = abs_x
|
||||
.mul_scalar(-1.0)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Mul failed: {e}")))?;
|
||||
let exp_part = neg_abs.exp()
|
||||
let exp_part = neg_abs
|
||||
.exp()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Exp failed: {e}")))?;
|
||||
let one_plus_exp = exp_part.add_scalar(1.0)
|
||||
let one_plus_exp = exp_part
|
||||
.add_scalar(1.0)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Add failed: {e}")))?;
|
||||
let log_part = one_plus_exp.log()
|
||||
let log_part = one_plus_exp
|
||||
.log()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Log failed: {e}")))?;
|
||||
|
||||
(&relu_part + &log_part)
|
||||
@@ -90,7 +96,8 @@ impl SimpleNorm {
|
||||
}
|
||||
|
||||
fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
self.linear.forward(x)
|
||||
self.linear
|
||||
.forward(x)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Linear forward failed: {e}")))
|
||||
}
|
||||
}
|
||||
@@ -152,8 +159,9 @@ impl DurationPredictor {
|
||||
bias: true,
|
||||
};
|
||||
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d creation failed: {e}")))?);
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Conv1d creation failed: {e}"))
|
||||
})?);
|
||||
|
||||
norms.push(SimpleNorm::new(config.hidden_dim, device)?);
|
||||
}
|
||||
@@ -177,38 +185,48 @@ impl DurationPredictor {
|
||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
// x: [batch, seq_len, hidden_dim]
|
||||
// Transpose to [batch, hidden_dim, seq_len] for Conv1d
|
||||
let mut hidden = x.transpose(1, 2)
|
||||
let mut hidden = x
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
for (conv, norm) in self.conv_layers.iter().zip(self.norms.iter()) {
|
||||
let conv_out = conv.forward(&hidden)
|
||||
let conv_out = conv
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d forward failed: {e}")))?;
|
||||
|
||||
// Transpose back for norm: [batch, seq_len, hidden_dim]
|
||||
let transposed = conv_out.transpose(1, 2)
|
||||
let transposed = conv_out
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let normed = norm.forward(&transposed)?;
|
||||
|
||||
let activated = self.activation.forward(&normed)
|
||||
let activated = self
|
||||
.activation
|
||||
.forward(&normed)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("ReLU forward failed: {e}")))?;
|
||||
|
||||
let dropped = self.dropout.forward(&activated)?;
|
||||
|
||||
// Transpose back to [batch, hidden_dim, seq_len] for next conv
|
||||
hidden = dropped.transpose(1, 2)
|
||||
hidden = dropped
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
}
|
||||
|
||||
// Transpose to [batch, seq_len, hidden_dim] for linear
|
||||
let hidden = hidden.transpose(1, 2)
|
||||
let hidden = hidden
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let output = self.linear.forward(&hidden)
|
||||
let output = self
|
||||
.linear
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Linear forward failed: {e}")))?;
|
||||
|
||||
// Squeeze last dimension: [batch, seq_len, 1] -> [batch, seq_len]
|
||||
let squeezed = output.squeeze(Some(2))
|
||||
let squeezed = output
|
||||
.squeeze(Some(2))
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Squeeze failed: {e}")))?;
|
||||
|
||||
// Apply softplus to ensure positive durations
|
||||
@@ -257,8 +275,9 @@ impl PitchPredictor {
|
||||
bias: true,
|
||||
};
|
||||
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d creation failed: {e}")))?);
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Conv1d creation failed: {e}"))
|
||||
})?);
|
||||
|
||||
norms.push(SimpleNorm::new(config.hidden_dim, device)?);
|
||||
}
|
||||
@@ -279,34 +298,44 @@ impl PitchPredictor {
|
||||
|
||||
/// Forward pass
|
||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
let mut hidden = x.transpose(1, 2)
|
||||
let mut hidden = x
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
for (conv, norm) in self.conv_layers.iter().zip(self.norms.iter()) {
|
||||
let conv_out = conv.forward(&hidden)
|
||||
let conv_out = conv
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d forward failed: {e}")))?;
|
||||
|
||||
let transposed = conv_out.transpose(1, 2)
|
||||
let transposed = conv_out
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let normed = norm.forward(&transposed)?;
|
||||
|
||||
let activated = self.activation.forward(&normed)
|
||||
let activated = self
|
||||
.activation
|
||||
.forward(&normed)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("ReLU forward failed: {e}")))?;
|
||||
|
||||
let dropped = self.dropout.forward(&activated)?;
|
||||
|
||||
hidden = dropped.transpose(1, 2)
|
||||
hidden = dropped
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
}
|
||||
|
||||
let hidden = hidden.transpose(1, 2)
|
||||
let hidden = hidden
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let output = self.linear.forward(&hidden)
|
||||
let output = self
|
||||
.linear
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Linear forward failed: {e}")))?;
|
||||
|
||||
output.squeeze(Some(2))
|
||||
output
|
||||
.squeeze(Some(2))
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Squeeze failed: {e}")))
|
||||
}
|
||||
|
||||
@@ -353,8 +382,9 @@ impl EnergyPredictor {
|
||||
bias: true,
|
||||
};
|
||||
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d creation failed: {e}")))?);
|
||||
conv_layers.push(Conv1d::from_config(conv_config, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Conv1d creation failed: {e}"))
|
||||
})?);
|
||||
|
||||
norms.push(SimpleNorm::new(config.hidden_dim, device)?);
|
||||
}
|
||||
@@ -376,34 +406,44 @@ impl EnergyPredictor {
|
||||
|
||||
/// Forward pass
|
||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||
let mut hidden = x.transpose(1, 2)
|
||||
let mut hidden = x
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
for (conv, norm) in self.conv_layers.iter().zip(self.norms.iter()) {
|
||||
let conv_out = conv.forward(&hidden)
|
||||
let conv_out = conv
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Conv1d forward failed: {e}")))?;
|
||||
|
||||
let transposed = conv_out.transpose(1, 2)
|
||||
let transposed = conv_out
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let normed = norm.forward(&transposed)?;
|
||||
|
||||
let activated = self.activation.forward(&normed)
|
||||
let activated = self
|
||||
.activation
|
||||
.forward(&normed)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("ReLU forward failed: {e}")))?;
|
||||
|
||||
let dropped = self.dropout.forward(&activated)?;
|
||||
|
||||
hidden = dropped.transpose(1, 2)
|
||||
hidden = dropped
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
}
|
||||
|
||||
let hidden = hidden.transpose(1, 2)
|
||||
let hidden = hidden
|
||||
.transpose(1, 2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Transpose failed: {e}")))?;
|
||||
|
||||
let output = self.linear.forward(&hidden)
|
||||
let output = self
|
||||
.linear
|
||||
.forward(&hidden)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Linear forward failed: {e}")))?;
|
||||
|
||||
let squeezed = output.squeeze(Some(2))
|
||||
let squeezed = output
|
||||
.squeeze(Some(2))
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Squeeze failed: {e}")))?;
|
||||
|
||||
// Apply softplus to ensure positive energy
|
||||
@@ -441,7 +481,8 @@ impl LengthRegulator {
|
||||
let hidden_dim = shape.dims()[2];
|
||||
|
||||
// Convert durations to integers
|
||||
let durations_data = durations.to_vec()
|
||||
let durations_data = durations
|
||||
.to_vec()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Failed to get durations: {e}")))?;
|
||||
|
||||
let mut outputs = Vec::new();
|
||||
@@ -456,8 +497,9 @@ impl LengthRegulator {
|
||||
let start_idx = b * seq_len * hidden_dim + i * hidden_dim;
|
||||
let end_idx = start_idx + hidden_dim;
|
||||
|
||||
let hidden_data = hidden.to_vec()
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Failed to get hidden: {e}")))?;
|
||||
let hidden_data = hidden.to_vec().map_err(|e| {
|
||||
crate::TtsError::TensorError(format!("Failed to get hidden: {e}"))
|
||||
})?;
|
||||
|
||||
let frame = &hidden_data[start_idx..end_idx];
|
||||
|
||||
@@ -471,7 +513,11 @@ impl LengthRegulator {
|
||||
}
|
||||
|
||||
// Find max length
|
||||
let max_len = outputs.iter().map(|x| x.len() / hidden_dim).max().unwrap_or(0);
|
||||
let max_len = outputs
|
||||
.iter()
|
||||
.map(|x| x.len() / hidden_dim)
|
||||
.max()
|
||||
.unwrap_or(0);
|
||||
|
||||
// Pad all sequences to max length
|
||||
let mut padded = Vec::new();
|
||||
@@ -482,8 +528,9 @@ impl LengthRegulator {
|
||||
padded.extend(seq);
|
||||
}
|
||||
|
||||
Tensor::from_vec(padded, &[batch_size, max_len, hidden_dim], &self.device)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Failed to create expanded tensor: {e}")))
|
||||
Tensor::from_vec(padded, &[batch_size, max_len, hidden_dim], &self.device).map_err(|e| {
|
||||
crate::TtsError::TensorError(format!("Failed to create expanded tensor: {e}"))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -508,11 +555,13 @@ impl VarianceAdaptor {
|
||||
let energy_predictor = EnergyPredictor::new(&config, device)?;
|
||||
let length_regulator = LengthRegulator::new(device);
|
||||
|
||||
let pitch_embedding = Linear::new(1, config.hidden_dim, true, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Pitch embedding creation failed: {e}")))?;
|
||||
let pitch_embedding = Linear::new(1, config.hidden_dim, true, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Pitch embedding creation failed: {e}"))
|
||||
})?;
|
||||
|
||||
let energy_embedding = Linear::new(1, config.hidden_dim, true, device)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Energy embedding creation failed: {e}")))?;
|
||||
let energy_embedding = Linear::new(1, config.hidden_dim, true, device).map_err(|e| {
|
||||
crate::TtsError::ModelError(format!("Energy embedding creation failed: {e}"))
|
||||
})?;
|
||||
|
||||
Ok(Self {
|
||||
duration_predictor,
|
||||
@@ -534,14 +583,20 @@ impl VarianceAdaptor {
|
||||
let energy = self.energy_predictor.forward(hidden)?;
|
||||
|
||||
// Embed pitch and energy
|
||||
let pitch_unsqueezed = pitch.unsqueeze(2)
|
||||
let pitch_unsqueezed = pitch
|
||||
.unsqueeze(2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Unsqueeze failed: {e}")))?;
|
||||
let energy_unsqueezed = energy.unsqueeze(2)
|
||||
let energy_unsqueezed = energy
|
||||
.unsqueeze(2)
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Unsqueeze failed: {e}")))?;
|
||||
|
||||
let pitch_emb = self.pitch_embedding.forward(&pitch_unsqueezed)
|
||||
let pitch_emb = self
|
||||
.pitch_embedding
|
||||
.forward(&pitch_unsqueezed)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Pitch embedding failed: {e}")))?;
|
||||
let energy_emb = self.energy_embedding.forward(&energy_unsqueezed)
|
||||
let energy_emb = self
|
||||
.energy_embedding
|
||||
.forward(&energy_unsqueezed)
|
||||
.map_err(|e| crate::TtsError::ModelError(format!("Energy embedding failed: {e}")))?;
|
||||
|
||||
// Add pitch and energy to hidden states
|
||||
@@ -551,7 +606,9 @@ impl VarianceAdaptor {
|
||||
.map_err(|e| crate::TtsError::TensorError(format!("Add energy failed: {e}")))?;
|
||||
|
||||
// Length regulation (expand by durations)
|
||||
let expanded = self.length_regulator.forward(&hidden_with_variance, &duration)?;
|
||||
let expanded = self
|
||||
.length_regulator
|
||||
.forward(&hidden_with_variance, &duration)?;
|
||||
|
||||
Ok((expanded, duration, pitch, energy))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user