1119 lines
52 KiB
Python
1119 lines
52 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import math
|
|
import random
|
|
import sys
|
|
import time
|
|
from dataclasses import asdict, dataclass
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
from torch import nn
|
|
import torch.nn.functional as F
|
|
import torchaudio
|
|
|
|
SCRIPT_ROOT = Path(__file__).resolve().parent
|
|
PROJECT_ROOT = SCRIPT_ROOT.parents[0]
|
|
TINY_ROOT = PROJECT_ROOT / "third_party" / "tiny-tts"
|
|
sys.path = [str(TINY_ROOT), str(SCRIPT_ROOT)] + [p for p in sys.path if p]
|
|
|
|
from train_hifigan_oracle_v1 import HifiGanConfig, HifiGanGenerator, MelFrontend
|
|
|
|
|
|
@dataclass
|
|
class MicroFastSpeechConfig:
|
|
vocab_size: int = 256
|
|
tone_size: int = 16
|
|
lang_size: int = 4
|
|
n_mels: int = 80
|
|
hidden: int = 168
|
|
encoder_layers: int = 5
|
|
decoder_layers: int = 6
|
|
decoder_ff_mult: int = 3
|
|
kernel_size: int = 7
|
|
speaker_count: int = 2
|
|
speaker_dim: int = 64
|
|
dropout: float = 0.08
|
|
sample_rate: int = 24000
|
|
max_frames: int = 1400
|
|
postnet_scale: float = 0.10
|
|
use_frame_pitch: bool = True
|
|
abs_frame_bins: int = 512
|
|
use_contextual_predictors: bool = False
|
|
use_group_duration_planner: bool = False
|
|
|
|
|
|
def count_parameters(model: nn.Module) -> int:
|
|
return sum(p.numel() for p in model.parameters())
|
|
|
|
|
|
def load_rows(path: Path, max_rows: int = 0) -> list[dict]:
|
|
rows = []
|
|
with path.open("r", encoding="utf-8") as f:
|
|
for line in f:
|
|
if line.strip():
|
|
row = json.loads(line)
|
|
if Path(str(row.get("target_audio") or "")).is_file():
|
|
rows.append(row)
|
|
if max_rows and len(rows) >= max_rows:
|
|
break
|
|
if not rows:
|
|
raise RuntimeError(f"No usable rows in {path}")
|
|
return rows
|
|
|
|
|
|
def load_audio(path: str, sample_rate: int, max_seconds: float) -> torch.Tensor:
|
|
wav, sr = torchaudio.load(path)
|
|
if wav.shape[0] > 1:
|
|
wav = wav.mean(dim=0, keepdim=True)
|
|
if sr != sample_rate:
|
|
wav = torchaudio.functional.resample(wav, sr, sample_rate)
|
|
return wav[:, : int(sample_rate * max_seconds)].squeeze(0).clamp(-1.0, 1.0)
|
|
|
|
|
|
def fit_durations(durations: list[int], target_frames: int) -> list[int]:
|
|
if sum(durations) == target_frames:
|
|
return list(durations)
|
|
total = max(1, sum(durations))
|
|
raw = [max(0.0, d * target_frames / total) for d in durations]
|
|
out = [int(math.floor(x)) for x in raw]
|
|
order = sorted(((raw[i] - out[i], i) for i in range(len(out))), reverse=True)
|
|
for _, idx in order[: max(0, target_frames - sum(out))]:
|
|
out[idx] += 1
|
|
while sum(out) > target_frames:
|
|
idx = max(range(len(out)), key=lambda i: out[i])
|
|
out[idx] -= 1
|
|
return out
|
|
|
|
|
|
def pad_1d(items: list[torch.Tensor], value: float = 0.0) -> torch.Tensor:
|
|
max_len = max(x.numel() for x in items)
|
|
out = torch.full((len(items), max_len), value, dtype=items[0].dtype)
|
|
for i, item in enumerate(items):
|
|
out[i, : item.numel()] = item
|
|
return out
|
|
|
|
|
|
def pad_2d(items: list[torch.Tensor], value: float = 0.0) -> torch.Tensor:
|
|
max_len = max(x.shape[0] for x in items)
|
|
dim = items[0].shape[1]
|
|
out = torch.full((len(items), max_len, dim), value, dtype=items[0].dtype)
|
|
for i, item in enumerate(items):
|
|
out[i, : item.shape[0]] = item
|
|
return out
|
|
|
|
|
|
def pad_mels(items: list[torch.Tensor]) -> tuple[torch.Tensor, torch.Tensor]:
|
|
max_len = max(x.shape[-1] for x in items)
|
|
n_mels = items[0].shape[0]
|
|
out = torch.zeros(len(items), n_mels, max_len, dtype=items[0].dtype)
|
|
mask = torch.zeros(len(items), max_len, dtype=torch.bool)
|
|
for i, mel in enumerate(items):
|
|
frames = mel.shape[-1]
|
|
out[i, :, :frames] = mel
|
|
mask[i, :frames] = True
|
|
return out, mask
|
|
|
|
|
|
def pad_wavs(items: list[torch.Tensor], frames: list[int], hop_size: int) -> torch.Tensor:
|
|
max_len = max(max(1, int(frame_count)) * hop_size for frame_count in frames)
|
|
out = torch.zeros(len(items), max_len, dtype=items[0].dtype)
|
|
for i, (wav, frame_count) in enumerate(zip(items, frames)):
|
|
length = max(1, int(frame_count)) * hop_size
|
|
cropped = wav[:length]
|
|
out[i, : cropped.numel()] = cropped
|
|
return out
|
|
|
|
|
|
def aggregate_token_features(mel: torch.Tensor, durations: list[int]) -> tuple[torch.Tensor, torch.Tensor]:
|
|
# mel: [80, frames], log-mel from the exact V2+ frontend.
|
|
frames = mel.shape[-1]
|
|
amp = torch.exp(mel).clamp_min(1e-5)
|
|
energy_frame = mel.mean(dim=0)
|
|
bins = torch.linspace(0.0, 1.0, mel.shape[0], device=mel.device).view(-1, 1)
|
|
bright_frame = (amp * bins).sum(dim=0) / amp.sum(dim=0).clamp_min(1e-5)
|
|
energies = []
|
|
brights = []
|
|
pos = 0
|
|
for dur in durations:
|
|
end = min(frames, pos + max(0, int(dur)))
|
|
if end > pos:
|
|
energies.append(energy_frame[pos:end].mean())
|
|
brights.append(bright_frame[pos:end].mean())
|
|
else:
|
|
energies.append(torch.zeros((), device=mel.device, dtype=mel.dtype))
|
|
brights.append(torch.zeros((), device=mel.device, dtype=mel.dtype))
|
|
pos = end
|
|
return torch.stack(energies), torch.stack(brights)
|
|
|
|
|
|
def aggregate_token_pitch(pitch_frame: torch.Tensor, durations: list[int]) -> torch.Tensor:
|
|
# pitch_frame: [2, frames] with normalized log-f0 and voiced flag.
|
|
frames = pitch_frame.shape[-1]
|
|
out = []
|
|
pos = 0
|
|
for dur in durations:
|
|
end = min(frames, pos + max(0, int(dur)))
|
|
if end > pos:
|
|
span = pitch_frame[:, pos:end]
|
|
voiced = span[1].mean()
|
|
voiced_mask = span[1] > 0.5
|
|
if bool(voiced_mask.any()):
|
|
log_f0 = span[0, voiced_mask].mean()
|
|
else:
|
|
log_f0 = torch.zeros((), dtype=pitch_frame.dtype)
|
|
out.append(torch.stack([log_f0, voiced]))
|
|
else:
|
|
out.append(torch.zeros(2, dtype=pitch_frame.dtype))
|
|
pos = end
|
|
return torch.stack(out, dim=0)
|
|
|
|
|
|
def extract_pitch_features(wav: torch.Tensor, sample_rate: int, frames: int) -> torch.Tensor:
|
|
# Returns [2, frames]: normalized log-f0 and voiced flag. The detector can
|
|
# produce octave spikes, so clip to speech range and median-smooth lightly.
|
|
pitch = torchaudio.functional.detect_pitch_frequency(
|
|
wav.unsqueeze(0).cpu(),
|
|
sample_rate,
|
|
frame_time=256 / sample_rate,
|
|
).squeeze(0)
|
|
if pitch.numel() < frames:
|
|
pitch = F.pad(pitch, (0, frames - pitch.numel()), value=0.0)
|
|
pitch = pitch[:frames]
|
|
voiced = ((pitch >= 55.0) & (pitch <= 420.0)).float()
|
|
pitch = pitch.clamp(55.0, 420.0)
|
|
# Median filter over 5 frames to reduce spurious jumps.
|
|
if pitch.numel() >= 5:
|
|
padded = F.pad(pitch.view(1, 1, -1), (2, 2), mode="replicate")
|
|
windows = padded.unfold(-1, 5, 1).squeeze(0).squeeze(0)
|
|
pitch = windows.median(dim=-1).values
|
|
log_f0 = (torch.log(pitch) - math.log(140.0)) / 0.45
|
|
log_f0 = log_f0.clamp(-3.0, 3.0) * voiced
|
|
return torch.stack([log_f0, voiced], dim=0)
|
|
|
|
|
|
class ConvFFNBlock(nn.Module):
|
|
def __init__(self, hidden: int, kernel_size: int, dropout: float, ff_mult: int = 4) -> None:
|
|
super().__init__()
|
|
pad = kernel_size // 2
|
|
self.norm1 = nn.LayerNorm(hidden)
|
|
self.depth = nn.Conv1d(hidden, hidden * 2, kernel_size, padding=pad, groups=hidden)
|
|
self.point = nn.Conv1d(hidden, hidden, 1)
|
|
self.drop = nn.Dropout(dropout)
|
|
self.norm2 = nn.LayerNorm(hidden)
|
|
self.ff = nn.Sequential(
|
|
nn.Linear(hidden, hidden * ff_mult),
|
|
nn.SiLU(),
|
|
nn.Dropout(dropout),
|
|
nn.Linear(hidden * ff_mult, hidden),
|
|
)
|
|
|
|
def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None) -> torch.Tensor:
|
|
y = self.norm1(x).transpose(1, 2)
|
|
a, b = self.depth(y).chunk(2, dim=1)
|
|
y = self.point(a * torch.sigmoid(b)).transpose(1, 2)
|
|
x = x + self.drop(y)
|
|
x = x + self.drop(self.ff(self.norm2(x)))
|
|
if mask is not None:
|
|
x = x * mask.unsqueeze(-1)
|
|
return x
|
|
|
|
|
|
class MicroFastSpeech(nn.Module):
|
|
def __init__(self, cfg: MicroFastSpeechConfig) -> None:
|
|
super().__init__()
|
|
self.cfg = cfg
|
|
# Phone id 0 is a real inserted blank/silence token from TinyTTS, not
|
|
# padding. Padding is tracked by duration masks instead.
|
|
self.phone = nn.Embedding(cfg.vocab_size, cfg.hidden)
|
|
self.tone = nn.Embedding(cfg.tone_size, cfg.hidden)
|
|
self.lang = nn.Embedding(cfg.lang_size, cfg.hidden)
|
|
self.speaker = nn.Embedding(cfg.speaker_count, cfg.speaker_dim)
|
|
self.speaker_proj = nn.Linear(cfg.speaker_dim, cfg.hidden)
|
|
self.encoder = nn.ModuleList([ConvFFNBlock(cfg.hidden, cfg.kernel_size, cfg.dropout) for _ in range(cfg.encoder_layers)])
|
|
self.duration_head = nn.Sequential(nn.LayerNorm(cfg.hidden), nn.Linear(cfg.hidden, cfg.hidden), nn.SiLU(), nn.Linear(cfg.hidden, 1))
|
|
self.energy_head = nn.Sequential(nn.LayerNorm(cfg.hidden), nn.Linear(cfg.hidden, cfg.hidden // 2), nn.SiLU(), nn.Linear(cfg.hidden // 2, 1))
|
|
self.bright_head = nn.Sequential(nn.LayerNorm(cfg.hidden), nn.Linear(cfg.hidden, cfg.hidden // 2), nn.SiLU(), nn.Linear(cfg.hidden // 2, 1))
|
|
self.pitch_head = nn.Sequential(nn.LayerNorm(cfg.hidden), nn.Linear(cfg.hidden, cfg.hidden), nn.SiLU(), nn.Linear(cfg.hidden, 2))
|
|
self.group_duration_delta = nn.Linear(cfg.hidden, 1) if cfg.use_group_duration_planner else None
|
|
if self.group_duration_delta is not None:
|
|
nn.init.zeros_(self.group_duration_delta.weight)
|
|
nn.init.zeros_(self.group_duration_delta.bias)
|
|
self.predictor_context = (
|
|
ConvFFNBlock(cfg.hidden, 5, cfg.dropout, 2) if cfg.use_contextual_predictors else nn.Identity()
|
|
)
|
|
self.duration_delta = nn.Linear(cfg.hidden, 1) if cfg.use_contextual_predictors else None
|
|
self.energy_delta = nn.Linear(cfg.hidden, 1) if cfg.use_contextual_predictors else None
|
|
self.bright_delta = nn.Linear(cfg.hidden, 1) if cfg.use_contextual_predictors else None
|
|
self.pitch_delta = nn.Linear(cfg.hidden, 2) if cfg.use_contextual_predictors else None
|
|
if cfg.use_contextual_predictors:
|
|
for layer in (self.duration_delta, self.energy_delta, self.bright_delta, self.pitch_delta):
|
|
nn.init.zeros_(layer.weight)
|
|
nn.init.zeros_(layer.bias)
|
|
self.energy_proj = nn.Linear(1, cfg.hidden)
|
|
self.bright_proj = nn.Linear(1, cfg.hidden)
|
|
self.pitch_proj = nn.Sequential(nn.Linear(2, cfg.hidden), nn.SiLU(), nn.Linear(cfg.hidden, cfg.hidden))
|
|
self.abs_frame = nn.Embedding(cfg.abs_frame_bins, cfg.hidden)
|
|
self.frame_proj = nn.Sequential(nn.Linear(8, cfg.hidden), nn.SiLU(), nn.Linear(cfg.hidden, cfg.hidden))
|
|
self.local_ctx = nn.Sequential(
|
|
nn.Linear(cfg.hidden * 3, cfg.hidden * 2),
|
|
nn.SiLU(),
|
|
nn.Linear(cfg.hidden * 2, cfg.hidden),
|
|
)
|
|
self.decoder = nn.ModuleList([ConvFFNBlock(cfg.hidden, cfg.kernel_size, cfg.dropout, cfg.decoder_ff_mult) for _ in range(cfg.decoder_layers)])
|
|
self.frame_gru = nn.GRU(cfg.hidden, cfg.hidden // 2, num_layers=1, batch_first=True, bidirectional=True)
|
|
self.mel_head = nn.Sequential(nn.LayerNorm(cfg.hidden), nn.Linear(cfg.hidden, cfg.hidden), nn.SiLU(), nn.Linear(cfg.hidden, cfg.n_mels))
|
|
self.postnet = nn.Sequential(
|
|
nn.Conv1d(cfg.n_mels, cfg.hidden, 5, padding=2),
|
|
nn.Tanh(),
|
|
nn.Conv1d(cfg.hidden, cfg.hidden, 5, padding=2),
|
|
nn.Tanh(),
|
|
nn.Conv1d(cfg.hidden, cfg.n_mels, 5, padding=2),
|
|
)
|
|
|
|
def encode(self, phone: torch.Tensor, tone: torch.Tensor, lang: torch.Tensor, speaker: torch.Tensor, token_mask: torch.Tensor) -> torch.Tensor:
|
|
x = self.phone(phone) + self.tone(tone.clamp_max(self.cfg.tone_size - 1)) + self.lang(lang.clamp_max(self.cfg.lang_size - 1))
|
|
x = x + self.speaker_proj(self.speaker(speaker)).unsqueeze(1)
|
|
x = x * token_mask.unsqueeze(-1)
|
|
for block in self.encoder:
|
|
x = block(x, token_mask)
|
|
return x
|
|
|
|
def regulate(self, encoded: torch.Tensor, durations: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
batch_frames = []
|
|
batch_meta = []
|
|
lengths = []
|
|
device = encoded.device
|
|
for b in range(encoded.shape[0]):
|
|
reps = []
|
|
meta = []
|
|
durs = durations[b].long().clamp_min(0)
|
|
token_count = max(1, int((durs > 0).sum().item()))
|
|
for i, dur_t in enumerate(durs.tolist()):
|
|
dur = int(dur_t)
|
|
if dur <= 0:
|
|
continue
|
|
reps.append(encoded[b, i].view(1, -1).expand(dur, -1))
|
|
rel = torch.linspace(0.0, 1.0, dur, device=device)
|
|
token_pos = torch.full((dur,), i / max(1, token_count - 1), device=device)
|
|
log_dur = torch.full((dur,), math.log1p(dur) / 6.0, device=device)
|
|
inv_rel = 1.0 - rel
|
|
center = 1.0 - torch.abs(rel * 2.0 - 1.0)
|
|
meta.append(
|
|
torch.stack(
|
|
[
|
|
rel,
|
|
inv_rel,
|
|
center,
|
|
torch.sin(rel * math.pi),
|
|
torch.cos(rel * math.pi),
|
|
token_pos,
|
|
log_dur,
|
|
torch.full_like(rel, dur / 40.0),
|
|
],
|
|
dim=-1,
|
|
)
|
|
)
|
|
if reps:
|
|
frames = torch.cat(reps, dim=0)
|
|
frame_meta = torch.cat(meta, dim=0)
|
|
else:
|
|
frames = encoded[b, :1]
|
|
frame_meta = torch.zeros(1, 8, device=device)
|
|
batch_frames.append(frames[: self.cfg.max_frames])
|
|
batch_meta.append(frame_meta[: self.cfg.max_frames])
|
|
lengths.append(min(frames.shape[0], self.cfg.max_frames))
|
|
max_len = max(lengths)
|
|
out = torch.zeros(encoded.shape[0], max_len, encoded.shape[-1], device=device)
|
|
meta_out = torch.zeros(encoded.shape[0], max_len, 8, device=device)
|
|
mask = torch.zeros(encoded.shape[0], max_len, dtype=torch.bool, device=device)
|
|
for b, frames in enumerate(batch_frames):
|
|
n = min(frames.shape[0], max_len)
|
|
out[b, :n] = frames[:n]
|
|
meta_out[b, :n] = batch_meta[b][:n]
|
|
mask[b, :n] = True
|
|
return out, meta_out, mask
|
|
|
|
def add_local_context(self, encoded: torch.Tensor, durations: torch.Tensor) -> torch.Tensor:
|
|
device = encoded.device
|
|
batch_frames = []
|
|
for b in range(encoded.shape[0]):
|
|
reps = []
|
|
durs = durations[b].long().clamp_min(0)
|
|
for i, dur_t in enumerate(durs.tolist()):
|
|
dur = int(dur_t)
|
|
if dur <= 0:
|
|
continue
|
|
prev_i = max(0, i - 1)
|
|
next_i = min(encoded.shape[1] - 1, i + 1)
|
|
ctx = torch.cat([encoded[b, prev_i], encoded[b, i], encoded[b, next_i]], dim=-1)
|
|
reps.append(ctx.view(1, -1).expand(dur, -1))
|
|
if reps:
|
|
frames = torch.cat(reps, dim=0)
|
|
else:
|
|
frames = torch.zeros(1, encoded.shape[-1] * 3, device=device)
|
|
batch_frames.append(frames[: self.cfg.max_frames])
|
|
max_len = max(x.shape[0] for x in batch_frames)
|
|
ctx_out = torch.zeros(encoded.shape[0], max_len, encoded.shape[-1] * 3, device=device)
|
|
for b, frames in enumerate(batch_frames):
|
|
ctx_out[b, : frames.shape[0]] = frames
|
|
return self.local_ctx(ctx_out)
|
|
|
|
def expand_token_feature(self, feature: torch.Tensor, durations: torch.Tensor) -> torch.Tensor:
|
|
device = feature.device
|
|
batch_frames = []
|
|
for b in range(feature.shape[0]):
|
|
reps = []
|
|
durs = durations[b].long().clamp_min(0)
|
|
for i, dur_t in enumerate(durs.tolist()):
|
|
dur = int(dur_t)
|
|
if dur <= 0:
|
|
continue
|
|
reps.append(feature[b, i].view(1, -1).expand(dur, -1))
|
|
if reps:
|
|
frames = torch.cat(reps, dim=0)
|
|
else:
|
|
frames = torch.zeros(1, feature.shape[-1], device=device)
|
|
batch_frames.append(frames[: self.cfg.max_frames])
|
|
max_len = max(x.shape[0] for x in batch_frames)
|
|
out = torch.zeros(feature.shape[0], max_len, feature.shape[-1], device=device)
|
|
for b, frames in enumerate(batch_frames):
|
|
out[b, : frames.shape[0]] = frames
|
|
return out
|
|
|
|
def forward(
|
|
self,
|
|
phone: torch.Tensor,
|
|
tone: torch.Tensor,
|
|
lang: torch.Tensor,
|
|
speaker: torch.Tensor,
|
|
durations: torch.Tensor,
|
|
energy_target: torch.Tensor | None = None,
|
|
bright_target: torch.Tensor | None = None,
|
|
pitch_frame: torch.Tensor | None = None,
|
|
predicted_prosody_mix: float = 0.0,
|
|
detach_mixed_predictions: bool = True,
|
|
) -> dict[str, torch.Tensor]:
|
|
token_mask = durations.gt(0)
|
|
encoded = self.encode(phone, tone, lang, speaker, token_mask)
|
|
log_dur, energy_pred, bright_pred, pitch_pred = self.predict_prosody(encoded, token_mask)
|
|
mixed_energy_pred = energy_pred.detach() if detach_mixed_predictions else energy_pred
|
|
mixed_bright_pred = bright_pred.detach() if detach_mixed_predictions else bright_pred
|
|
if energy_target is not None:
|
|
energy = torch.lerp(energy_target, mixed_energy_pred, predicted_prosody_mix)
|
|
else:
|
|
energy = energy_pred
|
|
if bright_target is not None:
|
|
bright = torch.lerp(bright_target, mixed_bright_pred, predicted_prosody_mix)
|
|
else:
|
|
bright = bright_pred
|
|
conditioned = encoded + self.energy_proj(energy.unsqueeze(-1)) + self.bright_proj(bright.unsqueeze(-1))
|
|
frames, frame_meta, frame_mask = self.regulate(conditioned, durations)
|
|
x = frames + self.frame_proj(frame_meta) + self.add_local_context(conditioned, durations)
|
|
pos = torch.arange(x.shape[1], device=x.device)
|
|
pos = torch.div(pos * self.cfg.abs_frame_bins, max(1, self.cfg.max_frames), rounding_mode="floor").clamp_max(
|
|
self.cfg.abs_frame_bins - 1
|
|
)
|
|
x = x + self.abs_frame(pos).unsqueeze(0)
|
|
if self.cfg.use_frame_pitch:
|
|
if pitch_frame is not None:
|
|
pitch_frame = pitch_frame[:, :, : x.shape[1]].transpose(1, 2)
|
|
if pitch_frame.shape[1] < x.shape[1]:
|
|
pitch_frame = F.pad(pitch_frame, (0, 0, 0, x.shape[1] - pitch_frame.shape[1]))
|
|
if predicted_prosody_mix > 0.0:
|
|
mixed_pitch_pred = pitch_pred.detach() if detach_mixed_predictions else pitch_pred
|
|
predicted_pitch_frame = self.expand_token_feature(mixed_pitch_pred, durations)[:, : x.shape[1]]
|
|
pitch_frame = torch.lerp(pitch_frame, predicted_pitch_frame, predicted_prosody_mix)
|
|
else:
|
|
pitch_frame = self.expand_token_feature(pitch_pred, durations)[:, : x.shape[1]]
|
|
x = x + self.pitch_proj(pitch_frame)
|
|
for block in self.decoder:
|
|
x = block(x, frame_mask)
|
|
x = x + self.frame_gru(x)[0]
|
|
mel = self.mel_head(x).transpose(1, 2)
|
|
mel = mel + self.cfg.postnet_scale * self.postnet(mel)
|
|
group_log_dur, group_mask = self.group_log_durations(phone, log_dur, encoded)
|
|
return {
|
|
"mel": mel,
|
|
"frame_mask": frame_mask,
|
|
"log_dur": log_dur,
|
|
"group_log_dur": group_log_dur,
|
|
"group_mask": group_mask,
|
|
"energy": energy_pred,
|
|
"bright": bright_pred,
|
|
"pitch": pitch_pred,
|
|
"token_mask": token_mask,
|
|
}
|
|
|
|
def predict_prosody(
|
|
self, encoded: torch.Tensor, token_mask: torch.Tensor
|
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
log_dur = self.duration_head(encoded).squeeze(-1)
|
|
energy = self.energy_head(encoded).squeeze(-1)
|
|
bright = self.bright_head(encoded).squeeze(-1)
|
|
pitch = self.pitch_head(encoded)
|
|
if self.cfg.use_contextual_predictors:
|
|
context = self.predictor_context(encoded, token_mask)
|
|
log_dur = log_dur + self.duration_delta(context).squeeze(-1)
|
|
energy = energy + self.energy_delta(context).squeeze(-1)
|
|
bright = bright + self.bright_delta(context).squeeze(-1)
|
|
pitch = pitch + self.pitch_delta(context)
|
|
return log_dur, energy, bright, pitch
|
|
|
|
def group_log_durations(
|
|
self, phone: torch.Tensor, log_dur: torch.Tensor, encoded: torch.Tensor
|
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
"""Predict stable blank-plus-phone region durations at visible phones."""
|
|
base_dur = torch.expm1(log_dur).clamp_min(0.05)
|
|
grouped = torch.zeros_like(log_dur)
|
|
group_mask = torch.zeros_like(phone, dtype=torch.bool)
|
|
delta = self.group_duration_delta(encoded).squeeze(-1) if self.group_duration_delta is not None else None
|
|
for batch_index in range(phone.shape[0]):
|
|
pending: list[torch.Tensor] = []
|
|
last_visible: int | None = None
|
|
for token_index in range(phone.shape[1]):
|
|
pending.append(base_dur[batch_index, token_index])
|
|
if int(phone[batch_index, token_index].item()) != 0:
|
|
value = torch.stack(pending).sum()
|
|
if delta is not None:
|
|
value = value * torch.exp(delta[batch_index, token_index].clamp(-1.5, 1.5))
|
|
grouped[batch_index, token_index] = torch.log1p(value)
|
|
group_mask[batch_index, token_index] = True
|
|
pending = []
|
|
last_visible = token_index
|
|
if pending and last_visible is not None:
|
|
value = torch.expm1(grouped[batch_index, last_visible]) + torch.stack(pending).sum()
|
|
grouped[batch_index, last_visible] = torch.log1p(value)
|
|
return grouped, group_mask
|
|
|
|
def apply_group_duration_plan(
|
|
self, phone: torch.Tensor, log_dur: torch.Tensor, encoded: torch.Tensor, length_scale: float, max_duration: int
|
|
) -> torch.Tensor:
|
|
base = torch.expm1(log_dur).clamp_min(0.05)
|
|
group_log, _ = self.group_log_durations(phone, log_dur, encoded)
|
|
planned = torch.zeros_like(base, dtype=torch.long)
|
|
for batch_index in range(phone.shape[0]):
|
|
pending: list[int] = []
|
|
last_visible: int | None = None
|
|
for token_index in range(phone.shape[1]):
|
|
pending.append(token_index)
|
|
if int(phone[batch_index, token_index].item()) != 0:
|
|
target = max(len(pending), int(round(float(torch.expm1(group_log[batch_index, token_index]) * length_scale))))
|
|
weights = base[batch_index, pending]
|
|
remaining = target - len(pending)
|
|
raw = weights / weights.sum().clamp_min(1e-6) * remaining
|
|
allocated = torch.ones_like(raw, dtype=torch.long) + torch.floor(raw).long()
|
|
remainder = target - int(allocated.sum().item())
|
|
if remainder > 0:
|
|
order = torch.argsort(raw - torch.floor(raw), descending=True)
|
|
allocated[order[:remainder]] += 1
|
|
planned[batch_index, pending] = allocated
|
|
pending = []
|
|
last_visible = token_index
|
|
if pending and last_visible is not None:
|
|
planned[batch_index, last_visible] += max(1, int(round(float(base[batch_index, pending].sum() * length_scale))))
|
|
return planned.clamp(0, max_duration)
|
|
|
|
@torch.no_grad()
|
|
def infer(
|
|
self,
|
|
phone: torch.Tensor,
|
|
tone: torch.Tensor,
|
|
lang: torch.Tensor,
|
|
speaker: torch.Tensor,
|
|
length_scale: float = 1.0,
|
|
min_duration: int = 1,
|
|
max_duration: int = 80,
|
|
pitch_scale: float = 1.0,
|
|
energy_scale: float = 1.0,
|
|
smooth_predictors: bool = False,
|
|
) -> torch.Tensor:
|
|
# In single-sample inference there is no padded tail; id 0 remains the
|
|
# explicit blank/pause token and must keep duration.
|
|
token_mask = torch.ones_like(phone, dtype=torch.bool)
|
|
encoded = self.encode(phone, tone, lang, speaker, token_mask)
|
|
log_dur, energy, bright, pitch = self.predict_prosody(encoded, token_mask)
|
|
if self.group_duration_delta is not None:
|
|
durations = self.apply_group_duration_plan(phone, log_dur, encoded, length_scale, max_duration)
|
|
durations = durations.masked_fill(~token_mask, 0)
|
|
else:
|
|
pred_dur = torch.expm1(log_dur).clamp(0, max_duration) * length_scale
|
|
durations = torch.round(pred_dur).long().clamp_min(min_duration).masked_fill(~token_mask, 0)
|
|
energy = energy * energy_scale
|
|
pitch = torch.stack([pitch[..., 0] * pitch_scale, pitch[..., 1].clamp(0.0, 1.0)], dim=-1)
|
|
if smooth_predictors and phone.shape[1] >= 3:
|
|
energy = F.avg_pool1d(energy.unsqueeze(1), 3, stride=1, padding=1).squeeze(1)
|
|
bright = F.avg_pool1d(bright.unsqueeze(1), 3, stride=1, padding=1).squeeze(1)
|
|
pitch_t = pitch.transpose(1, 2)
|
|
pitch = F.avg_pool1d(pitch_t, 3, stride=1, padding=1).transpose(1, 2)
|
|
conditioned = encoded + self.energy_proj(energy.unsqueeze(-1)) + self.bright_proj(bright.unsqueeze(-1))
|
|
frames, frame_meta, frame_mask = self.regulate(conditioned, durations)
|
|
x = frames + self.frame_proj(frame_meta) + self.add_local_context(conditioned, durations)
|
|
pos = torch.arange(x.shape[1], device=x.device)
|
|
pos = torch.div(pos * self.cfg.abs_frame_bins, max(1, self.cfg.max_frames), rounding_mode="floor").clamp_max(
|
|
self.cfg.abs_frame_bins - 1
|
|
)
|
|
x = x + self.abs_frame(pos).unsqueeze(0)
|
|
if self.cfg.use_frame_pitch:
|
|
pitch_frame = self.expand_token_feature(pitch, durations)[:, : x.shape[1]]
|
|
x = x + self.pitch_proj(pitch_frame)
|
|
for block in self.decoder:
|
|
x = block(x, frame_mask)
|
|
x = x + self.frame_gru(x)[0]
|
|
mel = self.mel_head(x).transpose(1, 2)
|
|
mel = mel + self.cfg.postnet_scale * self.postnet(mel)
|
|
return mel
|
|
|
|
|
|
def collate(batch: list[dict], cfg: MicroFastSpeechConfig, mel_frontend: MelFrontend, device: torch.device, max_seconds: float, hop_size: int):
|
|
phones = [torch.LongTensor(x["phone_ids"]) for x in batch]
|
|
tones = [torch.LongTensor(x["tone_ids"]) for x in batch]
|
|
langs = [torch.LongTensor(x["lang_ids"]) for x in batch]
|
|
durations_raw = [list(map(int, x["hifigan_durations"])) for x in batch]
|
|
speakers = torch.LongTensor([int(x["speaker_id"]) for x in batch])
|
|
phone = pad_1d(phones, 0).long()
|
|
tone = pad_1d(tones, 0).long()
|
|
lang = pad_1d(langs, 0).long()
|
|
mels = []
|
|
durations = []
|
|
energies = []
|
|
brights = []
|
|
pitches = []
|
|
token_pitches = []
|
|
wavs = []
|
|
frame_counts = []
|
|
with torch.no_grad():
|
|
for row, dur in zip(batch, durations_raw):
|
|
wav_1d = load_audio(str(row["target_audio"]), cfg.sample_rate, max_seconds)
|
|
wav = wav_1d.unsqueeze(0).to(device)
|
|
mel = mel_frontend(wav).squeeze(0).detach().cpu()
|
|
dur = fit_durations(dur[: len(row["phone_ids"])], min(mel.shape[-1], cfg.max_frames))
|
|
mel = mel[:, : sum(dur)]
|
|
energy, bright = aggregate_token_features(mel, dur)
|
|
pitch = extract_pitch_features(wav_1d, cfg.sample_rate, mel.shape[-1])
|
|
token_pitch = aggregate_token_pitch(pitch, dur)
|
|
mels.append(mel)
|
|
durations.append(torch.LongTensor(dur))
|
|
energies.append(energy)
|
|
brights.append(bright)
|
|
pitches.append(pitch)
|
|
token_pitches.append(token_pitch)
|
|
wavs.append(wav_1d)
|
|
frame_counts.append(mel.shape[-1])
|
|
duration = pad_1d(durations, 0).long()
|
|
energy = pad_1d(energies, 0.0).float()
|
|
bright = pad_1d(brights, 0.0).float()
|
|
token_pitch = pad_2d(token_pitches, 0.0).float()
|
|
target_mel, frame_mask = pad_mels(mels)
|
|
pitch_frame, _ = pad_mels(pitches)
|
|
target_wav = pad_wavs(wavs, frame_counts, hop_size)
|
|
return (
|
|
phone.to(device),
|
|
tone.to(device),
|
|
lang.to(device),
|
|
speakers.to(device),
|
|
duration.to(device),
|
|
energy.to(device),
|
|
bright.to(device),
|
|
token_pitch.to(device),
|
|
target_mel.to(device),
|
|
frame_mask.to(device),
|
|
pitch_frame.to(device),
|
|
target_wav.to(device),
|
|
)
|
|
|
|
|
|
def prepare_row_features(
|
|
row: dict,
|
|
cfg: MicroFastSpeechConfig,
|
|
mel_frontend: MelFrontend,
|
|
device: torch.device,
|
|
max_seconds: float,
|
|
) -> dict:
|
|
dur = list(map(int, row["hifigan_durations"]))
|
|
wav_1d = load_audio(str(row["target_audio"]), cfg.sample_rate, max_seconds)
|
|
with torch.no_grad():
|
|
wav = wav_1d.unsqueeze(0).to(device)
|
|
mel = mel_frontend(wav).squeeze(0).detach().cpu()
|
|
dur = fit_durations(dur[: len(row["phone_ids"])], min(mel.shape[-1], cfg.max_frames))
|
|
mel = mel[:, : sum(dur)]
|
|
energy, bright = aggregate_token_features(mel, dur)
|
|
pitch = extract_pitch_features(wav_1d, cfg.sample_rate, mel.shape[-1])
|
|
token_pitch = aggregate_token_pitch(pitch, dur)
|
|
return {
|
|
"phone": torch.LongTensor(row["phone_ids"]),
|
|
"tone": torch.LongTensor(row["tone_ids"]),
|
|
"lang": torch.LongTensor(row["lang_ids"]),
|
|
"speaker": int(row["speaker_id"]),
|
|
"duration": torch.LongTensor(dur),
|
|
"energy": energy.float(),
|
|
"bright": bright.float(),
|
|
"token_pitch": token_pitch.float(),
|
|
"target_mel": mel.float(),
|
|
"pitch_frame": pitch.float(),
|
|
"target_wav": wav_1d.float(),
|
|
"frame_count": int(mel.shape[-1]),
|
|
}
|
|
|
|
|
|
def collate_prepared(batch: list[dict], device: torch.device, hop_size: int):
|
|
phone = pad_1d([x["phone"] for x in batch], 0).long()
|
|
tone = pad_1d([x["tone"] for x in batch], 0).long()
|
|
lang = pad_1d([x["lang"] for x in batch], 0).long()
|
|
speakers = torch.LongTensor([int(x["speaker"]) for x in batch])
|
|
duration = pad_1d([x["duration"] for x in batch], 0).long()
|
|
energy = pad_1d([x["energy"] for x in batch], 0.0).float()
|
|
bright = pad_1d([x["bright"] for x in batch], 0.0).float()
|
|
token_pitch = pad_2d([x["token_pitch"] for x in batch], 0.0).float()
|
|
target_mel, frame_mask = pad_mels([x["target_mel"] for x in batch])
|
|
pitch_frame, _ = pad_mels([x["pitch_frame"] for x in batch])
|
|
target_wav = pad_wavs([x["target_wav"] for x in batch], [int(x["frame_count"]) for x in batch], hop_size)
|
|
return (
|
|
phone.to(device),
|
|
tone.to(device),
|
|
lang.to(device),
|
|
speakers.to(device),
|
|
duration.to(device),
|
|
energy.to(device),
|
|
bright.to(device),
|
|
token_pitch.to(device),
|
|
target_mel.to(device),
|
|
frame_mask.to(device),
|
|
pitch_frame.to(device),
|
|
target_wav.to(device),
|
|
)
|
|
|
|
|
|
def masked_l1(pred: torch.Tensor, target: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|
common = min(pred.shape[-1], target.shape[-1], mask.shape[-1])
|
|
pred = pred[..., :common]
|
|
target = target[..., :common]
|
|
mask = mask[:, :common].unsqueeze(1)
|
|
return (torch.abs(pred - target) * mask).sum() / (mask.sum() * pred.shape[1]).clamp_min(1.0)
|
|
|
|
|
|
def masked_mse(pred: torch.Tensor, target: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|
common = min(pred.shape[-1], target.shape[-1], mask.shape[-1])
|
|
pred = pred[..., :common]
|
|
target = target[..., :common]
|
|
mask = mask[:, :common].unsqueeze(1)
|
|
return (((pred - target) ** 2) * mask).sum() / (mask.sum() * pred.shape[1]).clamp_min(1.0)
|
|
|
|
|
|
def masked_delta_loss(pred: torch.Tensor, target: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|
common = min(pred.shape[-1], target.shape[-1], mask.shape[-1])
|
|
if common < 2:
|
|
return torch.zeros((), device=pred.device)
|
|
dp = pred[..., 1:common] - pred[..., : common - 1]
|
|
dt = target[..., 1:common] - target[..., : common - 1]
|
|
dm = (mask[:, 1:common] & mask[:, : common - 1]).unsqueeze(1)
|
|
return (torch.abs(dp - dt) * dm).sum() / (dm.sum() * pred.shape[1]).clamp_min(1.0)
|
|
|
|
|
|
def masked_accel_loss(pred: torch.Tensor, target: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|
common = min(pred.shape[-1], target.shape[-1], mask.shape[-1])
|
|
if common < 3:
|
|
return torch.zeros((), device=pred.device)
|
|
dp = pred[..., 2:common] - 2.0 * pred[..., 1 : common - 1] + pred[..., : common - 2]
|
|
dt = target[..., 2:common] - 2.0 * target[..., 1 : common - 1] + target[..., : common - 2]
|
|
dm = (mask[:, 2:common] & mask[:, 1 : common - 1] & mask[:, : common - 2]).unsqueeze(1)
|
|
return (torch.abs(dp - dt) * dm).sum() / (dm.sum() * pred.shape[1]).clamp_min(1.0)
|
|
|
|
|
|
def token_mse(pred: torch.Tensor, target: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|
common = min(pred.shape[1], target.shape[1], mask.shape[1])
|
|
pred = pred[:, :common]
|
|
target = target[:, :common]
|
|
mask = mask[:, :common]
|
|
return (((pred - target) ** 2) * mask).sum() / mask.sum().clamp_min(1.0)
|
|
|
|
|
|
def token_mse_nd(pred: torch.Tensor, target: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|
common = min(pred.shape[1], target.shape[1], mask.shape[1])
|
|
pred = pred[:, :common]
|
|
target = target[:, :common]
|
|
mask = mask[:, :common].unsqueeze(-1)
|
|
return (((pred - target) ** 2) * mask).sum() / (mask.sum() * pred.shape[-1]).clamp_min(1.0)
|
|
|
|
|
|
def group_duration_targets(phone: torch.Tensor, durations: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
|
grouped = torch.zeros_like(durations, dtype=torch.float32)
|
|
mask = torch.zeros_like(phone, dtype=torch.bool)
|
|
for batch_index in range(phone.shape[0]):
|
|
pending: list[torch.Tensor] = []
|
|
last_visible: int | None = None
|
|
for token_index in range(phone.shape[1]):
|
|
pending.append(durations[batch_index, token_index].float())
|
|
if int(phone[batch_index, token_index].item()) != 0:
|
|
grouped[batch_index, token_index] = torch.log1p(torch.stack(pending).sum())
|
|
mask[batch_index, token_index] = True
|
|
pending = []
|
|
last_visible = token_index
|
|
if pending and last_visible is not None:
|
|
value = torch.expm1(grouped[batch_index, last_visible]) + torch.stack(pending).sum()
|
|
grouped[batch_index, last_visible] = torch.log1p(value)
|
|
return grouped, mask
|
|
|
|
|
|
def masked_wav_l1(pred: torch.Tensor, target: torch.Tensor, frame_mask: torch.Tensor, hop_size: int) -> torch.Tensor:
|
|
if pred.dim() == 3:
|
|
pred = pred.squeeze(1)
|
|
common = min(pred.shape[-1], target.shape[-1], frame_mask.shape[-1] * hop_size)
|
|
pred = pred[:, :common]
|
|
target = target[:, :common]
|
|
sample_mask = frame_mask.repeat_interleave(hop_size, dim=1)[:, :common].to(pred.dtype)
|
|
return (torch.abs(pred - target) * sample_mask).sum() / sample_mask.sum().clamp_min(1.0)
|
|
|
|
|
|
def load_frozen_vocoder(path: Path, device: torch.device) -> tuple[HifiGanGenerator, HifiGanConfig]:
|
|
ckpt = torch.load(path, map_location=device, weights_only=False)
|
|
cfg_payload = ckpt.get("config") or {"variant": "v2plus"}
|
|
cfg = HifiGanConfig(**cfg_payload) if isinstance(cfg_payload, dict) else cfg_payload
|
|
vocoder = HifiGanGenerator(cfg).to(device)
|
|
vocoder.load_state_dict(ckpt["generator"])
|
|
vocoder.eval()
|
|
for param in vocoder.parameters():
|
|
param.requires_grad_(False)
|
|
return vocoder, cfg
|
|
|
|
|
|
def load_model_state_flexible(model: nn.Module, state: dict[str, torch.Tensor]) -> tuple[int, int]:
|
|
current = model.state_dict()
|
|
compatible = {key: value for key, value in state.items() if key in current and current[key].shape == value.shape}
|
|
model.load_state_dict(compatible, strict=False)
|
|
return len(compatible), len(state) - len(compatible)
|
|
|
|
|
|
def set_trainable_by_mode(model: MicroFastSpeech, mode: str) -> None:
|
|
if mode == "all":
|
|
for param in model.parameters():
|
|
param.requires_grad_(True)
|
|
return
|
|
for param in model.parameters():
|
|
param.requires_grad_(False)
|
|
prefixes: tuple[str, ...]
|
|
if mode == "duration":
|
|
prefixes = ("phone.", "tone.", "lang.", "speaker.", "speaker_proj.", "encoder.", "duration_head.")
|
|
elif mode == "predictors":
|
|
prefixes = (
|
|
"phone.",
|
|
"tone.",
|
|
"lang.",
|
|
"speaker.",
|
|
"speaker_proj.",
|
|
"encoder.",
|
|
"duration_head.",
|
|
"energy_head.",
|
|
"bright_head.",
|
|
"pitch_head.",
|
|
)
|
|
elif mode == "heads":
|
|
prefixes = (
|
|
"duration_head.",
|
|
"energy_head.",
|
|
"bright_head.",
|
|
"pitch_head.",
|
|
"predictor_context.",
|
|
"duration_delta.",
|
|
"energy_delta.",
|
|
"bright_delta.",
|
|
"pitch_delta.",
|
|
)
|
|
elif mode == "contextual":
|
|
prefixes = (
|
|
"predictor_context.",
|
|
"duration_delta.",
|
|
"energy_delta.",
|
|
"bright_delta.",
|
|
"pitch_delta.",
|
|
)
|
|
elif mode == "group_duration":
|
|
prefixes = ("group_duration_delta.",)
|
|
elif mode == "decoder_adapt":
|
|
prefixes = (
|
|
"energy_proj.",
|
|
"bright_proj.",
|
|
"pitch_proj.",
|
|
"abs_frame.",
|
|
"frame_proj.",
|
|
"local_ctx.",
|
|
"decoder.",
|
|
"frame_gru.",
|
|
"mel_head.",
|
|
"postnet.",
|
|
)
|
|
else:
|
|
raise ValueError(f"Unknown trainable mode: {mode}")
|
|
for name, param in model.named_parameters():
|
|
if name.startswith(prefixes):
|
|
param.requires_grad_(True)
|
|
|
|
|
|
def latest_checkpoint(out_dir: Path) -> Path | None:
|
|
found = []
|
|
for path in out_dir.glob("inflect-micro-fastspeech-*.pt"):
|
|
tail = path.stem.rsplit("-", 1)[-1]
|
|
if tail.isdigit():
|
|
found.append((int(tail), path))
|
|
return max(found)[1] if found else None
|
|
|
|
|
|
def save_checkpoint(path: Path, model: nn.Module, optim, cfg: MicroFastSpeechConfig, step: int, args, speakers: dict[str, int]) -> None:
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
tmp = path.with_suffix(path.suffix + ".tmp")
|
|
torch.save(
|
|
{
|
|
"model": model.state_dict(),
|
|
"optim": optim.state_dict(),
|
|
"config": asdict(cfg),
|
|
"step": step,
|
|
"speakers": speakers,
|
|
"args": {k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()},
|
|
"params": count_parameters(model),
|
|
},
|
|
tmp,
|
|
)
|
|
tmp.replace(path)
|
|
|
|
|
|
def train(args: argparse.Namespace) -> None:
|
|
device = torch.device(args.device)
|
|
rows = load_rows(args.durations_jsonl, args.max_rows)
|
|
speakers = {voice: idx for idx, voice in enumerate(sorted({str(r.get("voice_id") or "qwen3_mark") for r in rows}))}
|
|
max_phone_id = max(max(map(int, r["phone_ids"])) for r in rows)
|
|
max_tone_id = max(max(map(int, r["tone_ids"])) for r in rows)
|
|
max_lang_id = max(max(map(int, r["lang_ids"])) for r in rows)
|
|
cfg = MicroFastSpeechConfig(
|
|
vocab_size=max(256, max_phone_id + 1),
|
|
tone_size=max(16, max_tone_id + 1),
|
|
lang_size=max(4, max_lang_id + 1),
|
|
speaker_count=max(2, len(speakers)),
|
|
hidden=args.hidden,
|
|
encoder_layers=args.encoder_layers,
|
|
decoder_layers=args.decoder_layers,
|
|
decoder_ff_mult=args.decoder_ff_mult,
|
|
max_frames=args.max_frames,
|
|
postnet_scale=args.postnet_scale,
|
|
abs_frame_bins=args.abs_frame_bins,
|
|
use_contextual_predictors=args.contextual_predictors,
|
|
use_group_duration_planner=args.group_duration_planner,
|
|
)
|
|
for row in rows:
|
|
row["speaker_id"] = speakers[str(row.get("voice_id") or "qwen3_mark")]
|
|
random.Random(args.seed).shuffle(rows)
|
|
|
|
model = MicroFastSpeech(cfg).to(device)
|
|
start_step = 0
|
|
if args.init_checkpoint and not args.resume:
|
|
ckpt = torch.load(args.init_checkpoint, map_location=device, weights_only=False)
|
|
copied, skipped = load_model_state_flexible(model, ckpt["model"])
|
|
print(f"Initialized model from {args.init_checkpoint} ({copied} tensors copied, {skipped} skipped)")
|
|
set_trainable_by_mode(model, args.trainable)
|
|
trainable_params = [param for param in model.parameters() if param.requires_grad]
|
|
optim = torch.optim.AdamW(trainable_params, lr=args.lr, betas=(0.9, 0.98), weight_decay=args.weight_decay)
|
|
if args.resume:
|
|
ckpt_path = latest_checkpoint(args.out_dir)
|
|
if ckpt_path:
|
|
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
|
|
model.load_state_dict(ckpt["model"])
|
|
optim.load_state_dict(ckpt["optim"])
|
|
start_step = int(ckpt.get("step") or 0)
|
|
print(f"Resumed {ckpt_path} at step {start_step}")
|
|
|
|
hifi_cfg = HifiGanConfig(variant="v2plus")
|
|
mel_frontend = MelFrontend(hifi_cfg).to(device)
|
|
prepared_rows = None
|
|
if args.preload_features:
|
|
print("Preloading audio/mel/pitch features...", flush=True)
|
|
prepared_rows = [prepare_row_features(row, cfg, mel_frontend, device, args.max_seconds) for row in rows]
|
|
total_frames = sum(int(row["frame_count"]) for row in prepared_rows)
|
|
print(f"Preloaded {len(prepared_rows)} rows ({total_frames:,} frames)", flush=True)
|
|
consistency_vocoder = None
|
|
if args.vocoder_checkpoint:
|
|
consistency_vocoder, consistency_cfg = load_frozen_vocoder(args.vocoder_checkpoint, device)
|
|
if consistency_cfg.hop_size != hifi_cfg.hop_size:
|
|
raise RuntimeError(f"Vocoder hop mismatch: {consistency_cfg.hop_size} != {hifi_cfg.hop_size}")
|
|
print(f"Loaded frozen vocoder consistency checkpoint: {args.vocoder_checkpoint}")
|
|
if (args.vocoder_wav_weight > 0.0 or args.vocoder_mel_weight > 0.0) and consistency_vocoder is None:
|
|
raise RuntimeError("--vocoder-checkpoint is required when vocoder consistency losses are enabled")
|
|
args.out_dir.mkdir(parents=True, exist_ok=True)
|
|
(args.out_dir / "config.json").write_text(
|
|
json.dumps({"config": asdict(cfg), "speakers": speakers, "rows": len(rows), "params": count_parameters(model)}, indent=2),
|
|
encoding="utf-8",
|
|
)
|
|
|
|
print(f"Rows: {len(rows)}")
|
|
print(f"Speakers: {speakers}")
|
|
print(f"Acoustic params: {count_parameters(model):,} ({count_parameters(model)/1_000_000:.3f}M)")
|
|
print(f"Trainable params: {sum(p.numel() for p in model.parameters() if p.requires_grad):,} mode={args.trainable}")
|
|
print(f"Total with V2+ vocoder: {(count_parameters(model)+1_426_842):,} ({(count_parameters(model)+1_426_842)/1_000_000:.3f}M)")
|
|
|
|
rng = random.Random(args.seed + start_step)
|
|
step = start_step
|
|
started = time.time()
|
|
while step < args.steps:
|
|
source_rows = prepared_rows if prepared_rows is not None else rows
|
|
batch = [source_rows[rng.randrange(len(source_rows))] for _ in range(args.batch_size)]
|
|
if prepared_rows is not None:
|
|
phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_token_t, target_mel, frame_mask, pitch_frame, target_wav = collate_prepared(
|
|
batch, device, hifi_cfg.hop_size
|
|
)
|
|
else:
|
|
phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_token_t, target_mel, frame_mask, pitch_frame, target_wav = collate(
|
|
batch, cfg, mel_frontend, device, args.max_seconds, hifi_cfg.hop_size
|
|
)
|
|
out = model(phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_frame)
|
|
token_mask = out["token_mask"]
|
|
log_dur_t = torch.log1p(durations.float())
|
|
group_log_dur_t, group_mask = group_duration_targets(phone, durations)
|
|
mel_l1 = masked_l1(out["mel"], target_mel, frame_mask)
|
|
mel_mse = masked_mse(out["mel"], target_mel, frame_mask)
|
|
delta = masked_delta_loss(out["mel"], target_mel, frame_mask)
|
|
accel = masked_accel_loss(out["mel"], target_mel, frame_mask)
|
|
dur_loss = token_mse(out["log_dur"], log_dur_t, token_mask)
|
|
group_dur_loss = token_mse(out["group_log_dur"], group_log_dur_t, group_mask)
|
|
energy_loss = token_mse(out["energy"], energy_t, token_mask)
|
|
bright_loss = token_mse(out["bright"], bright_t, token_mask)
|
|
pitch_loss = token_mse_nd(out["pitch"], pitch_token_t, token_mask)
|
|
predicted_prosody_mel_loss = torch.zeros((), device=device)
|
|
predicted_prosody_delta_loss = torch.zeros((), device=device)
|
|
if args.predicted_prosody_mel_weight > 0.0 or args.predicted_prosody_delta_weight > 0.0:
|
|
# Train the predictor heads against the acoustic result they produce at
|
|
# inference, while retaining teacher durations so this path remains
|
|
# differentiable and isolates prosody exposure bias.
|
|
predicted_conditioning = model(phone, tone, lang, speaker, durations)
|
|
if args.predicted_prosody_mel_weight > 0.0:
|
|
predicted_prosody_mel_loss = masked_l1(predicted_conditioning["mel"], target_mel, frame_mask)
|
|
if args.predicted_prosody_delta_weight > 0.0:
|
|
predicted_prosody_delta_loss = masked_delta_loss(predicted_conditioning["mel"], target_mel, frame_mask)
|
|
robust_prosody_mel_loss = torch.zeros((), device=device)
|
|
robust_prosody_delta_loss = torch.zeros((), device=device)
|
|
if args.robust_prosody_mel_weight > 0.0 or args.robust_prosody_delta_weight > 0.0:
|
|
robust_conditioning = model(
|
|
phone,
|
|
tone,
|
|
lang,
|
|
speaker,
|
|
durations,
|
|
energy_t,
|
|
bright_t,
|
|
pitch_frame,
|
|
predicted_prosody_mix=args.robust_prosody_mix,
|
|
detach_mixed_predictions=True,
|
|
)
|
|
if args.robust_prosody_mel_weight > 0.0:
|
|
robust_prosody_mel_loss = masked_l1(robust_conditioning["mel"], target_mel, frame_mask)
|
|
if args.robust_prosody_delta_weight > 0.0:
|
|
robust_prosody_delta_loss = masked_delta_loss(robust_conditioning["mel"], target_mel, frame_mask)
|
|
voc_wav_loss = torch.zeros((), device=device)
|
|
voc_mel_loss = torch.zeros((), device=device)
|
|
if consistency_vocoder is not None and (args.vocoder_wav_weight > 0.0 or args.vocoder_mel_weight > 0.0):
|
|
pred_wav = consistency_vocoder(out["mel"].clamp(-12.0, 2.0))
|
|
if args.vocoder_wav_weight > 0.0:
|
|
voc_wav_loss = masked_wav_l1(pred_wav, target_wav, frame_mask, hifi_cfg.hop_size)
|
|
if args.vocoder_mel_weight > 0.0:
|
|
pred_recon_mel = mel_frontend(pred_wav.squeeze(1))
|
|
voc_mel_loss = masked_l1(pred_recon_mel, target_mel, frame_mask)
|
|
loss = (
|
|
mel_l1
|
|
+ args.mse_weight * mel_mse
|
|
+ args.delta_weight * delta
|
|
+ args.accel_weight * accel
|
|
+ args.duration_weight * dur_loss
|
|
+ args.group_duration_weight * group_dur_loss
|
|
+ args.energy_weight * energy_loss
|
|
+ args.bright_weight * bright_loss
|
|
+ args.pitch_weight * pitch_loss
|
|
+ args.predicted_prosody_mel_weight * predicted_prosody_mel_loss
|
|
+ args.predicted_prosody_delta_weight * predicted_prosody_delta_loss
|
|
+ args.robust_prosody_mel_weight * robust_prosody_mel_loss
|
|
+ args.robust_prosody_delta_weight * robust_prosody_delta_loss
|
|
+ args.vocoder_wav_weight * voc_wav_loss
|
|
+ args.vocoder_mel_weight * voc_mel_loss
|
|
)
|
|
optim.zero_grad(set_to_none=True)
|
|
loss.backward()
|
|
grad = torch.nn.utils.clip_grad_norm_(trainable_params, args.grad_clip)
|
|
optim.step()
|
|
step += 1
|
|
|
|
if step == 1 or step % args.log_interval == 0:
|
|
elapsed = max(1e-6, time.time() - started)
|
|
speed = (step - start_step) / elapsed
|
|
eta = (args.steps - step) / max(1e-6, speed)
|
|
print(
|
|
f"step={step}/{args.steps} loss={loss.item():.4f} mel={mel_l1.item():.4f} "
|
|
f"mse={mel_mse.item():.4f} delta={delta.item():.4f} accel={accel.item():.4f} dur={dur_loss.item():.4f} "
|
|
f"gdur={group_dur_loss.item():.4f} "
|
|
f"energy={energy_loss.item():.4f} bright={bright_loss.item():.4f} "
|
|
f"pitch={pitch_loss.item():.4f} pmel={predicted_prosody_mel_loss.item():.4f} "
|
|
f"pdelta={predicted_prosody_delta_loss.item():.4f} rmel={robust_prosody_mel_loss.item():.4f} "
|
|
f"rdelta={robust_prosody_delta_loss.item():.4f} vwav={voc_wav_loss.item():.4f} "
|
|
f"vmel={voc_mel_loss.item():.4f} grad={float(grad):.2f} "
|
|
f"speed={speed:.3f} step/s eta={eta/60:.1f}m",
|
|
flush=True,
|
|
)
|
|
if step % args.save_interval == 0 or step >= args.steps:
|
|
save_checkpoint(args.out_dir / f"inflect-micro-fastspeech-{step}.pt", model, optim, cfg, step, args, speakers)
|
|
save_checkpoint(args.out_dir / "inflect-micro-fastspeech-latest.pt", model, optim, cfg, step, args, speakers)
|
|
|
|
print(f"Done. {args.out_dir}")
|
|
|
|
|
|
def main() -> None:
|
|
ap = argparse.ArgumentParser(description="Train Inflect Micro duration-conditioned acoustic model.")
|
|
ap.add_argument("--durations-jsonl", type=Path, required=True)
|
|
ap.add_argument("--out-dir", type=Path, required=True)
|
|
ap.add_argument("--max-rows", type=int, default=0)
|
|
ap.add_argument("--steps", type=int, default=20000)
|
|
ap.add_argument("--batch-size", type=int, default=6)
|
|
ap.add_argument("--lr", type=float, default=2.0e-4)
|
|
ap.add_argument("--weight-decay", type=float, default=1.0e-4)
|
|
ap.add_argument("--hidden", type=int, default=168)
|
|
ap.add_argument("--encoder-layers", type=int, default=5)
|
|
ap.add_argument("--decoder-layers", type=int, default=6)
|
|
ap.add_argument("--decoder-ff-mult", type=int, default=3)
|
|
ap.add_argument("--max-seconds", type=float, default=12.0)
|
|
ap.add_argument("--max-frames", type=int, default=1400)
|
|
ap.add_argument("--mse-weight", type=float, default=0.25)
|
|
ap.add_argument("--delta-weight", type=float, default=0.18)
|
|
ap.add_argument("--accel-weight", type=float, default=0.0)
|
|
ap.add_argument("--duration-weight", type=float, default=0.08)
|
|
ap.add_argument("--group-duration-weight", type=float, default=0.0)
|
|
ap.add_argument("--energy-weight", type=float, default=0.04)
|
|
ap.add_argument("--bright-weight", type=float, default=0.04)
|
|
ap.add_argument("--pitch-weight", type=float, default=0.04)
|
|
ap.add_argument("--predicted-prosody-mel-weight", type=float, default=0.0)
|
|
ap.add_argument("--predicted-prosody-delta-weight", type=float, default=0.0)
|
|
ap.add_argument("--robust-prosody-mix", type=float, default=0.0)
|
|
ap.add_argument("--robust-prosody-mel-weight", type=float, default=0.0)
|
|
ap.add_argument("--robust-prosody-delta-weight", type=float, default=0.0)
|
|
ap.add_argument("--grad-clip", type=float, default=5.0)
|
|
ap.add_argument("--postnet-scale", type=float, default=0.10)
|
|
ap.add_argument("--abs-frame-bins", type=int, default=512)
|
|
ap.add_argument("--init-checkpoint", type=Path)
|
|
ap.add_argument("--vocoder-checkpoint", type=Path)
|
|
ap.add_argument("--vocoder-wav-weight", type=float, default=0.0)
|
|
ap.add_argument("--vocoder-mel-weight", type=float, default=0.0)
|
|
ap.add_argument("--save-interval", type=int, default=2000)
|
|
ap.add_argument("--log-interval", type=int, default=50)
|
|
ap.add_argument("--seed", type=int, default=42)
|
|
ap.add_argument("--resume", action="store_true")
|
|
ap.add_argument("--preload-features", action="store_true", help="Cache decoded audio, mels, pitch, and token features in RAM before training.")
|
|
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
|
ap.add_argument(
|
|
"--trainable",
|
|
choices=["all", "duration", "predictors", "heads", "contextual", "group_duration", "decoder_adapt"],
|
|
default="all",
|
|
)
|
|
ap.add_argument("--contextual-predictors", action="store_true")
|
|
ap.add_argument("--group-duration-planner", action="store_true")
|
|
args = ap.parse_args()
|
|
train(args)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|