| import argparse | |
| from pathlib import Path | |
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import DataLoader | |
| from torchvision import transforms | |
| from sklearn.metrics import ( | |
| accuracy_score, | |
| precision_score, | |
| recall_score, | |
| f1_score, | |
| roc_auc_score, | |
| confusion_matrix, | |
| ) | |
| from video_data import VideoDataset | |
| from video_model import ResNetLSTM | |
| def build_transforms(): | |
| return transforms.Compose( | |
| [ | |
| transforms.Resize((224, 224)), | |
| transforms.ToTensor(), | |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), | |
| ] | |
| ) | |
| def accuracy_from_logits(logits, labels): | |
| preds = torch.argmax(logits, dim=1) | |
| return (preds == labels).float().mean().item() | |
| def compute_classification_metrics(y_true, y_pred, y_prob_fake): | |
| cm = confusion_matrix(y_true, y_pred, labels=[0, 1]) | |
| tn, fp, fn, tp = cm.ravel() | |
| metrics = { | |
| "accuracy": accuracy_score(y_true, y_pred), | |
| "precision_real": precision_score(y_true, y_pred, pos_label=0, zero_division=0), | |
| "precision_fake": precision_score(y_true, y_pred, pos_label=1, zero_division=0), | |
| "recall_real": recall_score(y_true, y_pred, pos_label=0, zero_division=0), | |
| "recall_fake": recall_score(y_true, y_pred, pos_label=1, zero_division=0), | |
| "f1_real": f1_score(y_true, y_pred, pos_label=0, zero_division=0), | |
| "f1_fake": f1_score(y_true, y_pred, pos_label=1, zero_division=0), | |
| "f1_macro": f1_score(y_true, y_pred, average="macro", zero_division=0), | |
| "f1_weighted": f1_score(y_true, y_pred, average="weighted", zero_division=0), | |
| "tn": int(tn), | |
| "fp": int(fp), | |
| "fn": int(fn), | |
| "tp": int(tp), | |
| "specificity": tn / (tn + fp) if (tn + fp) > 0 else 0.0, | |
| "sensitivity": tp / (tp + fn) if (tp + fn) > 0 else 0.0, | |
| } | |
| try: | |
| metrics["auc_roc"] = roc_auc_score(y_true, y_prob_fake) | |
| except Exception: | |
| metrics["auc_roc"] = float("nan") | |
| return metrics | |
| def train_one_epoch(model, loader, optimizer, criterion, device, frame_loss_weight): | |
| model.train() | |
| total_loss = 0.0 | |
| total_acc = 0.0 | |
| for videos, labels in loader: | |
| videos = videos.to(device) | |
| labels = labels.to(device) | |
| frame_logits, video_logits = model(videos) | |
| loss_video = criterion(video_logits, labels) | |
| b, t, c = frame_logits.shape | |
| frame_targets = labels.repeat_interleave(t) | |
| loss_frame = criterion(frame_logits.view(b * t, c), frame_targets) | |
| loss = loss_video + frame_loss_weight * loss_frame | |
| optimizer.zero_grad() | |
| loss.backward() | |
| optimizer.step() | |
| total_loss += loss.item() | |
| total_acc += accuracy_from_logits(video_logits, labels) | |
| return total_loss / max(1, len(loader)), total_acc / max(1, len(loader)) | |
| def evaluate(model, loader, criterion, device): | |
| model.eval() | |
| total_loss = 0.0 | |
| total_acc = 0.0 | |
| all_labels = [] | |
| all_preds = [] | |
| all_prob_fake = [] | |
| with torch.no_grad(): | |
| for videos, labels in loader: | |
| videos = videos.to(device) | |
| labels = labels.to(device) | |
| _, video_logits = model(videos) | |
| loss = criterion(video_logits, labels) | |
| probs = torch.softmax(video_logits, dim=1) | |
| preds = torch.argmax(video_logits, dim=1) | |
| total_loss += loss.item() | |
| total_acc += accuracy_from_logits(video_logits, labels) | |
| all_labels.extend(labels.cpu().tolist()) | |
| all_preds.extend(preds.cpu().tolist()) | |
| all_prob_fake.extend(probs[:, 1].cpu().tolist()) | |
| metrics = compute_classification_metrics(all_labels, all_preds, all_prob_fake) | |
| return total_loss / max(1, len(loader)), total_acc / max(1, len(loader)), metrics | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--dataset", type=str, default="SDFVD/SDFVD") | |
| parser.add_argument("--out", type=str, default="video_xception_lstm.pt") | |
| parser.add_argument("--epochs", type=int, default=50) | |
| parser.add_argument("--resume", type=str, default=None) | |
| parser.add_argument("--override_config", action="store_true") | |
| parser.add_argument("--batch_size", type=int, default=2) | |
| parser.add_argument("--lr", type=float, default=1e-4) | |
| parser.add_argument("--frames_per_video", type=int, default=16) | |
| parser.add_argument("--frame_stride", type=int, default=4) | |
| parser.add_argument("--max_videos_per_class", type=int, default=None) | |
| parser.add_argument("--val_split", type=float, default=0.2) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--num_workers", type=int, default=0) | |
| parser.add_argument("--no_face_detection", action="store_true") | |
| parser.add_argument("--frame_loss_weight", type=float, default=0.3) | |
| parser.add_argument("--temporal_pool", type=str, choices=["mean", "last"], default="mean") | |
| parser.add_argument("--backbone", type=str, choices=["resnet50", "xception"], default="xception") | |
| parser.add_argument("--hidden_size", type=int, default=256) | |
| parser.add_argument("--num_layers", type=int, default=1) | |
| parser.add_argument("--no_bidirectional", action="store_false", dest="bidirectional") | |
| parser.add_argument("--no_pretrained", action="store_true") | |
| parser.add_argument("--early_stop_patience", type=int, default=8) | |
| parser.add_argument("--early_stop_min_delta", type=float, default=1e-4) | |
| parser.add_argument("--cpu", action="store_true") | |
| args = parser.parse_args() | |
| device = torch.device("cpu" if args.cpu or not torch.cuda.is_available() else "cuda") | |
| resume_path = Path(args.resume) if args.resume else None | |
| checkpoint = None | |
| if resume_path and resume_path.exists(): | |
| checkpoint = torch.load(resume_path, map_location="cpu") | |
| if not args.override_config and isinstance(checkpoint, dict) and "config" in checkpoint: | |
| cfg = checkpoint["config"] | |
| args.hidden_size = cfg.get("hidden_size", args.hidden_size) | |
| args.num_layers = cfg.get("num_layers", args.num_layers) | |
| args.bidirectional = cfg.get("bidirectional", args.bidirectional) | |
| args.temporal_pool = cfg.get("temporal_pool", args.temporal_pool) | |
| args.backbone = cfg.get("backbone", args.backbone) | |
| args.no_pretrained = not cfg.get("pretrained", not args.no_pretrained) | |
| args.frames_per_video = cfg.get("frames_per_video", args.frames_per_video) | |
| args.frame_stride = cfg.get("frame_stride", args.frame_stride) | |
| args.no_face_detection = not cfg.get("face_detection", not args.no_face_detection) | |
| transform = build_transforms() | |
| train_ds = VideoDataset( | |
| root_dir=args.dataset, | |
| split="train", | |
| val_split=args.val_split, | |
| seed=args.seed, | |
| max_videos_per_class=args.max_videos_per_class, | |
| frames_per_video=args.frames_per_video, | |
| frame_stride=args.frame_stride, | |
| face_detection=not args.no_face_detection, | |
| transform=transform, | |
| ) | |
| val_ds = VideoDataset( | |
| root_dir=args.dataset, | |
| split="val", | |
| val_split=args.val_split, | |
| seed=args.seed, | |
| max_videos_per_class=args.max_videos_per_class, | |
| frames_per_video=args.frames_per_video, | |
| frame_stride=args.frame_stride, | |
| face_detection=not args.no_face_detection, | |
| transform=transform, | |
| ) | |
| train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers) | |
| val_loader = DataLoader(val_ds, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers) | |
| model = ResNetLSTM( | |
| hidden_size=args.hidden_size, | |
| num_layers=args.num_layers, | |
| bidirectional=args.bidirectional, | |
| temporal_pool=args.temporal_pool, | |
| pretrained=not args.no_pretrained, | |
| backbone_name=args.backbone, | |
| ) | |
| model.to(device) | |
| if checkpoint and "model_state" in checkpoint: | |
| model.load_state_dict(checkpoint["model_state"], strict=True) | |
| print(f"Resumed model weights from {resume_path}") | |
| criterion = nn.CrossEntropyLoss() | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr) | |
| best_val_acc = 0.0 | |
| start_epoch = 0 | |
| if checkpoint: | |
| best_val_acc = float(checkpoint.get("best_val_acc", best_val_acc)) | |
| start_epoch = int(checkpoint.get("epoch", start_epoch)) | |
| out_path = Path(args.out) | |
| out_path.parent.mkdir(parents=True, exist_ok=True) | |
| total_epochs = start_epoch + args.epochs | |
| no_improve_epochs = 0 | |
| for epoch in range(args.epochs): | |
| train_loss, train_acc = train_one_epoch( | |
| model, train_loader, optimizer, criterion, device, args.frame_loss_weight | |
| ) | |
| val_loss, val_acc, val_metrics = evaluate(model, val_loader, criterion, device) | |
| print( | |
| f"Epoch {start_epoch + epoch + 1}/{total_epochs} | " | |
| f"Train Loss: {train_loss:.4f} Acc: {train_acc:.3f} | " | |
| f"Val Loss: {val_loss:.4f} Acc: {val_acc:.3f}" | |
| ) | |
| print( | |
| " Accuracy | " | |
| f"Train Accuracy: {train_acc:.4f} | " | |
| f"Validation Accuracy: {val_acc:.4f}" | |
| ) | |
| print( | |
| " Val Metrics | " | |
| f"AUC: {val_metrics['auc_roc']:.4f} | " | |
| f"F1(macro): {val_metrics['f1_macro']:.4f} | " | |
| f"F1(weighted): {val_metrics['f1_weighted']:.4f}" | |
| ) | |
| print( | |
| " Real Class | " | |
| f"Precision: {val_metrics['precision_real']:.4f} | " | |
| f"Recall: {val_metrics['recall_real']:.4f} | " | |
| f"F1: {val_metrics['f1_real']:.4f}" | |
| ) | |
| print( | |
| " Fake Class | " | |
| f"Precision: {val_metrics['precision_fake']:.4f} | " | |
| f"Recall: {val_metrics['recall_fake']:.4f} | " | |
| f"F1: {val_metrics['f1_fake']:.4f}" | |
| ) | |
| print( | |
| " Derived | " | |
| f"Sensitivity(TPR): {val_metrics['sensitivity']:.4f} | " | |
| f"Specificity(TNR): {val_metrics['specificity']:.4f}" | |
| ) | |
| print( | |
| " Confusion | " | |
| f"TN: {val_metrics['tn']} FP: {val_metrics['fp']} " | |
| f"FN: {val_metrics['fn']} TP: {val_metrics['tp']}" | |
| ) | |
| if val_acc >= best_val_acc + args.early_stop_min_delta: | |
| best_val_acc = val_acc | |
| no_improve_epochs = 0 | |
| torch.save( | |
| { | |
| "model_state": model.state_dict(), | |
| "epoch": start_epoch + epoch + 1, | |
| "best_val_acc": best_val_acc, | |
| "config": { | |
| "hidden_size": args.hidden_size, | |
| "num_layers": args.num_layers, | |
| "bidirectional": args.bidirectional, | |
| "temporal_pool": args.temporal_pool, | |
| "backbone": args.backbone, | |
| "pretrained": not args.no_pretrained, | |
| "frames_per_video": args.frames_per_video, | |
| "frame_stride": args.frame_stride, | |
| "face_detection": not args.no_face_detection, | |
| }, | |
| }, | |
| out_path, | |
| ) | |
| print(f"Saved best model to {out_path}") | |
| else: | |
| no_improve_epochs += 1 | |
| if no_improve_epochs >= args.early_stop_patience: | |
| print( | |
| "Early stopping triggered: " | |
| f"no validation improvement for {args.early_stop_patience} epoch(s)." | |
| ) | |
| break | |
| if __name__ == "__main__": | |
| main() | |