crowncode-backend / app /training /ablation_flatness.py
Rthur2003's picture
feat: implement automated training diagnostics, dataset bias analysis, and revision figure generation pipelines
906c392
Raw History Blame Contribute Delete
5.11 kB
"""
Spectral flatness ablation study (reviewer priority #3).
The paper reports spectral_flatness_std as the single most important
LightGBM feature and attributes AI/human separability to it. The reviewer
asks for a direct ablation: does performance meaningfully depend on this one
feature, or would the model do just as well without it (in which case the
"AI music has flatter spectra" narrative should be stated more cautiously)?
Three configurations, each evaluated with 5-fold stratified CV + LightGBM:
(a) all 47 features
(b) 46 features, spectral_flatness_std removed
(c) spectral_flatness_std alone (1 feature)
Usage:
python -m app.training.ablation_flatness
"""
from __future__ import annotations
import csv
import sys
import warnings
from pathlib import Path
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
import lightgbm as lgb
from sklearn.exceptions import ConvergenceWarning
from sklearn.metrics import balanced_accuracy_score, f1_score, roc_auc_score, roc_curve
from sklearn.model_selection import StratifiedKFold, cross_val_predict
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
from app.training.evaluate import load_features_csv
FEATURES_CSV = Path("D:/CrownCode/DataSet/features.csv")
OUTPUT_CSV = Path(__file__).resolve().parents[3] / "docs/academic/paper/real_tables/flatness_ablation.csv"
_EXCLUDED_COLUMNS = {"file_path", "label_int", "duration_sec", "sample_rate"}
TARGET_FEATURE = "spectral_flatness_std"
_LGBM_PARAMS = dict(
n_estimators=300, max_depth=-1, learning_rate=0.05, num_leaves=31,
subsample=0.8, colsample_bytree=0.8, min_child_samples=20,
reg_alpha=0.1, reg_lambda=1.0, class_weight="balanced",
random_state=42, verbose=-1,
)
def _feature_columns() -> list[str]:
with open(FEATURES_CSV, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
return [c for c in (reader.fieldnames or []) if c not in _EXCLUDED_COLUMNS]
def _optimal_threshold(y_true: np.ndarray, y_prob: np.ndarray) -> float:
fpr, tpr, thresholds = roc_curve(y_true, y_prob)
return float(thresholds[np.argmax(tpr - fpr)])
def _evaluate_feature_set(X: np.ndarray, y: np.ndarray, n_folds: int = 5) -> dict:
cv = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=42)
pipeline = Pipeline([
("scaler", StandardScaler()),
("model", lgb.LGBMClassifier(**_LGBM_PARAMS)),
])
with warnings.catch_warnings():
warnings.simplefilter("ignore", category=ConvergenceWarning)
y_prob = cross_val_predict(pipeline, X, y, cv=cv, method="predict_proba")[:, 1]
threshold = _optimal_threshold(y, y_prob)
y_pred = (y_prob >= threshold).astype(int)
return {
"roc_auc": round(float(roc_auc_score(y, y_prob)), 4),
"f1": round(float(f1_score(y, y_pred, zero_division=0)), 4),
"balanced_accuracy": round(float(balanced_accuracy_score(y, y_pred)), 4),
"threshold": round(threshold, 4),
}
def run() -> None:
X_full, y = load_features_csv(FEATURES_CSV)
X_full = np.nan_to_num(X_full, nan=0.0, posinf=1.0, neginf=-1.0)
feature_cols = _feature_columns()
if TARGET_FEATURE not in feature_cols:
raise RuntimeError(f"{TARGET_FEATURE} not found in {FEATURES_CSV} columns: {feature_cols}")
target_idx = feature_cols.index(TARGET_FEATURE)
keep_idx = [i for i in range(len(feature_cols)) if i != target_idx]
configs = {
"47_features_all": X_full,
"46_features_without_flatness_std": X_full[:, keep_idx],
"1_feature_flatness_std_only": X_full[:, [target_idx]],
}
results = []
for name, X in configs.items():
print(f"\nEvaluating: {name} ({X.shape[1]} features)")
metrics = _evaluate_feature_set(X, y)
metrics["config"] = name
metrics["n_features"] = X.shape[1]
results.append(metrics)
print(f" AUC={metrics['roc_auc']:.4f} F1={metrics['f1']:.4f} "
f"BalAcc={metrics['balanced_accuracy']:.4f}")
OUTPUT_CSV.parent.mkdir(parents=True, exist_ok=True)
fieldnames = ["config", "n_features", "roc_auc", "f1", "balanced_accuracy", "threshold"]
with open(OUTPUT_CSV, "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
for r in results:
writer.writerow({k: r[k] for k in fieldnames})
print(f"\nOutput: {OUTPUT_CSV}")
full_auc = next(r["roc_auc"] for r in results if r["config"] == "47_features_all")
no_flatness_auc = next(r["roc_auc"] for r in results if r["config"] == "46_features_without_flatness_std")
flatness_only_auc = next(r["roc_auc"] for r in results if r["config"] == "1_feature_flatness_std_only")
print("\n" + "=" * 60)
print(f" 47 features: AUC={full_auc:.4f}")
print(f" 46 features (no flatness): AUC={no_flatness_auc:.4f} (diff={full_auc - no_flatness_auc:+.4f})")
print(f" flatness_std alone: AUC={flatness_only_auc:.4f}")
print("=" * 60)
if __name__ == "__main__":
run()