mirror of
https://github.com/Nighthawk42/soprano-factory.git
synced 2026-08-30 04:30:21 +00:00
130 lines
4.6 KiB
Python
130 lines
4.6 KiB
Python
# train_codec.py - Stage 0: Codec Training (Standardized for V3 - Unicode Safe)
|
|
|
|
import argparse
|
|
import glob
|
|
import os
|
|
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"]:
|
|
print(f"[Setup] Found local FFmpeg at: {local_ffmpeg}")
|
|
os.environ["PATH"] = local_ffmpeg + os.pathsep + os.environ["PATH"]
|
|
|
|
setup_audio_backend()
|
|
|
|
class SpectralLoss(torch.nn.Module):
|
|
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:
|
|
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'])
|
|
args = parser.parse_args()
|
|
|
|
save_dir = cfg.codec['save_dir']
|
|
os.makedirs(save_dir, exist_ok=True)
|
|
device = cfg.device
|
|
print(f"Launching Codec Training (Float32) on {device}...")
|
|
|
|
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()
|
|
|
|
decoder = Decoder(
|
|
input_channels=cfg.codec['bottleneck'],
|
|
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)
|
|
|
|
for epoch in range(args.epochs):
|
|
encoder.train()
|
|
decoder.train()
|
|
pbar = tqdm(dl)
|
|
for wav in pbar:
|
|
wav = wav.to(device)
|
|
z = encoder(wav)
|
|
rec = decoder(z)
|
|
loss = criterion(rec, wav)
|
|
opt.zero_grad()
|
|
loss.backward()
|
|
opt.step()
|
|
pbar.set_description(f"Ep {epoch+1}/{args.epochs} | Loss: {loss.item():.4f}")
|
|
|
|
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() |