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:
osobh
2026-08-10 07:09:36 -07:00
co-authored by Claude Sonnet 5
parent ad6405663f
commit 4aaa36a57a
305 changed files with 25537 additions and 18337 deletions
@@ -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());
}
+10 -11
View File
@@ -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(())
+186 -82
View File
@@ -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))
}