expre-gen/train.py
2025-10-17 09:47:23 +08:00

246 lines
8.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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()