246 lines
8.2 KiB
Python
246 lines
8.2 KiB
Python
|
|
"""
|
|||
|
|
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()
|