use std::collections::HashMap;
use std::path::Path;
use anyhow::{Context, Result};
use rayon::prelude::*;
use crate::weights::LazySt;
fn linear(x: &[f32], w: &[f32], b: &[f32], m: usize, n: usize, k: usize) -> Vec<f32> {
let mut out = vec![0f32; m * n];
out.par_chunks_mut(n).enumerate().for_each(|(i, orow)| {
let xr = &x[i * k..][..k];
for o in 0..n {
let wr = &w[o * k..][..k];
let mut acc = b[o];
for c in 0..k {
acc += xr[c] * wr[c];
}
orow[o] = acc;
}
});
out
}
fn layer_norm(x: &[f32], d: usize, w: &[f32], b: &[f32], eps: f32) -> Vec<f32> {
let mut out = vec![0f32; x.len()];
out.par_chunks_mut(d)
.zip(x.par_chunks(d))
.for_each(|(orow, row)| {
let mean = row.iter().sum::<f32>() / d as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / d as f32;
let inv = 1.0 / (var + eps).sqrt();
for i in 0..d {
orow[i] = (row[i] - mean) * inv * w[i] + b[i];
}
});
out
}
fn gelu(v: f32) -> f32 {
0.5 * v * (1.0 + libm::erff(v * std::f32::consts::FRAC_1_SQRT_2))
}
struct Linear {
w: Vec<f32>,
b: Vec<f32>,
n: usize,
k: usize,
}
impl Linear {
fn load(st: &LazySt, prefix: &str, n: usize, k: usize) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
let b = st.tensor_f32(&format!("{prefix}.bias"))?;
anyhow::ensure!(w.len() == n * k, "{prefix}.weight {} != {n}x{k}", w.len());
Ok(Self { w, b, n, k })
}
fn forward(&self, x: &[f32], m: usize) -> Vec<f32> {
linear(x, &self.w, &self.b, m, self.n, self.k)
}
}
struct Norm {
w: Vec<f32>,
b: Vec<f32>,
eps: f32,
}
impl Norm {
fn load(st: &LazySt, prefix: &str, eps: f32) -> Result<Self> {
Ok(Self {
w: st.tensor_f32(&format!("{prefix}.weight"))?,
b: st.tensor_f32(&format!("{prefix}.bias"))?,
eps,
})
}
fn forward(&self, x: &[f32], d: usize) -> Vec<f32> {
layer_norm(x, d, &self.w, &self.b, self.eps)
}
}
struct EncBlock {
q: Linear,
k: Linear,
v: Linear,
out: Linear,
ln1: Norm, fc1: Linear,
fc2: Linear,
ln2: Norm, }
struct DecBlock {
sa_q: Linear,
sa_k: Linear,
sa_v: Linear,
sa_out: Linear,
sa_ln: Norm, ca_q: Linear,
ca_k: Linear,
ca_v: Linear,
ca_out: Linear,
ca_ln: Norm, fc1: Linear,
fc2: Linear,
ln2: Norm, }
struct PostnetLayer {
w: Vec<f32>, bn_w: Vec<f32>,
bn_b: Vec<f32>,
bn_rm: Vec<f32>,
bn_rv: Vec<f32>,
out_ch: usize,
in_ch: usize,
k: usize,
tanh: bool,
}
struct Conv1d {
w: Vec<f32>,
b: Vec<f32>,
out_ch: usize,
in_ch: usize,
k: usize,
dilation: usize,
pad: usize,
}
impl Conv1d {
fn load(
st: &LazySt,
prefix: &str,
out_ch: usize,
in_ch: usize,
k: usize,
dilation: usize,
) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == out_ch * in_ch * k,
"{prefix}.weight {} != {out_ch}x{in_ch}x{k}",
w.len()
);
Ok(Self {
w,
b: st.tensor_f32(&format!("{prefix}.bias"))?,
out_ch,
in_ch,
k,
dilation,
pad: (k * dilation - dilation) / 2,
})
}
fn forward(&self, x: &[f32], t: usize) -> Vec<f32> {
let mut y = vec![0f32; self.out_ch * t];
y.par_chunks_mut(t).enumerate().for_each(|(o, orow)| {
for (ti, oval) in orow.iter_mut().enumerate() {
let mut acc = self.b[o];
for ic in 0..self.in_ch {
let xrow = &x[ic * t..][..t];
let wrow = &self.w[(o * self.in_ch + ic) * self.k..][..self.k];
for (kk, wv) in wrow.iter().enumerate() {
let src = ti + kk * self.dilation;
if src >= self.pad && src - self.pad < t {
acc += xrow[src - self.pad] * wv;
}
}
}
*oval = acc;
}
});
y
}
}
struct ConvT1d {
w: Vec<f32>,
b: Vec<f32>,
out_ch: usize,
in_ch: usize,
k: usize,
stride: usize,
pad: usize,
}
impl ConvT1d {
fn load(
st: &LazySt,
prefix: &str,
in_ch: usize,
out_ch: usize,
k: usize,
stride: usize,
) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == in_ch * out_ch * k,
"{prefix}.weight {} != {in_ch}x{out_ch}x{k}",
w.len()
);
Ok(Self {
w,
b: st.tensor_f32(&format!("{prefix}.bias"))?,
out_ch,
in_ch,
k,
stride,
pad: (k - stride) / 2,
})
}
fn forward(&self, x: &[f32], t: usize) -> (Vec<f32>, usize) {
let t_out = (t - 1) * self.stride + self.k - 2 * self.pad;
let mut y = vec![0f32; self.out_ch * t_out];
y.par_chunks_mut(t_out).enumerate().for_each(|(o, orow)| {
for (ti, oval) in orow.iter_mut().enumerate() {
let mut acc = self.b[o];
for kk in 0..self.k {
let up = ti + self.pad;
if up < kk || !(up - kk).is_multiple_of(self.stride) {
continue;
}
let src = (up - kk) / self.stride;
if src >= t {
continue;
}
for ic in 0..self.in_ch {
acc += x[ic * t + src] * self.w[(ic * self.out_ch + o) * self.k + kk];
}
}
*oval = acc;
}
});
(y, t_out)
}
}
struct ResBlock {
convs1: Vec<Conv1d>,
convs2: Vec<Conv1d>,
}
impl ResBlock {
fn forward(&self, x: &[f32], ch: usize, t: usize, slope: f32) -> Vec<f32> {
let lrelu = |v: &mut Vec<f32>| {
for e in v.iter_mut() {
if *e < 0.0 {
*e *= slope;
}
}
};
let mut x = x.to_vec();
for (c1, c2) in self.convs1.iter().zip(&self.convs2) {
let mut h = x.clone();
lrelu(&mut h);
let mut h = c1.forward(&h, t);
lrelu(&mut h);
let h = c2.forward(&h, t);
for i in 0..ch * t {
x[i] += h[i];
}
}
x
}
}
struct HifiGan {
mean: Vec<f32>,
scale: Vec<f32>,
conv_pre: Conv1d,
upsampler: Vec<ConvT1d>,
resblocks: Vec<ResBlock>, conv_post: Conv1d,
num_kernels: usize,
slope: f32,
}
impl HifiGan {
fn load(st: &LazySt, cfg: &serde_json::Value) -> Result<Self> {
let g = |k: &str| cfg.get(k).and_then(|x| x.as_u64()).unwrap_or(0) as usize;
let arr = |k: &str| -> Vec<usize> {
cfg.get(k)
.and_then(|x| x.as_array())
.map(|a| {
a.iter()
.filter_map(|v| v.as_u64().map(|u| u as usize))
.collect()
})
.unwrap_or_default()
};
let ch0 = g("upsample_initial_channel");
let rates = arr("upsample_rates");
let kernels = arr("upsample_kernel_sizes");
let res_k = arr("resblock_kernel_sizes");
let dils: Vec<Vec<usize>> = cfg
.get("resblock_dilation_sizes")
.and_then(|x| x.as_array())
.map(|a| {
a.iter()
.map(|row| {
row.as_array()
.map(|r| {
r.iter()
.filter_map(|v| v.as_u64().map(|u| u as usize))
.collect()
})
.unwrap_or_default()
})
.collect()
})
.unwrap_or_default();
let upsampler = rates
.iter()
.zip(&kernels)
.enumerate()
.map(|(i, (&r, &k))| {
ConvT1d::load(
st,
&format!("hifigan.upsampler.{i}"),
ch0 >> i,
ch0 >> (i + 1),
k,
r,
)
})
.collect::<Result<Vec<_>>>()?;
let mut resblocks = Vec::new();
for i in 0..rates.len() {
let ch = ch0 >> (i + 1);
for (j, (&k, dil)) in res_k.iter().zip(&dils).enumerate() {
let b = format!("hifigan.resblocks.{}", i * res_k.len() + j);
resblocks.push(ResBlock {
convs1: dil
.iter()
.enumerate()
.map(|(n, &d)| Conv1d::load(st, &format!("{b}.convs1.{n}"), ch, ch, k, d))
.collect::<Result<Vec<_>>>()?,
convs2: dil
.iter()
.enumerate()
.map(|(n, _)| Conv1d::load(st, &format!("{b}.convs2.{n}"), ch, ch, k, 1))
.collect::<Result<Vec<_>>>()?,
});
}
}
let last_ch = ch0 >> rates.len();
Ok(Self {
mean: st.tensor_f32("hifigan.mean")?,
scale: st.tensor_f32("hifigan.scale")?,
conv_pre: Conv1d::load(st, "hifigan.conv_pre", ch0, g("model_in_dim"), 7, 1)?,
upsampler,
resblocks,
conv_post: Conv1d::load(st, "hifigan.conv_post", 1, last_ch, 7, 1)?,
num_kernels: res_k.len(),
slope: cfg
.get("leaky_relu_slope")
.and_then(|x| x.as_f64())
.unwrap_or(0.1) as f32,
})
}
fn forward(&self, mel: &[f32], t: usize) -> Vec<f32> {
let nm = self.mean.len();
let mut x = vec![0f32; nm * t];
for i in 0..t {
for c in 0..nm {
x[c * t + i] = (mel[i * nm + c] - self.mean[c]) / self.scale[c];
}
}
let mut x = self.conv_pre.forward(&x, t);
let mut t = t;
for (i, up) in self.upsampler.iter().enumerate() {
for e in x.iter_mut() {
if *e < 0.0 {
*e *= self.slope;
}
}
let (xu, tu) = up.forward(&x, t);
t = tu;
let ch = up.out_ch;
let mut acc = self.resblocks[i * self.num_kernels].forward(&xu, ch, t, self.slope);
for j in 1..self.num_kernels {
let r = self.resblocks[i * self.num_kernels + j].forward(&xu, ch, t, self.slope);
for (a, b) in acc.iter_mut().zip(r) {
*a += b;
}
}
let inv = 1.0 / self.num_kernels as f32;
for a in acc.iter_mut() {
*a *= inv;
}
x = acc;
}
for e in x.iter_mut() {
if *e < 0.0 {
*e *= 0.01;
}
}
let y = self.conv_post.forward(&x, t);
y.iter().map(|v| v.tanh()).collect()
}
}
pub struct MelTrace {
pub prenet: Vec<f32>,
pub dec_h: Vec<f32>,
pub stop_logits: Vec<f32>,
pub mel_pre: Vec<f32>,
pub mel_post: Vec<f32>,
pub steps: usize,
}
pub struct SpeechT5Tts {
vocab: HashMap<String, u32>,
unk_id: u32,
eos_id: u32,
embed_tokens: Vec<f32>, enc_alpha: f32,
enc_pe: Vec<f32>, enc_ln0: Norm, pe_k: Vec<f32>, max_rel: i64,
enc_blocks: Vec<EncBlock>,
dec_alpha: f32,
dec_pe: Vec<f32>, prenet: Vec<Linear>, prenet_final: Linear, speaker_layer: Linear, dec_blocks: Vec<DecBlock>,
feat_out: Linear, prob_out: Linear, postnet: Vec<PostnetLayer>,
hifigan: HifiGan,
n_mels: usize,
reduction: usize,
hidden: usize,
heads: usize,
hd: usize,
sampling_rate: u32,
}
impl SpeechT5Tts {
pub fn load(dir: &Path) -> Result<Self> {
let v: serde_json::Value =
serde_json::from_slice(&std::fs::read(dir.join("config.json"))?).context("config")?;
let g = |k: &str| v.get(k).and_then(|x| x.as_u64()).unwrap_or(0) as usize;
let st = LazySt::open(dir)?;
let hidden = g("hidden_size");
let heads = g("encoder_attention_heads");
let inter = g("encoder_ffn_dim");
let eps = v
.get("layer_norm_eps")
.and_then(|x| x.as_f64())
.unwrap_or(1e-5) as f32;
let vocab = v
.get("vocab")
.and_then(|m| m.as_object())
.context("config vocab")?
.iter()
.filter_map(|(k, val)| val.as_u64().map(|i| (k.clone(), i as u32)))
.collect();
let enc = "speecht5.encoder";
let enc_blocks = (0..g("encoder_layers"))
.map(|i| {
let b = format!("{enc}.wrapped_encoder.layers.{i}");
Ok(EncBlock {
q: Linear::load(&st, &format!("{b}.attention.q_proj"), hidden, hidden)?,
k: Linear::load(&st, &format!("{b}.attention.k_proj"), hidden, hidden)?,
v: Linear::load(&st, &format!("{b}.attention.v_proj"), hidden, hidden)?,
out: Linear::load(&st, &format!("{b}.attention.out_proj"), hidden, hidden)?,
ln1: Norm::load(&st, &format!("{b}.layer_norm"), eps)?,
fc1: Linear::load(
&st,
&format!("{b}.feed_forward.intermediate_dense"),
inter,
hidden,
)?,
fc2: Linear::load(
&st,
&format!("{b}.feed_forward.output_dense"),
hidden,
inter,
)?,
ln2: Norm::load(&st, &format!("{b}.final_layer_norm"), eps)?,
})
})
.collect::<Result<Vec<_>>>()?;
let dec = "speecht5.decoder";
let dec_blocks = (0..g("decoder_layers"))
.map(|i| {
let b = format!("{dec}.wrapped_decoder.layers.{i}");
Ok(DecBlock {
sa_q: Linear::load(&st, &format!("{b}.self_attn.q_proj"), hidden, hidden)?,
sa_k: Linear::load(&st, &format!("{b}.self_attn.k_proj"), hidden, hidden)?,
sa_v: Linear::load(&st, &format!("{b}.self_attn.v_proj"), hidden, hidden)?,
sa_out: Linear::load(&st, &format!("{b}.self_attn.out_proj"), hidden, hidden)?,
sa_ln: Norm::load(&st, &format!("{b}.self_attn_layer_norm"), eps)?,
ca_q: Linear::load(&st, &format!("{b}.encoder_attn.q_proj"), hidden, hidden)?,
ca_k: Linear::load(&st, &format!("{b}.encoder_attn.k_proj"), hidden, hidden)?,
ca_v: Linear::load(&st, &format!("{b}.encoder_attn.v_proj"), hidden, hidden)?,
ca_out: Linear::load(
&st,
&format!("{b}.encoder_attn.out_proj"),
hidden,
hidden,
)?,
ca_ln: Norm::load(&st, &format!("{b}.encoder_attn_layer_norm"), eps)?,
fc1: Linear::load(
&st,
&format!("{b}.feed_forward.intermediate_dense"),
g("decoder_ffn_dim"),
hidden,
)?,
fc2: Linear::load(
&st,
&format!("{b}.feed_forward.output_dense"),
hidden,
g("decoder_ffn_dim"),
)?,
ln2: Norm::load(&st, &format!("{b}.final_layer_norm"), eps)?,
})
})
.collect::<Result<Vec<_>>>()?;
let (n_mels, units, rf) = (
g("num_mel_bins"),
g("speech_decoder_prenet_units"),
g("reduction_factor"),
);
let prenet = (0..g("speech_decoder_prenet_layers"))
.map(|i| {
Linear::load(
&st,
&format!("{dec}.prenet.layers.{i}"),
units,
if i == 0 { n_mels } else { units },
)
})
.collect::<Result<Vec<_>>>()?;
let post_units = g("speech_decoder_postnet_units");
let post_layers = g("speech_decoder_postnet_layers");
let postnet = (0..post_layers)
.map(|i| {
let b = format!("speech_decoder_postnet.layers.{i}");
let (in_ch, out_ch) = (
if i == 0 { n_mels } else { post_units },
if i == post_layers - 1 {
n_mels
} else {
post_units
},
);
let w = st.tensor_f32(&format!("{b}.conv.weight"))?;
let k = g("speech_decoder_postnet_kernel");
anyhow::ensure!(w.len() == out_ch * in_ch * k, "{b}.conv.weight shape");
Ok(PostnetLayer {
w,
bn_w: st.tensor_f32(&format!("{b}.batch_norm.weight"))?,
bn_b: st.tensor_f32(&format!("{b}.batch_norm.bias"))?,
bn_rm: st.tensor_f32(&format!("{b}.batch_norm.running_mean"))?,
bn_rv: st.tensor_f32(&format!("{b}.batch_norm.running_var"))?,
out_ch,
in_ch,
k,
tanh: i < post_layers - 1,
})
})
.collect::<Result<Vec<_>>>()?;
Ok(Self {
vocab,
unk_id: g("unk_token_id") as u32,
eos_id: g("eos_token_id") as u32,
embed_tokens: st.tensor_f32(&format!("{enc}.prenet.embed_tokens.weight"))?,
enc_alpha: st.tensor_f32(&format!("{enc}.prenet.encode_positions.alpha"))?[0],
enc_pe: st.tensor_f32("enc_pe")?,
enc_ln0: Norm::load(&st, &format!("{enc}.wrapped_encoder.layer_norm"), eps)?,
pe_k: st.tensor_f32(&format!(
"{enc}.wrapped_encoder.embed_positions.pe_k.weight"
))?,
max_rel: g("encoder_max_relative_position") as i64,
enc_blocks,
dec_alpha: st.tensor_f32(&format!("{dec}.prenet.encode_positions.alpha"))?[0],
dec_pe: st.tensor_f32("dec_pe")?,
prenet,
prenet_final: Linear::load(&st, &format!("{dec}.prenet.final_layer"), hidden, units)?,
speaker_layer: Linear::load(
&st,
&format!("{dec}.prenet.speaker_embeds_layer"),
hidden,
hidden + g("speaker_embedding_dim"),
)?,
dec_blocks,
feat_out: Linear::load(&st, "speech_decoder_postnet.feat_out", n_mels * rf, hidden)?,
prob_out: Linear::load(&st, "speech_decoder_postnet.prob_out", rf, hidden)?,
postnet,
hifigan: HifiGan::load(&st, v.get("hifigan").context("config hifigan")?)?,
n_mels,
reduction: rf,
hidden,
heads,
hd: hidden / heads,
sampling_rate: v
.get("hifigan")
.and_then(|h| h.get("sampling_rate"))
.and_then(|x| x.as_u64())
.unwrap_or(16000) as u32,
})
}
pub fn sampling_rate(&self) -> u32 {
self.sampling_rate
}
pub fn tokenize(&self, text: &str) -> Vec<u32> {
let sep = self.vocab.get("▁").copied().unwrap_or(self.unk_id);
let mut ids = Vec::with_capacity(text.len() + 2);
let mut pending_sep = true; for ch in text.chars() {
if ch.is_whitespace() {
pending_sep = true;
continue;
}
if pending_sep {
ids.push(sep);
pending_sep = false;
}
ids.push(
*self
.vocab
.get(ch.to_string().as_str())
.unwrap_or(&self.unk_id),
);
}
ids.push(self.eos_id); ids
}
fn enc_attention(&self, blk: &EncBlock, x: &[f32], t: usize) -> Vec<f32> {
let (h, heads, hd) = (self.hidden, self.heads, self.hd);
let q = blk.q.forward(x, t);
let k = blk.k.forward(x, t);
let v = blk.v.forward(x, t);
let scale = 1.0 / (hd as f32).sqrt();
let ctx: Vec<Vec<f32>> = (0..heads)
.into_par_iter()
.map(|head| {
let mut oh = vec![0f32; t * hd];
let mut srow = vec![0f32; t];
for i in 0..t {
let qi = &q[i * h + head * hd..][..hd];
for j in 0..t {
let kj = &k[j * h + head * hd..][..hd];
let rel = (i as i64 - j as i64).clamp(-self.max_rel, self.max_rel - 1)
+ self.max_rel;
let pe = &self.pe_k[rel as usize * hd..][..hd];
let mut dot = 0f32;
let mut bias = 0f32;
for c in 0..hd {
dot += qi[c] * kj[c];
bias += qi[c] * pe[c];
}
srow[j] = (dot + bias) * scale;
}
let max = srow.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for s in srow.iter_mut() {
*s = (*s - max).exp();
sum += *s;
}
let inv = 1.0 / sum;
let orow = &mut oh[i * hd..][..hd];
for j in 0..t {
let w = srow[j] * inv;
let vj = &v[j * h + head * hd..][..hd];
for c in 0..hd {
orow[c] += w * vj[c];
}
}
}
oh
})
.collect();
let mut merged = vec![0f32; t * h];
for (head, oh) in ctx.iter().enumerate() {
for i in 0..t {
merged[i * h + head * hd..][..hd].copy_from_slice(&oh[i * hd..][..hd]);
}
}
blk.out.forward(&merged, t)
}
pub fn encode_text(&self, ids: &[u32]) -> Vec<f32> {
let (h, t) = (self.hidden, ids.len());
let mut x = vec![0f32; t * h];
for (i, &id) in ids.iter().enumerate() {
let emb = &self.embed_tokens[id as usize * h..][..h];
let pe = &self.enc_pe[i * h..][..h];
for c in 0..h {
x[i * h + c] = emb[c] + self.enc_alpha * pe[c];
}
}
let mut x = self.enc_ln0.forward(&x, h);
for blk in &self.enc_blocks {
let a = self.enc_attention(blk, &x, t);
for (xi, ai) in x.iter_mut().zip(a) {
*xi += ai;
}
x = blk.ln1.forward(&x, h);
let mut mid = blk.fc1.forward(&x, t);
for m in mid.iter_mut() {
*m = gelu(*m);
}
let f = blk.fc2.forward(&mid, t);
for (xi, fi) in x.iter_mut().zip(f) {
*xi += fi;
}
x = blk.ln2.forward(&x, h);
}
x
}
fn attend_one(&self, q: &[f32], kc: &[f32], vc: &[f32], t: usize) -> Vec<f32> {
let (h, hd) = (self.hidden, self.hd);
let scale = 1.0 / (hd as f32).sqrt();
let mut out = vec![0f32; h];
out.par_chunks_mut(hd).enumerate().for_each(|(head, orow)| {
let qh = &q[head * hd..][..hd];
let mut srow = vec![0f32; t];
for j in 0..t {
let kj = &kc[j * h + head * hd..][..hd];
srow[j] = qh.iter().zip(kj).map(|(a, b)| a * b).sum::<f32>() * scale;
}
let max = srow.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for s in srow.iter_mut() {
*s = (*s - max).exp();
sum += *s;
}
let inv = 1.0 / sum;
for j in 0..t {
let w = srow[j] * inv;
let vj = &vc[j * h + head * hd..][..hd];
for c in 0..hd {
orow[c] += w * vj[c];
}
}
});
out
}
fn postnet_forward(&self, mel: &[f32], t: usize) -> Vec<f32> {
let nm = self.n_mels;
let mut x = vec![0f32; nm * t];
for i in 0..t {
for c in 0..nm {
x[c * t + i] = mel[i * nm + c];
}
}
for layer in &self.postnet {
let pad = (layer.k - 1) / 2;
let mut y = vec![0f32; layer.out_ch * t];
y.par_chunks_mut(t).enumerate().for_each(|(o, orow)| {
let inv = 1.0 / (layer.bn_rv[o] + 1e-5).sqrt(); for (ti, oval) in orow.iter_mut().enumerate() {
let mut acc = 0f32;
for ic in 0..layer.in_ch {
let xrow = &x[ic * t..][..t];
let wrow = &layer.w[(o * layer.in_ch + ic) * layer.k..][..layer.k];
for (kk, wv) in wrow.iter().enumerate() {
let src = ti + kk;
if src >= pad && src - pad < t {
acc += xrow[src - pad] * wv;
}
}
}
let bn = (acc - layer.bn_rm[o]) * inv * layer.bn_w[o] + layer.bn_b[o];
*oval = if layer.tanh { bn.tanh() } else { bn };
}
});
x = y;
}
let mut out = mel.to_vec();
for i in 0..t {
for c in 0..nm {
out[i * nm + c] += x[c * t + i];
}
}
out
}
pub fn decode_mel(&self, enc: &[f32], speaker: &[f32]) -> MelTrace {
let (h, nm, rf) = (self.hidden, self.n_mels, self.reduction);
let t_enc = enc.len() / h;
let maxlen = t_enc * 20 / rf;
let norm = speaker.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
let spk: Vec<f32> = speaker.iter().map(|v| v / norm).collect();
let cross_kv: Vec<(Vec<f32>, Vec<f32>)> = self
.dec_blocks
.iter()
.map(|b| (b.ca_k.forward(enc, t_enc), b.ca_v.forward(enc, t_enc)))
.collect();
let mut sa_k: Vec<Vec<f32>> = vec![Vec::new(); self.dec_blocks.len()];
let mut sa_v: Vec<Vec<f32>> = vec![Vec::new(); self.dec_blocks.len()];
let mut tr = MelTrace {
prenet: Vec::new(),
dec_h: Vec::new(),
stop_logits: Vec::new(),
mel_pre: Vec::new(),
mel_post: Vec::new(),
steps: 0,
};
let mut prev = vec![0f32; nm]; for step in 0..maxlen {
let mut x = prev.clone();
for l in &self.prenet {
x = l.forward(&x, 1);
for v in x.iter_mut() {
*v = v.max(0.0);
}
}
let mut x = self.prenet_final.forward(&x, 1);
for (c, xv) in x.iter_mut().enumerate() {
*xv += self.dec_alpha * self.dec_pe[step * h + c];
}
x.extend_from_slice(&spk);
let mut x = self.speaker_layer.forward(&x, 1);
for v in x.iter_mut() {
*v = v.max(0.0);
}
tr.prenet.extend_from_slice(&x);
for (li, blk) in self.dec_blocks.iter().enumerate() {
let q = blk.sa_q.forward(&x, 1);
sa_k[li].extend(blk.sa_k.forward(&x, 1));
sa_v[li].extend(blk.sa_v.forward(&x, 1));
let a = self.attend_one(&q, &sa_k[li], &sa_v[li], step + 1);
let a = blk.sa_out.forward(&a, 1);
for (xi, ai) in x.iter_mut().zip(a) {
*xi += ai;
}
x = blk.sa_ln.forward(&x, h);
let q = blk.ca_q.forward(&x, 1);
let a = self.attend_one(&q, &cross_kv[li].0, &cross_kv[li].1, t_enc);
let a = blk.ca_out.forward(&a, 1);
for (xi, ai) in x.iter_mut().zip(a) {
*xi += ai;
}
x = blk.ca_ln.forward(&x, h);
let mut mid = blk.fc1.forward(&x, 1);
for m in mid.iter_mut() {
*m = gelu(*m);
}
let f = blk.fc2.forward(&mid, 1);
for (xi, fi) in x.iter_mut().zip(f) {
*xi += fi;
}
x = blk.ln2.forward(&x, h);
}
tr.dec_h.extend_from_slice(&x);
let feat = self.feat_out.forward(&x, 1); let stop = self.prob_out.forward(&x, 1); tr.mel_pre.extend_from_slice(&feat);
tr.stop_logits.extend_from_slice(&stop);
prev.copy_from_slice(&feat[(rf - 1) * nm..]); tr.steps = step + 1;
let prob_sum: f32 = stop.iter().map(|l| 1.0 / (1.0 + (-l).exp())).sum();
if prob_sum >= 0.5 {
break;
}
}
tr.mel_post = self.postnet_forward(&tr.mel_pre, tr.steps * rf);
tr
}
pub fn vocode(&self, mel: &[f32]) -> Vec<f32> {
self.hifigan.forward(mel, mel.len() / self.n_mels)
}
pub fn synthesize(&self, text: &str, speaker: &[f32]) -> Vec<f32> {
let ids = self.tokenize(text);
let enc = self.encode_text(&ids);
let tr = self.decode_mel(&enc, speaker);
self.vocode(&tr.mel_post)
}
}