""" python train.py \ --train_audio ./data/synth_dataset/train/audio \ --train_bs ./data/synth_dataset/train/bs \ --val_audio ./data/synth_dataset/val/audio \ --val_bs ./data/synth_dataset/val/bs \ --bs_dim 24 \ --batch_size 4 \ --epochs 5 \ --save_dir ./data/ckpt_synth24 """ import argparse, os, time, math, json from pathlib import Path import numpy as np import librosa from tqdm import tqdm import torch from torch.utils.data import Dataset, DataLoader from torch.utils.tensorboard import SummaryWriter from a2fmodel import A2BModel # --------------------------------------------------------------------- # ============ 1. 数据集 ============ class FaceDataset(Dataset): """ 读取配对的 (wav, bs.npy) - wav 采样率统一到 16k - bs.npy 形状 [T, N], N = --bs_dim """ def __init__(self, audio_dir, bs_dir): self.audio_paths = sorted(Path(audio_dir).glob("*.wav")) self.bs_dir = Path(bs_dir) assert len(self.audio_paths) > 0, f"no wav in {audio_dir}" def __len__(self): return len(self.audio_paths) def __getitem__(self, idx): wav_p = self.audio_paths[idx] bs_p = self.bs_dir / f"{wav_p.stem}.npy" assert bs_p.exists(), f"missing {bs_p}" audio, _ = librosa.load(wav_p, sr=16000) # (L,) audio = torch.from_numpy(audio).float().unsqueeze(0) # (1, L) bs = torch.from_numpy(np.load(bs_p)).float() # (T, N) return {"audio": audio, "bs": bs} # ============ 2. collate(变长 pad) ============ def collate_fn(batch): # --------- audio pad ---------- max_a = max(b["audio"].shape[1] for b in batch) audios = torch.zeros(len(batch), 1, max_a) for i, b in enumerate(batch): L = b["audio"].shape[1] audios[i, 0, :L] = b["audio"] # --------- bs pad ------------- max_t = max(b["bs"].shape[0] for b in batch) N = batch[0]["bs"].shape[1] bss = torch.zeros(len(batch), max_t, N) lens = torch.zeros(len(batch), dtype=torch.long) for i, b in enumerate(batch): T = b["bs"].shape[0] bss[i, :T] = b["bs"] lens[i] = T # A2BModel 期望的字段(3 份音频/target 可直接复用) data = { "input11": audios, # (B,1,L) "input12": audios.clone(), "input21": audios.clone(), "target11": bss, # (B,T,N) "target12": bss.clone(), "target21": bss.clone(), "level": torch.zeros(len(batch), dtype=torch.long), # 未使用,默认全 0 "person": torch.zeros(len(batch), dtype=torch.long), # 未使用,默认全 0 "lengths": lens, # 每条样本的有效 T } return data # ============ 3. 损失 ============ def mse_loss(pred, gt, mask): """ pred/gt: (B, T, N) mask: (B, T) (1=valid, 0=pad) """ diff = (pred - gt) ** 2 # (B,T,N) diff = diff.sum(-1) # (B,T) diff = diff * mask return diff.sum() / mask.sum().clamp_min(1.0) # ============ 4. 训练 / 验证 ============ def train_one_epoch(model, loader, optimizer, device): model.train() total_loss = 0.0 for batch in tqdm(loader, desc="Train", leave=False): for k in batch: if isinstance(batch[k], torch.Tensor): batch[k] = batch[k].to(device, non_blocking=True) bs_pred1, bs_pred2, _ = model(batch) # (B,T,N) tgt = batch["target11"] # (B,T,N) # 构造 mask (B,T) B, T, _ = tgt.shape mask = (torch.arange(T, device=device)[None, :] < batch["lengths"][:, None]).float() loss = mse_loss(bs_pred1, tgt, mask) + mse_loss(bs_pred2, tgt, mask) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() return total_loss / max(len(loader), 1) @torch.no_grad() def evaluate(model, loader, device): model.eval() total_loss = 0.0 for batch in loader: for k in batch: if isinstance(batch[k], torch.Tensor): batch[k] = batch[k].to(device, non_blocking=True) bs_pred1, bs_pred2, _ = model(batch) tgt = batch["target11"] B, T, _ = tgt.shape mask = (torch.arange(T, device=device)[None, :] < batch["lengths"][:, None]).float() loss = mse_loss(bs_pred1, tgt, mask) + mse_loss(bs_pred2, tgt, mask) total_loss += loss.item() return total_loss / max(len(loader), 1) # ============ 5. CLI / 入口 ============ def main(): parser = argparse.ArgumentParser() # ---------- 数据路径 ---------- parser.add_argument("--train_audio", required=True, type=str) parser.add_argument("--train_bs", required=True, type=str) parser.add_argument("--val_audio", type=str, default=None) parser.add_argument("--val_bs", type=str, default=None) # ---------- 训练超参 ---------- parser.add_argument("--epochs", type=int, default=20) parser.add_argument("--batch_size", type=int, default=2) parser.add_argument("--lr", type=float, default=1e-4) parser.add_argument("--device", type=str, default="cuda:0") parser.add_argument("--save_dir", type=str, default="./checkpoints") # ---------- A2BModel 超参 ---------- parser.add_argument("--bs_dim", type=int, default=52) # 可改为你的 DOF 数 parser.add_argument("--feature_dim", type=int, default=832) parser.add_argument("--period", type=int, default=30) parser.add_argument("--max_seq_len", type=int, default=5000) parser.add_argument("--emo_guide", action="store_true") args = parser.parse_args() # ---------- 目录 / 日志 ---------- Path(args.save_dir).mkdir(parents=True, exist_ok=True) with open(Path(args.save_dir) / "hparams.json", "w") as f: json.dump(vars(args), f, indent=2, ensure_ascii=False) writer = SummaryWriter(log_dir=os.path.join(args.save_dir, "tb")) device = torch.device(args.device if torch.cuda.is_available() else "cpu") # ---------- DataLoader ---------- train_set = FaceDataset(args.train_audio, args.train_bs) train_loader = DataLoader( train_set, batch_size=args.batch_size, shuffle=True, num_workers=4, pin_memory=True, collate_fn=collate_fn ) val_loader = None if args.val_audio and args.val_bs: val_set = FaceDataset(args.val_audio, args.val_bs) val_loader = DataLoader( val_set, batch_size=args.batch_size, shuffle=False, num_workers=2, pin_memory=True, collate_fn=collate_fn ) # ---------- 模型 ---------- model_cfg = argparse.Namespace( bs_dim=args.bs_dim, feature_dim=args.feature_dim, period=args.period, max_seq_len=args.max_seq_len, emo_guide=args.emo_guide, device=args.device, # A2BModel 内部会用 .to(self.device) batch_size=args.batch_size, ) model = A2BModel(model_cfg).to(device) # ---------- 优化器 / 学习率计划 ---------- optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs) # ---------- 训练循环 ---------- best_val = math.inf for epoch in range(1, args.epochs + 1): t0 = time.time() train_loss = train_one_epoch(model, train_loader, optimizer, device) writer.add_scalar("loss/train", train_loss, epoch) msg = f"[{epoch:03d}/{args.epochs}] train {train_loss:.4f}" if val_loader is not None: val_loss = evaluate(model, val_loader, device) writer.add_scalar("loss/val", val_loss, epoch) msg += f" | val {val_loss:.4f}" if val_loss < best_val: best_val = val_loss torch.save(model.state_dict(), Path(args.save_dir) / "best.pth") # 保存最近模型 torch.save(model.state_dict(), Path(args.save_dir) / "last.pth") scheduler.step() dt = time.time() - t0 print(msg + f" | time {dt:.1f}s") writer.close() if best_val < math.inf: print(f"Training finished. Best val loss = {best_val:.4f}") else: print("Training finished.") if __name__ == "__main__": main()