Spaces:
Running on Zero
Running on Zero
Download src/train.py from AyoPrince/AuralGuard: direct link, hf CLI and curl.
- Browser
- Download file 6.18 kB
-
https://huggingface.co/spaces/AyoPrince/AuralGuard/resolve/main/src/train.py
- Command line
-
hf download hf://spaces/AyoPrince/AuralGuard/src/train.py
-
curl -L -o train.py https://huggingface.co/spaces/AyoPrince/AuralGuard/resolve/main/src/train.py
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() | |