Files
soprano-factory/train_codec.py
T
2026-02-09 01:01:44 -05:00

154 lines
5.2 KiB
Python

# train_codec.py - Stage 0: Codec Training (Standardized for V3)
import argparse
import glob
import os
import sys
# Windows UTF-8 Support
sys.stdout.reconfigure(encoding='utf-8')
import torch
import torch.nn.functional as F
import torchaudio
import soundfile as sf
from torch.utils.data import DataLoader, Dataset
from torch.optim import AdamW
from tqdm import tqdm
# Import our architecture
from model.encoder import Encoder
from model.decoder import Decoder
from utils.config import cfg
# --- Windows / Audio Backend Setup ---
def setup_audio_backend():
cwd = os.getcwd()
local_ffmpeg = os.path.join(cwd, "tools", "ffmpeg")
if os.path.exists(local_ffmpeg) and os.path.isdir(local_ffmpeg):
if local_ffmpeg not in os.environ["PATH"]:
os.environ["PATH"] = local_ffmpeg + os.pathsep + os.environ["PATH"]
setup_audio_backend()
class SpectralLoss(torch.nn.Module):
"""
Computes time-domain L1 loss and frequency-domain Mel loss.
Forced to Float32 for numerical stability.
"""
def __init__(self):
super().__init__()
self.mel = torchaudio.transforms.MelSpectrogram(
sample_rate=cfg.common['sample_rate'],
n_mels=cfg.codec['input_mels'],
n_fft=2048, hop_length=512
)
def forward(self, pred, target):
pred, target = pred.float(), target.float()
min_len = min(pred.shape[-1], target.shape[-1])
pred = pred[..., :min_len]
target = target[..., :min_len]
loss_time = F.l1_loss(pred, target)
if self.mel.mel_scale.fb.device != pred.device:
self.mel = self.mel.to(pred.device)
loss_mel = F.l1_loss(self.mel(pred), self.mel(target))
return loss_time + loss_mel
class WavDataset(Dataset):
def __init__(self, glob_pattern, segment_size):
self.files = glob.glob(glob_pattern, recursive=True)
self.segment_size = segment_size
print(f"Stage 0: Found {len(self.files)} wav files for training.")
def __len__(self):
return len(self.files)
def __getitem__(self, idx):
try:
wav_np, sr = sf.read(self.files[idx])
wav = torch.from_numpy(wav_np).float()
if wav.ndim == 1: wav = wav.unsqueeze(0)
else: wav = wav.t()
if sr != cfg.common['sample_rate']:
wav = torchaudio.functional.resample(wav, sr, cfg.common['sample_rate'])
if wav.shape[0] > 1: wav = wav.mean(dim=0, keepdim=True)
if wav.size(-1) < self.segment_size:
wav = F.pad(wav, (0, self.segment_size - wav.size(-1)))
if wav.size(-1) > self.segment_size:
start = torch.randint(0, wav.size(-1) - self.segment_size, (1,))
wav = wav[..., start : start + self.segment_size]
return wav
except Exception as e:
return torch.zeros(1, self.segment_size)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--wav-dir", type=str, required=True, help="Path to wavs")
parser.add_argument("--epochs", type=int, default=cfg.codec['epochs'])
parser.add_argument("--save-dir", type=str, default=None, help="Override default save directory")
args = parser.parse_args()
save_dir = args.save_dir if args.save_dir else cfg.codec['save_dir']
os.makedirs(save_dir, exist_ok=True)
device = cfg.device
print(f"Launching Codec Training (Float32) on {device}...")
# 1. Initialize Encoder
encoder = Encoder(
num_input_mels=cfg.codec['input_mels'],
encoder_dim=cfg.codec['dim'],
encoder_layers=cfg.codec['layers'],
bottleneck_channels=cfg.codec['bottleneck']
).to(device).float()
# 2. Initialize Decoder (SET INPUT TO 5 FOR STAGE 0)
decoder = Decoder(
input_channels=cfg.codec['bottleneck'], # <--- Specifically uses the 5 bottleneck channels
decoder_dim=cfg.codec['dim'],
decoder_layers=cfg.codec['layers']
).to(device).float()
opt = AdamW(list(encoder.parameters()) + list(decoder.parameters()), lr=float(cfg.codec['lr']))
criterion = SpectralLoss()
ds = WavDataset(args.wav_dir, segment_size=cfg.codec['segment_size'])
dl = DataLoader(ds, batch_size=cfg.codec['batch_size'], shuffle=True, num_workers=0, pin_memory=True)
# 3. Training Loop
for epoch in range(args.epochs):
encoder.train()
decoder.train()
pbar = tqdm(dl)
for wav in pbar:
wav = wav.to(device)
# Forward
z = encoder(wav)
rec = decoder(z)
loss = criterion(rec, wav)
# Backward
opt.zero_grad()
loss.backward()
opt.step()
pbar.set_description(f"Ep {epoch+1}/{args.epochs} | Loss: {loss.item():.4f}")
# Save checkpoints
if (epoch + 1) % 5 == 0 or epoch == args.epochs - 1:
torch.save(encoder.state_dict(), f"{save_dir}/encoder.pth")
torch.save(decoder.state_dict(), f"{save_dir}/decoder.pth")
print(f"Weights saved to {save_dir}")
if __name__ == "__main__":
main()