AuralGuard / src /train.py
AyoPrince's picture
Upload folder using huggingface_hub
127b976 verified
Raw History Blame Contribute Delete
6.18 kB
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Dict
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm
from .aasist_loader import build_aasist_backbone
from .dataset import AuralGuardDataset, AudioConfig
from .model_aasistpp import AuralGuardAASISTPP, compute_multitask_loss
from .metrics import classification_report_dict
def parse_args():
p = argparse.ArgumentParser(description="Train AuralGuard-AASIST++")
p.add_argument("--train-csv", required=True)
p.add_argument("--val-csv", required=True)
p.add_argument("--aasist-root", default="external/aasist")
p.add_argument("--aasist-config", default="external/aasist/config/AASIST.conf")
p.add_argument("--aasist-checkpoint", default=None, help="Optional pretrained AASIST checkpoint")
p.add_argument("--out-dir", default="results/run1")
p.add_argument("--epochs", type=int, default=10)
p.add_argument("--batch-size", type=int, default=8)
p.add_argument("--lr", type=float, default=1e-4)
p.add_argument("--sample-rate", type=int, default=16000)
p.add_argument("--duration-sec", type=float, default=4.0)
p.add_argument("--feature-dim", type=int, default=160)
p.add_argument("--num-workers", type=int, default=2)
p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
p.add_argument("--freeze-backbone", action="store_true")
return p.parse_args()
def move_batch_to_device(batch: Dict, device: torch.device) -> Dict:
out = {}
for k, v in batch.items():
out[k] = v.to(device) if torch.is_tensor(v) else v
return out
def run_epoch(model, loader, optimizer, device, train: bool):
model.train(train)
total_loss = 0.0
y_true, y_pred, fake_scores = [], [], []
attack_true, attack_pred = [], []
exp_true, exp_pred = [], []
loop = tqdm(loader, desc="train" if train else "val", leave=False)
for batch in loop:
batch = move_batch_to_device(batch, device)
with torch.set_grad_enabled(train):
outputs = model(batch["wav"], freq_aug=train)
losses = compute_multitask_loss(
outputs,
batch["binary_label"],
batch["attack_label"],
batch["explanation_label"],
)
if train:
optimizer.zero_grad(set_to_none=True)
losses["loss"].backward()
optimizer.step()
probs = torch.softmax(outputs["binary_logits"], dim=-1)
fake_prob = probs[:, 1]
pred = torch.argmax(outputs["binary_logits"], dim=-1)
attack_p = torch.argmax(outputs["attack_logits"], dim=-1)
exp_p = torch.argmax(outputs["explanation_logits"], dim=-1)
bs = batch["wav"].shape[0]
total_loss += losses["loss"].item() * bs
y_true.extend(batch["binary_label"].detach().cpu().tolist())
y_pred.extend(pred.detach().cpu().tolist())
fake_scores.extend(fake_prob.detach().cpu().tolist())
attack_true.extend(batch["attack_label"].detach().cpu().tolist())
attack_pred.extend(attack_p.detach().cpu().tolist())
exp_true.extend(batch["explanation_label"].detach().cpu().tolist())
exp_pred.extend(exp_p.detach().cpu().tolist())
loop.set_postfix(loss=losses["loss"].item())
metrics = classification_report_dict(
y_true, y_pred, fake_scores, attack_true, attack_pred, exp_true, exp_pred
)
metrics["loss"] = total_loss / max(len(loader.dataset), 1)
return metrics
def main():
args = parse_args()
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
device = torch.device(args.device)
audio_train = AudioConfig(args.sample_rate, args.duration_sec, random_crop=True)
audio_eval = AudioConfig(args.sample_rate, args.duration_sec, random_crop=False)
train_ds = AuralGuardDataset(args.train_csv, audio_config=audio_train)
val_ds = AuralGuardDataset(args.val_csv, audio_config=audio_eval)
train_loader = DataLoader(
train_ds, batch_size=args.batch_size, shuffle=True,
num_workers=args.num_workers, pin_memory=(device.type == "cuda")
)
val_loader = DataLoader(
val_ds, batch_size=args.batch_size, shuffle=False,
num_workers=args.num_workers, pin_memory=(device.type == "cuda")
)
backbone = build_aasist_backbone(
args.aasist_root, args.aasist_config, checkpoint=args.aasist_checkpoint, device=device
)
model = AuralGuardAASISTPP(
backbone,
feature_dim=args.feature_dim,
freeze_backbone=args.freeze_backbone,
).to(device)
optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr=args.lr)
best_eer = float("inf")
history = []
for epoch in range(1, args.epochs + 1):
print(f"\nEpoch {epoch}/{args.epochs}")
train_metrics = run_epoch(model, train_loader, optimizer, device, train=True)
val_metrics = run_epoch(model, val_loader, optimizer, device, train=False)
record = {"epoch": epoch, "train": train_metrics, "val": val_metrics}
history.append(record)
print("train:", json.dumps(train_metrics, indent=2))
print("val:", json.dumps(val_metrics, indent=2))
ckpt = {
"epoch": epoch,
"model": model.state_dict(),
"args": vars(args),
"val_metrics": val_metrics,
}
torch.save(ckpt, out_dir / "last.pt")
val_eer = val_metrics.get("eer", float("inf"))
if val_eer < best_eer:
best_eer = val_eer
torch.save(ckpt, out_dir / "best.pt")
print(f"Saved new best checkpoint with EER={best_eer:.4f}")
with (out_dir / "history.json").open("w", encoding="utf-8") as f:
json.dump(history, f, indent=2)
print(f"Done. Best checkpoint: {out_dir / 'best.pt'}")
if __name__ == "__main__":
main()