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:
@@ -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