crowncode-backend / app /training /visualize_results.py
Rthur2003's picture
feat: enhance visualization pipeline with additional figures and manifest support
b06ac73
Raw History Blame Contribute Delete
40.6 kB
"""
Publication-quality visualization pipeline for AURIS.
Generates all figures required for an academic paper / conference
submission on AI-generated music detection:
1. ROC curves (per-model overlay)
2. Precision-Recall curves (per-model overlay)
3. Confusion matrices (heatmap per model)
4. Model comparison bar chart (Accuracy, F1, AUC side-by-side)
5. Feature importance (top-N horizontal bar)
6. Feature correlation heatmap
7. Feature distribution violin plots (AI vs Human)
8. Training summary table (LaTeX-ready)
Usage:
python -m app.training.visualize_results \\
--results models/training_results.json \\
--features data/training/features.csv \\
--output figures/
All figures are saved at 300 DPI in both PNG and PDF formats
for direct inclusion in LaTeX / Word documents.
"""
from __future__ import annotations
import argparse
import csv
import json
import sys
from pathlib import Path
from typing import Any
import numpy as np
try:
import matplotlib
matplotlib.use("Agg") # Non-interactive backend for server/CI
import matplotlib.pyplot as plt
import matplotlib.ticker as mticker
from matplotlib.colors import LinearSegmentedColormap
except ImportError:
print("ERROR: matplotlib required. pip install matplotlib")
sys.exit(1)
try:
import seaborn as sns
HAS_SEABORN = True
except ImportError:
HAS_SEABORN = False
from sklearn.metrics import (
auc,
confusion_matrix,
precision_recall_curve,
roc_curve,
)
# ═══════════════════════════════════════════════════════════════════════
# Style configuration — academic paper quality
# ═══════════════════════════════════════════════════════════════════════
# Color palette — distinct, colorblind-friendly
MODEL_COLORS = {
"Logistic Regression": "#4363d8",
"Random Forest": "#3cb44b",
"Gradient Boosting": "#e6194b",
"SVM (RBF)": "#f58231",
"MLP Neural Network": "#911eb4",
"XGBoost": "#42d4f4",
"LightGBM": "#f032e6",
"Heuristic (no vocals)": "#808080",
"Heuristic + Vocals": "#a9a9a9",
}
AURIS_BLUE = "#1a73e8"
AURIS_RED = "#e8431a"
plt.rcParams.update({
"font.family": "serif",
"font.size": 11,
"axes.titlesize": 13,
"axes.labelsize": 12,
"xtick.labelsize": 10,
"ytick.labelsize": 10,
"legend.fontsize": 9,
"figure.dpi": 150,
"savefig.dpi": 300,
"savefig.bbox": "tight",
"savefig.pad_inches": 0.1,
"axes.grid": True,
"grid.alpha": 0.3,
"axes.spines.top": False,
"axes.spines.right": False,
})
def _save_fig(fig: plt.Figure, output_dir: Path, name: str) -> None:
"""Save figure in both PNG and PDF formats."""
fig.savefig(output_dir / f"{name}.png", format="png")
fig.savefig(output_dir / f"{name}.pdf", format="pdf")
plt.close(fig)
print(f" Saved: {name}.png / .pdf")
def _get_color(name: str) -> str:
return MODEL_COLORS.get(name, "#333333")
# ═══════════════════════════════════════════════════════════════════════
# Figure 1: ROC Curves
# ═══════════════════════════════════════════════════════════════════════
def plot_roc_curves(
results: dict[str, Any],
output_dir: Path,
) -> None:
"""Plot ROC curves for all models on the same axes."""
fig, ax = plt.subplots(figsize=(7, 6))
for name, data in results.items():
if name.startswith("_"):
continue
y_true = np.array(data["y_true"])
y_prob = np.array(data["y_prob"])
fpr, tpr, _ = roc_curve(y_true, y_prob)
roc_auc = auc(fpr, tpr)
ax.plot(
fpr, tpr,
color=_get_color(name),
linewidth=2,
label=f"{name} (AUC = {roc_auc:.3f})",
)
# Diagonal reference
ax.plot([0, 1], [0, 1], "k--", linewidth=1, alpha=0.5, label="Random (AUC = 0.500)")
ax.set_xlim([-0.02, 1.02])
ax.set_ylim([-0.02, 1.02])
ax.set_xlabel("False Positive Rate")
ax.set_ylabel("True Positive Rate")
ax.set_title("ROC Curves — AURIS Model Comparison")
ax.legend(loc="lower right", framealpha=0.9)
ax.set_aspect("equal")
_save_fig(fig, output_dir, "fig1_roc_curves")
# ═══════════════════════════════════════════════════════════════════════
# Figure 2: Precision-Recall Curves
# ═══════════════════════════════════════════════════════════════════════
def plot_pr_curves(
results: dict[str, Any],
output_dir: Path,
) -> None:
"""Plot Precision-Recall curves for all models."""
fig, ax = plt.subplots(figsize=(7, 6))
for name, data in results.items():
if name.startswith("_"):
continue
y_true = np.array(data["y_true"])
y_prob = np.array(data["y_prob"])
precision, recall, _ = precision_recall_curve(y_true, y_prob)
pr_auc = auc(recall, precision)
ax.plot(
recall, precision,
color=_get_color(name),
linewidth=2,
label=f"{name} (AP = {pr_auc:.3f})",
)
# Baseline: proportion of positives
all_y = []
for name, data in results.items():
if not name.startswith("_"):
all_y = data["y_true"]
break
baseline = np.mean(all_y) if all_y else 0.5
ax.axhline(y=baseline, color="k", linestyle="--", linewidth=1, alpha=0.5,
label=f"Baseline ({baseline:.2f})")
ax.set_xlim([-0.02, 1.02])
ax.set_ylim([-0.02, 1.05])
ax.set_xlabel("Recall")
ax.set_ylabel("Precision")
ax.set_title("Precision-Recall Curves — AURIS Model Comparison")
ax.legend(loc="lower left", framealpha=0.9)
_save_fig(fig, output_dir, "fig2_pr_curves")
# ═══════════════════════════════════════════════════════════════════════
# Figure 3: Confusion Matrices
# ═══════════════════════════════════════════════════════════════════════
def plot_confusion_matrices(
results: dict[str, Any],
output_dir: Path,
) -> None:
"""Plot confusion matrix heatmap for each model."""
model_names = [k for k in results if not k.startswith("_")]
n_models = len(model_names)
cols = min(3, n_models)
rows = (n_models + cols - 1) // cols
fig, axes = plt.subplots(rows, cols, figsize=(5 * cols, 4.5 * rows))
if n_models == 1:
axes = np.array([axes])
axes = axes.flatten()
cmap = LinearSegmentedColormap.from_list("auris", ["#ffffff", AURIS_BLUE])
for idx, name in enumerate(model_names):
ax = axes[idx]
data = results[name]
y_true = np.array(data["y_true"])
y_pred = np.array(data["y_pred"])
cm = confusion_matrix(y_true, y_pred)
cm_norm = cm.astype(float) / cm.sum(axis=1, keepdims=True)
im = ax.imshow(cm_norm, interpolation="nearest", cmap=cmap, vmin=0, vmax=1)
# Annotate cells with count and percentage
for i in range(2):
for j in range(2):
color = "white" if cm_norm[i, j] > 0.6 else "black"
ax.text(j, i, f"{cm[i, j]}\n({cm_norm[i, j]:.1%})",
ha="center", va="center", fontsize=12, color=color,
fontweight="bold")
ax.set_xticks([0, 1])
ax.set_yticks([0, 1])
ax.set_xticklabels(["Human", "AI"])
ax.set_yticklabels(["Human", "AI"])
ax.set_xlabel("Predicted")
ax.set_ylabel("Actual")
ax.set_title(name, fontsize=11)
# Hide unused axes
for idx in range(n_models, len(axes)):
axes[idx].set_visible(False)
fig.suptitle("Confusion Matrices — AURIS Model Comparison", fontsize=14, y=1.02)
fig.tight_layout()
_save_fig(fig, output_dir, "fig3_confusion_matrices")
# ═══════════════════════════════════════════════════════════════════════
# Figure 4: Model Comparison Bar Chart
# ═══════════════════════════════════════════════════════════════════════
def plot_model_comparison(
results: dict[str, Any],
output_dir: Path,
) -> None:
"""Bar chart comparing Accuracy, F1, Precision, Recall, AUC across models."""
model_names = [k for k in results if not k.startswith("_")]
metrics = ["accuracy", "precision", "recall", "f1", "roc_auc"]
metric_labels = ["Accuracy", "Precision", "Recall", "F1 Score", "ROC-AUC"]
x = np.arange(len(model_names))
width = 0.15
metric_colors = ["#4363d8", "#3cb44b", "#e6194b", "#f58231", "#911eb4"]
fig, ax = plt.subplots(figsize=(max(10, len(model_names) * 2), 6))
for i, (metric, label, color) in enumerate(zip(metrics, metric_labels, metric_colors)):
values = []
for name in model_names:
val = results[name].get(metric, 0)
values.append(val if val is not None else 0)
bars = ax.bar(x + i * width, values, width, label=label, color=color, alpha=0.85)
# Value labels on bars
for bar, val in zip(bars, values):
ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.01,
f"{val:.2f}", ha="center", va="bottom", fontsize=7)
ax.set_xlabel("Model")
ax.set_ylabel("Score")
ax.set_title("Model Performance Comparison — AURIS")
ax.set_xticks(x + width * 2)
ax.set_xticklabels(model_names, rotation=25, ha="right")
ax.set_ylim([0, 1.12])
ax.legend(loc="upper right", ncol=5, framealpha=0.9)
ax.yaxis.set_major_formatter(mticker.PercentFormatter(1.0))
fig.tight_layout()
_save_fig(fig, output_dir, "fig4_model_comparison")
# ═══════════════════════════════════════════════════════════════════════
# Figure 5: Feature Importance
# ═══════════════════════════════════════════════════════════════════════
def plot_feature_importance(
results: dict[str, Any],
output_dir: Path,
top_n: int = 20,
) -> None:
"""Horizontal bar chart of top-N feature importances."""
importance = results.get("_feature_importance")
if not importance:
print(" Skipping feature importance — no data available")
return
# Sort and take top N
sorted_features = sorted(importance.items(), key=lambda x: x[1], reverse=True)[:top_n]
names = [f[0] for f in reversed(sorted_features)]
values = [f[1] for f in reversed(sorted_features)]
fig, ax = plt.subplots(figsize=(8, max(6, top_n * 0.35)))
colors = plt.cm.Blues(np.linspace(0.3, 0.9, len(names)))
bars = ax.barh(names, values, color=colors, edgecolor="white", linewidth=0.5)
# Value labels
for bar, val in zip(bars, values):
ax.text(bar.get_width() + 0.002, bar.get_y() + bar.get_height() / 2,
f"{val:.4f}", ha="left", va="center", fontsize=9)
ax.set_xlabel("Relative Importance")
ax.set_title(f"Top {top_n} Feature Importances — AURIS ({results.get('_best_model', '')})")
ax.set_xlim([0, max(values) * 1.15])
fig.tight_layout()
_save_fig(fig, output_dir, "fig5_feature_importance")
# ═══════════════════════════════════════════════════════════════════════
# Figure 6: Feature Correlation Heatmap
# ═══════════════════════════════════════════════════════════════════════
def plot_correlation_heatmap(
features_csv: Path,
output_dir: Path,
) -> None:
"""Correlation heatmap of all features."""
X, y, feature_cols = _load_features_with_names(features_csv)
corr = np.corrcoef(X.T)
fig, ax = plt.subplots(figsize=(max(12, len(feature_cols) * 0.4),
max(10, len(feature_cols) * 0.35)))
cmap = "RdBu_r" if HAS_SEABORN else "coolwarm"
if HAS_SEABORN:
sns.heatmap(
corr, xticklabels=feature_cols, yticklabels=feature_cols,
cmap=cmap, center=0, vmin=-1, vmax=1,
square=True, linewidths=0.5, ax=ax,
cbar_kws={"shrink": 0.8, "label": "Pearson Correlation"},
)
else:
im = ax.imshow(corr, cmap=cmap, vmin=-1, vmax=1, aspect="auto")
ax.set_xticks(range(len(feature_cols)))
ax.set_yticks(range(len(feature_cols)))
ax.set_xticklabels(feature_cols, rotation=90, fontsize=7)
ax.set_yticklabels(feature_cols, fontsize=7)
fig.colorbar(im, ax=ax, shrink=0.8, label="Pearson Correlation")
ax.set_title("Feature Correlation Matrix — AURIS", fontsize=14, pad=20)
fig.tight_layout()
_save_fig(fig, output_dir, "fig6_correlation_heatmap")
# ═══════════════════════════════════════════════════════════════════════
# Figure 7: Feature Distribution (Violin / Box Plots)
# ═══════════════════════════════════════════════════════════════════════
def plot_feature_distributions(
features_csv: Path,
output_dir: Path,
results: dict[str, Any] | None = None,
top_n: int = 12,
) -> None:
"""Violin plots showing feature distributions for AI vs Human."""
X, y, feature_cols = _load_features_with_names(features_csv)
# Select top features by importance, or first N
if results and "_feature_importance" in results:
sorted_feats = sorted(
results["_feature_importance"].items(),
key=lambda x: x[1], reverse=True,
)
selected = [f[0] for f in sorted_feats[:top_n] if f[0] in feature_cols]
else:
selected = feature_cols[:top_n]
n_features = len(selected)
cols = 3
rows = (n_features + cols - 1) // cols
fig, axes = plt.subplots(rows, cols, figsize=(5 * cols, 3.5 * rows))
axes = axes.flatten()
for idx, feat_name in enumerate(selected):
ax = axes[idx]
col_idx = feature_cols.index(feat_name)
human_vals = X[y == 0, col_idx]
ai_vals = X[y == 1, col_idx]
if HAS_SEABORN:
data_list = []
for val in human_vals:
data_list.append({"Feature": feat_name, "Value": val, "Class": "Human"})
for val in ai_vals:
data_list.append({"Feature": feat_name, "Value": val, "Class": "AI"})
import pandas as pd
df = pd.DataFrame(data_list)
sns.violinplot(
data=df, x="Class", y="Value",
palette={"Human": AURIS_BLUE, "AI": AURIS_RED},
ax=ax, inner="quartile", linewidth=1,
)
else:
parts = ax.violinplot(
[human_vals, ai_vals],
positions=[0, 1],
showmeans=True, showmedians=True,
)
ax.set_xticks([0, 1])
ax.set_xticklabels(["Human", "AI"])
ax.set_title(feat_name, fontsize=10)
ax.set_xlabel("")
for idx in range(n_features, len(axes)):
axes[idx].set_visible(False)
fig.suptitle("Feature Distributions — AI vs Human", fontsize=14, y=1.01)
fig.tight_layout()
_save_fig(fig, output_dir, "fig7_feature_distributions")
# ═══════════════════════════════════════════════════════════════════════
# Figure 8: Training Summary Table (LaTeX)
# ═══════════════════════════════════════════════════════════════════════
def generate_latex_table(
results: dict[str, Any],
output_dir: Path,
) -> None:
"""Generate LaTeX-ready comparison table."""
model_names = [k for k in results if not k.startswith("_")]
best_model = results.get("_best_model", "")
lines = [
r"\begin{table}[htbp]",
r"\centering",
r"\caption{AURIS Model Performance Comparison}",
r"\label{tab:model-comparison}",
r"\begin{tabular}{lccccc}",
r"\toprule",
r"Model & Accuracy & Precision & Recall & F1 & ROC-AUC \\",
r"\midrule",
]
for name in model_names:
data = results[name]
acc = data.get("accuracy", 0)
prec = data.get("precision", 0)
rec = data.get("recall", 0)
f1 = data.get("f1", 0)
roc = data.get("roc_auc", 0)
# Bold the best model
prefix = r"\textbf{" if name == best_model else ""
suffix = "}" if name == best_model else ""
row = (
f" {prefix}{name}{suffix} & "
f"{prefix}{acc:.4f}{suffix} & "
f"{prefix}{prec:.4f}{suffix} & "
f"{prefix}{rec:.4f}{suffix} & "
f"{prefix}{f1:.4f}{suffix} & "
f"{prefix}{roc:.4f}{suffix} \\\\"
)
lines.append(row)
lines.extend([
r"\bottomrule",
r"\end{tabular}",
r"\end{table}",
])
latex_content = "\n".join(lines)
table_path = output_dir / "table1_model_comparison.tex"
table_path.write_text(latex_content, encoding="utf-8")
print(f" Saved: table1_model_comparison.tex")
# Also save as markdown for README
md_lines = [
"| Model | Accuracy | Precision | Recall | F1 | ROC-AUC |",
"|-------|----------|-----------|--------|-----|---------|",
]
for name in model_names:
data = results[name]
bold = "**" if name == best_model else ""
md_lines.append(
f"| {bold}{name}{bold} | "
f"{data.get('accuracy', 0):.4f} | "
f"{data.get('precision', 0):.4f} | "
f"{data.get('recall', 0):.4f} | "
f"{data.get('f1', 0):.4f} | "
f"{data.get('roc_auc', 0):.4f} |"
)
md_path = output_dir / "table1_model_comparison.md"
md_path.write_text("\n".join(md_lines), encoding="utf-8")
print(f" Saved: table1_model_comparison.md")
# ═══════════════════════════════════════════════════════════════════════
# Utilities
# ═══════════════════════════════════════════════════════════════════════
def _load_features_with_names(
features_csv: Path,
) -> tuple[np.ndarray, np.ndarray, list[str]]:
"""Load features CSV returning X, y, and column names."""
rows = []
labels = []
with open(features_csv, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
feature_cols = [
c for c in reader.fieldnames
if c not in ("file_path", "label_int")
]
for row in reader:
feat_values = []
for col in feature_cols:
try:
feat_values.append(float(row[col]))
except (ValueError, KeyError):
feat_values.append(0.0)
rows.append(feat_values)
labels.append(int(row["label_int"]))
X = np.nan_to_num(np.array(rows, dtype=np.float32), nan=0.0)
y = np.array(labels, dtype=np.int32)
return X, y, feature_cols
def _load_manifest(manifest_csv: Path) -> list[dict]:
"""Load manifest.csv rows as dicts."""
if not manifest_csv.exists():
return []
with open(manifest_csv, "r", encoding="utf-8") as f:
return list(csv.DictReader(f))
# ═══════════════════════════════════════════════════════════════════════
# Figure 9: Dataset Statistics
# ═══════════════════════════════════════════════════════════════════════
def plot_dataset_statistics(
manifest_csv: Path,
output_dir: Path,
) -> None:
"""Dataset composition: class distribution, source distribution, duration."""
rows = _load_manifest(manifest_csv)
if not rows:
print(" Skipping dataset stats — no manifest")
return
from collections import Counter
labels = [r.get("label", "unknown") for r in rows]
sources = [r.get("generator", "unknown") for r in rows]
durations = []
for r in rows:
try:
durations.append(float(r.get("duration_sec", 0)))
except ValueError:
continue
label_counts = Counter(labels)
source_counts = Counter(sources)
fig, axes = plt.subplots(1, 3, figsize=(16, 5))
# (a) Class distribution
ax = axes[0]
colors = [AURIS_BLUE if l == "human" else AURIS_RED for l in label_counts.keys()]
bars = ax.bar(label_counts.keys(), label_counts.values(),
color=colors, edgecolor="black", linewidth=0.5)
for bar, val in zip(bars, label_counts.values()):
pct = val / sum(label_counts.values()) * 100
ax.text(bar.get_x() + bar.get_width() / 2,
bar.get_height() + max(label_counts.values()) * 0.02,
f"{val:,}\n({pct:.1f}%)",
ha="center", va="bottom", fontsize=10)
ax.set_title("Class Distribution")
ax.set_ylabel("Samples")
ax.set_ylim([0, max(label_counts.values()) * 1.18])
# (b) Source distribution
ax = axes[1]
sorted_sources = sorted(source_counts.items(), key=lambda x: x[1], reverse=True)
src_names = [s[0] for s in sorted_sources]
src_values = [s[1] for s in sorted_sources]
bars = ax.barh(src_names, src_values, color=plt.cm.tab10.colors[:len(src_names)])
for bar, val in zip(bars, src_values):
ax.text(bar.get_width() + max(src_values) * 0.01,
bar.get_y() + bar.get_height() / 2,
f"{val:,}", ha="left", va="center", fontsize=9)
ax.set_title("Dataset Source Distribution")
ax.set_xlabel("Samples")
ax.set_xlim([0, max(src_values) * 1.15])
ax.invert_yaxis()
# (c) Duration histogram
ax = axes[2]
if durations:
ax.hist(durations, bins=40, color=AURIS_BLUE, edgecolor="white", alpha=0.8)
ax.axvline(np.median(durations), color="red", linestyle="--",
label=f"Median: {np.median(durations):.1f}s")
ax.axvline(np.mean(durations), color="orange", linestyle="--",
label=f"Mean: {np.mean(durations):.1f}s")
ax.legend()
ax.set_title("Audio Duration Distribution")
ax.set_xlabel("Duration (seconds)")
ax.set_ylabel("Count")
fig.suptitle(
f"AURIS Dataset — {sum(label_counts.values()):,} samples",
fontsize=14, y=1.02,
)
fig.tight_layout()
_save_fig(fig, output_dir, "fig9_dataset_statistics")
# ═══════════════════════════════════════════════════════════════════════
# Figure 10: Per-Source Class Balance
# ═══════════════════════════════════════════════════════════════════════
def plot_per_source_balance(
manifest_csv: Path,
output_dir: Path,
) -> None:
"""Stacked bar: AI vs Human count per source."""
rows = _load_manifest(manifest_csv)
if not rows:
return
from collections import defaultdict
source_label_counts: dict[str, dict[str, int]] = defaultdict(
lambda: {"ai": 0, "human": 0}
)
for r in rows:
src = r.get("generator", "unknown")
lbl = r.get("label", "unknown")
if lbl in ("ai", "human"):
source_label_counts[src][lbl] += 1
sources = sorted(source_label_counts.keys(),
key=lambda s: sum(source_label_counts[s].values()),
reverse=True)
ai_counts = [source_label_counts[s]["ai"] for s in sources]
human_counts = [source_label_counts[s]["human"] for s in sources]
fig, ax = plt.subplots(figsize=(10, 6))
x = np.arange(len(sources))
width = 0.38
b1 = ax.bar(x - width/2, human_counts, width,
label="Human", color=AURIS_BLUE, edgecolor="black", linewidth=0.3)
b2 = ax.bar(x + width/2, ai_counts, width,
label="AI", color=AURIS_RED, edgecolor="black", linewidth=0.3)
for bars in (b1, b2):
for bar in bars:
h = bar.get_height()
if h > 0:
ax.text(bar.get_x() + bar.get_width() / 2, h + 20,
f"{int(h)}", ha="center", va="bottom", fontsize=8)
ax.set_xticks(x)
ax.set_xticklabels(sources, rotation=20, ha="right")
ax.set_ylabel("Samples")
ax.set_title("Class Balance per Source")
ax.legend(loc="upper right")
fig.tight_layout()
_save_fig(fig, output_dir, "fig10_per_source_balance")
# ═══════════════════════════════════════════════════════════════════════
# Figure 11: t-SNE / PCA Feature Embedding
# ═══════════════════════════════════════════════════════════════════════
def plot_feature_embedding(
features_csv: Path,
output_dir: Path,
max_samples: int = 2000,
) -> None:
"""2D embedding (PCA + t-SNE) of features colored by class."""
try:
from sklearn.decomposition import PCA
from sklearn.manifold import TSNE
from sklearn.preprocessing import StandardScaler
except ImportError:
print(" Skipping embedding — sklearn not available")
return
X, y, _ = _load_features_with_names(features_csv)
if len(X) > max_samples:
rng = np.random.default_rng(42)
idx = rng.choice(len(X), max_samples, replace=False)
X = X[idx]
y = y[idx]
X_scaled = StandardScaler().fit_transform(X)
fig, axes = plt.subplots(1, 2, figsize=(14, 6))
# PCA
pca = PCA(n_components=2, random_state=42)
X_pca = pca.fit_transform(X_scaled)
ax = axes[0]
for cls, color, label in [(0, AURIS_BLUE, "Human"), (1, AURIS_RED, "AI")]:
mask = y == cls
ax.scatter(X_pca[mask, 0], X_pca[mask, 1],
c=color, s=12, alpha=0.5, edgecolors="none", label=label)
var_exp = pca.explained_variance_ratio_
ax.set_xlabel(f"PC1 ({var_exp[0] * 100:.1f}%)")
ax.set_ylabel(f"PC2 ({var_exp[1] * 100:.1f}%)")
ax.set_title(f"PCA Projection (total var = {sum(var_exp) * 100:.1f}%)")
ax.legend(loc="best")
# t-SNE
try:
perplexity = min(30, max(5, len(X) // 20))
tsne = TSNE(n_components=2, perplexity=perplexity,
random_state=42, max_iter=500, init="pca")
X_tsne = tsne.fit_transform(X_scaled)
ax = axes[1]
for cls, color, label in [(0, AURIS_BLUE, "Human"), (1, AURIS_RED, "AI")]:
mask = y == cls
ax.scatter(X_tsne[mask, 0], X_tsne[mask, 1],
c=color, s=12, alpha=0.5, edgecolors="none", label=label)
ax.set_xlabel("t-SNE 1")
ax.set_ylabel("t-SNE 2")
ax.set_title(f"t-SNE Projection (perplexity={perplexity})")
ax.legend(loc="best")
except Exception as e:
print(f" t-SNE failed: {e}")
axes[1].set_visible(False)
fig.suptitle("Feature Space Visualization — AI vs Human", fontsize=14, y=1.01)
fig.tight_layout()
_save_fig(fig, output_dir, "fig11_feature_embedding")
# ═══════════════════════════════════════════════════════════════════════
# Figure 12: Probability Calibration
# ═══════════════════════════════════════════════════════════════════════
def plot_calibration_curves(
results: dict[str, Any],
output_dir: Path,
n_bins: int = 10,
) -> None:
"""Reliability diagram: predicted probability vs observed frequency."""
try:
from sklearn.calibration import calibration_curve
except ImportError:
return
fig, ax = plt.subplots(figsize=(7, 6))
for name, data in results.items():
if name.startswith("_"):
continue
y_true = np.array(data.get("y_true", []))
y_prob = np.array(data.get("y_prob", []))
if len(y_true) == 0:
continue
try:
frac_pos, mean_pred = calibration_curve(y_true, y_prob, n_bins=n_bins)
ax.plot(mean_pred, frac_pos, "o-",
color=_get_color(name), label=name, linewidth=1.5, markersize=5)
except Exception:
continue
ax.plot([0, 1], [0, 1], "k--", alpha=0.5, label="Perfect calibration")
ax.set_xlabel("Mean Predicted Probability")
ax.set_ylabel("Fraction of Positives")
ax.set_title("Probability Calibration (Reliability Diagram)")
ax.legend(loc="lower right", fontsize=8)
ax.set_xlim([0, 1])
ax.set_ylim([0, 1])
ax.set_aspect("equal")
fig.tight_layout()
_save_fig(fig, output_dir, "fig12_calibration")
# ═══════════════════════════════════════════════════════════════════════
# Figure 13: Prediction Score Distribution
# ═══════════════════════════════════════════════════════════════════════
def plot_score_distribution(
results: dict[str, Any],
output_dir: Path,
) -> None:
"""Best model's predicted probability distribution per class."""
best_model = results.get("_best_model", "")
if not best_model or best_model not in results:
candidates = [k for k in results if not k.startswith("_")]
if not candidates:
return
best_model = candidates[0]
data = results[best_model]
y_true = np.array(data.get("y_true", []))
y_prob = np.array(data.get("y_prob", []))
if len(y_true) == 0:
return
fig, ax = plt.subplots(figsize=(9, 5))
ax.hist(y_prob[y_true == 0], bins=40, alpha=0.6,
color=AURIS_BLUE, label="Human (true class)", edgecolor="white")
ax.hist(y_prob[y_true == 1], bins=40, alpha=0.6,
color=AURIS_RED, label="AI (true class)", edgecolor="white")
ax.axvline(0.5, color="black", linestyle="--", linewidth=1, alpha=0.7,
label="Decision threshold (0.5)")
ax.set_xlabel(f"Predicted P(AI) — {best_model}")
ax.set_ylabel("Count")
ax.set_title("Prediction Score Distribution — Best Model")
ax.legend(loc="upper center")
fig.tight_layout()
_save_fig(fig, output_dir, "fig13_score_distribution")
# ═══════════════════════════════════════════════════════════════════════
# Extended Summary Statistics Table (Markdown)
# ═══════════════════════════════════════════════════════════════════════
def generate_extended_summary(
results: dict[str, Any],
manifest_csv: Path,
features_csv: Path,
output_dir: Path,
) -> None:
"""Comprehensive markdown summary for the academic report."""
rows = _load_manifest(manifest_csv)
X, y, feature_cols = _load_features_with_names(features_csv)
from collections import Counter
labels = Counter(r.get("label", "?") for r in rows)
sources = Counter(r.get("generator", "?") for r in rows)
model_names = [k for k in results if not k.startswith("_")]
best_model = results.get("_best_model", "")
lines = []
lines.append("# AURIS — Dataset & Results Summary\n")
lines.append("## Dataset\n")
lines.append(f"- **Total samples:** {len(rows):,}")
lines.append(f"- **Features extracted:** {len(feature_cols)}")
lines.append(f"- **Samples with extracted features:** {len(X):,}")
lines.append("")
lines.append("### Class distribution")
lines.append("| Class | Count | % |")
lines.append("|-------|-------|---|")
total = sum(labels.values())
for lbl in sorted(labels):
pct = labels[lbl] / total * 100 if total else 0
lines.append(f"| {lbl} | {labels[lbl]:,} | {pct:.1f}% |")
lines.append("")
lines.append("### Source distribution")
lines.append("| Source | Count | % |")
lines.append("|--------|-------|---|")
for src in sorted(sources, key=lambda s: sources[s], reverse=True):
pct = sources[src] / total * 100 if total else 0
lines.append(f"| {src} | {sources[src]:,} | {pct:.1f}% |")
lines.append("")
lines.append("## Results\n")
lines.append(f"- **Best model:** {best_model}")
n_folds = results.get("_n_folds", "?")
lines.append(f"- **Cross-validation folds:** {n_folds}")
lines.append("")
lines.append("### Model performance")
lines.append("| Model | Accuracy | Precision | Recall | F1 | ROC-AUC |")
lines.append("|-------|----------|-----------|--------|-----|---------|")
for name in model_names:
d = results[name]
bold = "**" if name == best_model else ""
lines.append(
f"| {bold}{name}{bold} | "
f"{d.get('accuracy', 0):.4f} | "
f"{d.get('precision', 0):.4f} | "
f"{d.get('recall', 0):.4f} | "
f"{d.get('f1', 0):.4f} | "
f"{d.get('roc_auc', 0):.4f} |"
)
lines.append("")
if "_feature_importance" in results:
lines.append("### Top-15 Feature Importances")
lines.append("| Rank | Feature | Importance |")
lines.append("|------|---------|------------|")
sorted_imp = sorted(
results["_feature_importance"].items(),
key=lambda x: x[1], reverse=True,
)[:15]
for i, (feat, imp) in enumerate(sorted_imp, 1):
lines.append(f"| {i} | {feat} | {imp:.4f} |")
lines.append("")
output_path = output_dir / "REPORT_SUMMARY.md"
output_path.write_text("\n".join(lines), encoding="utf-8")
print(f" Saved: REPORT_SUMMARY.md")
# ═══════════════════════════════════════════════════════════════════════
# Main entry point
# ═══════════════════════════════════════════════════════════════════════
def generate_all_figures(
results_path: str | Path,
features_csv: str | Path,
output_dir: str | Path = "figures",
manifest_csv: str | Path | None = None,
) -> None:
"""Generate all publication-quality figures."""
results_path = Path(results_path)
features_csv = Path(features_csv)
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
if manifest_csv is None:
manifest_csv = features_csv.parent / "manifest.csv"
manifest_csv = Path(manifest_csv)
print(f"\nLoading results from {results_path}...", flush=True)
with open(results_path, "r") as f:
saved_results = json.load(f)
has_arrays = any(
"y_true" in v for k, v in saved_results.items()
if isinstance(v, dict) and not k.startswith("_")
)
print(f"\nGenerating figures in {output_dir}/...\n", flush=True)
if has_arrays:
print("[1/13] ROC Curves", flush=True)
plot_roc_curves(saved_results, output_dir)
print("[2/13] Precision-Recall Curves", flush=True)
plot_pr_curves(saved_results, output_dir)
print("[3/13] Confusion Matrices", flush=True)
plot_confusion_matrices(saved_results, output_dir)
else:
print("[1-3/13] Skipping ROC/PR/CM — no per-sample predictions stored", flush=True)
print("[4/13] Model Comparison Bar Chart", flush=True)
plot_model_comparison(saved_results, output_dir)
print("[5/13] Feature Importance", flush=True)
plot_feature_importance(saved_results, output_dir)
print("[6/13] Feature Correlation Heatmap", flush=True)
plot_correlation_heatmap(features_csv, output_dir)
print("[7/13] Feature Distributions (AI vs Human)", flush=True)
plot_feature_distributions(features_csv, output_dir, saved_results)
print("[8/13] LaTeX / Markdown Tables", flush=True)
generate_latex_table(saved_results, output_dir)
print("[9/13] Dataset Statistics", flush=True)
plot_dataset_statistics(manifest_csv, output_dir)
print("[10/13] Per-Source Class Balance", flush=True)
plot_per_source_balance(manifest_csv, output_dir)
print("[11/13] Feature Space Embedding (PCA + t-SNE)", flush=True)
plot_feature_embedding(features_csv, output_dir)
if has_arrays:
print("[12/13] Calibration Curves", flush=True)
plot_calibration_curves(saved_results, output_dir)
print("[13/13] Prediction Score Distribution", flush=True)
plot_score_distribution(saved_results, output_dir)
else:
print("[12-13/13] Skipping calibration/score dist — no probabilities stored", flush=True)
print("\n[EXTRA] Report Summary (Markdown)", flush=True)
generate_extended_summary(saved_results, manifest_csv, features_csv, output_dir)
print(f"\nAll figures saved to {output_dir}/", flush=True)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="AURIS Visualization Pipeline")
parser.add_argument(
"--results", default="models/training_results.json",
help="Path to training_results.json",
)
parser.add_argument(
"--features", default="data/training/features.csv",
help="Path to features.csv",
)
parser.add_argument(
"--output", default="figures",
help="Output directory for figures",
)
args = parser.parse_args()
generate_all_figures(args.results, args.features, args.output)