multisense_df / src /evaluation /metrics.py
vijay-2132's picture
Clean deployment without large checkpoints
a94bd47
Raw
History Blame Contribute Delete
4.87 kB
"""
MultiSense-DF — Evaluation Metrics
AUC, Accuracy, F1, EER, per-class breakdown
"""
import numpy as np
import torch
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.metrics import (
roc_auc_score, accuracy_score, f1_score,
roc_curve, confusion_matrix, classification_report
)
from pathlib import Path
def compute_eer(y_true, y_score):
"""Equal Error Rate — standard audio anti-spoofing metric."""
fpr, tpr, thresholds = roc_curve(y_true, y_score, pos_label=1)
fnr = 1 - tpr
eer_idx = np.nanargmin(np.abs(fnr - fpr))
eer = float(np.mean([fpr[eer_idx], fnr[eer_idx]]))
return eer, thresholds[eer_idx]
def evaluate_model(model, dataloader, device='cuda', save_dir='results'):
"""
Run model over dataloader and compute full evaluation metrics.
Returns dict of metrics and saves plots.
"""
model.eval()
all_probs, all_labels = [], []
per_mod_probs = {'visual': [], 'audio': [], 'lipsync': []}
with torch.no_grad():
for batch in dataloader:
frames = batch['frames'].to(device)
waveform = batch['waveform'].to(device)
mouth_crops = batch['mouth_crops'].to(device)
mel_specs = batch['mel_specs'].to(device)
labels = batch['label']
outputs = model(frames, waveform, mouth_crops, mel_specs)
global_prob = torch.sigmoid(outputs['global_logit']).cpu().numpy()
all_probs.extend(global_prob.flatten())
all_labels.extend(labels.numpy())
for mod in ['visual', 'audio', 'lipsync']:
p = torch.sigmoid(outputs['per_mod_logits'][mod]).cpu().numpy()
per_mod_probs[mod].extend(p.flatten())
y = np.array(all_labels)
p = np.array(all_probs)
preds = (p >= 0.5).astype(int)
auc = roc_auc_score(y, p)
acc = accuracy_score(y, preds)
f1 = f1_score(y, preds)
eer, eer_thresh = compute_eer(y, p)
per_mod_metrics = {}
for mod, mp in per_mod_probs.items():
mp = np.array(mp)
per_mod_metrics[mod] = {
'auc': roc_auc_score(y, mp),
'acc': accuracy_score(y, (mp >= 0.5).astype(int)),
}
print(f'\n{"="*50}')
print(f' MultiSense-DF Evaluation Results')
print(f'{"="*50}')
print(f' AUC-ROC : {auc:.4f}')
print(f' Accuracy : {acc:.4f}')
print(f' F1 Score : {f1:.4f}')
print(f' EER : {eer:.4f} (threshold={eer_thresh:.3f})')
print(f'\n Per-modality AUC:')
for mod, m in per_mod_metrics.items():
print(f' {mod:8s} → AUC={m["auc"]:.4f} Acc={m["acc"]:.4f}')
print(f'{"="*50}\n')
# Save plots
save_dir = Path(save_dir)
save_dir.mkdir(parents=True, exist_ok=True)
_plot_roc(y, p, auc, save_dir)
_plot_confusion(y, preds, save_dir)
_plot_per_mod(per_mod_metrics, auc, save_dir)
return {
'auc': auc, 'accuracy': acc, 'f1': f1, 'eer': eer,
'per_modality': per_mod_metrics
}
def _plot_roc(y, probs, auc, save_dir):
fpr, tpr, _ = roc_curve(y, probs)
fig, ax = plt.subplots(figsize=(6, 5))
ax.plot(fpr, tpr, lw=2, label=f'MultiSense-DF (AUC={auc:.3f})', color='#6C63FF')
ax.plot([0, 1], [0, 1], '--', color='gray', lw=1)
ax.set(xlabel='False Positive Rate', ylabel='True Positive Rate',
title='ROC Curve — MultiSense-DF')
ax.legend()
plt.tight_layout()
plt.savefig(save_dir / 'roc_curve.png', dpi=150)
plt.close()
def _plot_confusion(y, preds, save_dir):
cm = confusion_matrix(y, preds)
fig, ax = plt.subplots(figsize=(4, 4))
sns.heatmap(cm, annot=True, fmt='d', cmap='Purples',
xticklabels=['Real', 'Fake'],
yticklabels=['Real', 'Fake'], ax=ax)
ax.set(xlabel='Predicted', ylabel='True', title='Confusion Matrix')
plt.tight_layout()
plt.savefig(save_dir / 'confusion_matrix.png', dpi=150)
plt.close()
def _plot_per_mod(per_mod_metrics, global_auc, save_dir):
labels = ['Visual', 'Audio', 'Lip-Sync', 'Fusion\n(Global)']
aucs = [per_mod_metrics['visual']['auc'],
per_mod_metrics['audio']['auc'],
per_mod_metrics['lipsync']['auc'],
global_auc]
colors = ['#FF6B6B', '#4ECDC4', '#45B7D1', '#6C63FF']
fig, ax = plt.subplots(figsize=(7, 4))
bars = ax.bar(labels, aucs, color=colors, edgecolor='white', width=0.5)
ax.set_ylim(0.5, 1.0)
ax.set_ylabel('AUC-ROC')
ax.set_title('Per-Modality vs Fusion AUC Comparison')
for bar, val in zip(bars, aucs):
ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.005,
f'{val:.3f}', ha='center', va='bottom', fontsize=10, fontweight='bold')
plt.tight_layout()
plt.savefig(save_dir / 'per_modality_auc.png', dpi=150)
plt.close()
print(f' Plots saved to {save_dir}')