ravel / scripts /generate_validation_predictions.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
11.7 kB
#!/usr/bin/env python3
"""Generate validation logits and temperature-scaling artifacts for RAVEL runs."""
from __future__ import annotations
import argparse
import json
import math
import os
import sys
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from transformers.utils import logging as hf_logging
PROJECT_ROOT = Path(__file__).resolve().parents[1]
SCRIPT_DIR = Path(__file__).resolve().parent
for path in [PROJECT_ROOT, SCRIPT_DIR]:
if str(path) not in sys.path:
sys.path.insert(0, str(path))
from run_revised_experiments import ( # noqa: E402
METHODS,
apply_method_trainability,
evaluate,
freeze_non_lora,
load_dataset,
set_seed,
)
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
hf_logging.set_verbosity_error()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Generate validation predictions and temperature scaling.")
parser.add_argument("--output-root", default="ravel_revision_results")
parser.add_argument("--datasets", nargs="+", default=None)
parser.add_argument("--methods", nargs="+", default=None, choices=sorted(METHODS))
parser.add_argument("--seeds", nargs="+", type=int, default=None)
parser.add_argument("--device", default="cuda")
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--max-length", type=int, default=None)
parser.add_argument("--num-workers", type=int, default=0)
parser.add_argument("--hfm-deleak-manifest", default="ravel_revision_results/data_audit/hfm_split_manifest_deleaked.csv")
parser.add_argument("--overwrite", action="store_true")
parser.add_argument("--max-runs", type=int, default=None)
parser.add_argument("--limit-train-samples", type=int, default=None)
parser.add_argument("--limit-val-samples", type=int, default=None)
parser.add_argument("--limit-test-samples", type=int, default=None)
return parser.parse_args()
def softmax_np(logits: np.ndarray) -> np.ndarray:
logits = logits.astype(float)
logits = logits - np.max(logits, axis=1, keepdims=True)
exp = np.exp(logits)
return exp / np.clip(exp.sum(axis=1, keepdims=True), 1e-12, None)
def ece_equal_width(probs: np.ndarray, y_true: np.ndarray, bins: int = 15) -> float:
conf = probs.max(axis=1)
pred = probs.argmax(axis=1)
correct = (pred == y_true).astype(float)
edges = np.linspace(0.0, 1.0, bins + 1)
ece = 0.0
for idx, (low, high) in enumerate(zip(edges[:-1], edges[1:])):
mask = (conf > low) & (conf <= high) if idx > 0 else (conf >= low) & (conf <= high)
if not mask.any():
continue
ece += float(mask.mean()) * abs(float(correct[mask].mean()) - float(conf[mask].mean()))
return float(ece)
def nll_score(probs: np.ndarray, y_true: np.ndarray) -> float:
return float(-np.mean(np.log(np.clip(probs[np.arange(len(y_true)), y_true], 1e-12, 1.0))))
def fit_temperature(logits: np.ndarray, labels: np.ndarray) -> Tuple[float, float, float]:
raw_probs = softmax_np(logits)
before = nll_score(raw_probs, labels)
try:
from scipy.optimize import minimize_scalar
def objective(log_temp: float) -> float:
temp = float(math.exp(log_temp))
return nll_score(softmax_np(logits / temp), labels)
result = minimize_scalar(objective, bounds=(math.log(0.05), math.log(20.0)), method="bounded")
temperature = float(math.exp(float(result.x)))
except Exception:
grid = np.exp(np.linspace(math.log(0.05), math.log(20.0), 300))
losses = np.array([nll_score(softmax_np(logits / temp), labels) for temp in grid])
temperature = float(grid[int(losses.argmin())])
after = nll_score(softmax_np(logits / max(temperature, 1e-8)), labels)
return temperature, before, after
def logits_from_prediction_df(df: pd.DataFrame) -> Tuple[np.ndarray, np.ndarray]:
logit_cols = sorted(
[col for col in df.columns if col.startswith("logit_class_")],
key=lambda col: int(col.rsplit("_", 1)[1]),
)
if not logit_cols:
logit_cols = sorted(
[col for col in df.columns if col.startswith("logit_") and col.removeprefix("logit_").isdigit()],
key=lambda col: int(col.rsplit("_", 1)[1]),
)
logits = df[logit_cols].astype(float).to_numpy()
labels = df["true_label"].astype(int).to_numpy()
return logits, labels
def prediction_rows_to_frame(rows: List[Dict[str, Any]], num_classes: int) -> pd.DataFrame:
df = pd.DataFrame(rows)
for c in range(num_classes):
df[f"logit_{c}"] = pd.to_numeric(df[f"logit_class_{c}"], errors="coerce")
df[f"prob_{c}"] = pd.to_numeric(df[f"prob_class_{c}"], errors="coerce")
df["confidence"] = df[[f"prob_{c}" for c in range(num_classes)]].max(axis=1)
minimal = [
"sample_id",
"dataset",
"split",
"seed",
"method",
"true_label",
"predicted_label",
*[f"logit_{c}" for c in range(num_classes)],
*[f"prob_{c}" for c in range(num_classes)],
"confidence",
]
original_cols = [col for col in df.columns if col not in minimal]
return df[minimal + original_cols]
def discover_runs(args: argparse.Namespace) -> List[Tuple[str, str, int, Path]]:
root = Path(args.output_root)
runs: List[Tuple[str, str, int, Path]] = []
for metrics_path in sorted((root / "runs").glob("*/*/seed_*/metrics.json")):
dataset, method, seed_dir = metrics_path.parts[-4], metrics_path.parts[-3], metrics_path.parts[-2]
seed = int(seed_dir.replace("seed_", ""))
if args.datasets and dataset not in args.datasets:
continue
if args.methods and method not in args.methods:
continue
if args.seeds and seed not in args.seeds:
continue
if method not in METHODS:
continue
runs.append((dataset, method, seed, metrics_path.parent))
if args.max_runs is not None:
runs = runs[: max(0, int(args.max_runs))]
return runs
def load_checkpoint(path: Path) -> Dict[str, Any]:
return torch.load(path, map_location="cpu")
def process_run(args: argparse.Namespace, dataset: str, method_key: str, seed: int, run_dir: Path) -> bool:
validation_path = run_dir / "validation_predictions.csv"
temp_path = run_dir / "temperature_scaling.json"
if validation_path.exists() and temp_path.exists() and not args.overwrite:
print(f"SKIP validation+temperature {dataset} {method_key} seed={seed}", flush=True)
return True
ckpt_path = run_dir / "checkpoint.pt"
test_pred_path = Path(args.output_root) / "predictions" / dataset / f"{method_key}_seed_{seed}.csv"
if not ckpt_path.exists() or not test_pred_path.exists():
print(f"MISS checkpoint/test predictions {dataset} {method_key} seed={seed}", flush=True)
return False
method = METHODS[method_key]
ckpt = load_checkpoint(ckpt_path)
ckpt_cfg = ckpt.get("cfg", {}) if isinstance(ckpt, dict) else {}
max_length = int(args.max_length or ckpt_cfg.get("max_length") or 96)
batch_size = int(args.batch_size or ckpt_cfg.get("batch_size") or 8)
set_seed(seed)
device = torch.device(args.device if torch.cuda.is_available() or args.device == "cpu" else "cpu")
(
model_cls,
cfg,
_train_loader,
val_loader,
_test_loader,
_train_samples,
val_samples,
_test_samples,
num_classes,
_label_names,
) = load_dataset(
dataset_key=dataset,
seed=seed,
batch_size=batch_size,
max_length=max_length,
num_workers=args.num_workers,
method=method,
hfm_deleak_manifest=args.hfm_deleak_manifest,
limits=(args.limit_train_samples, args.limit_val_samples, args.limit_test_samples),
)
cfg.update(ckpt_cfg)
cfg.update(
{
"architecture": method.architecture,
"enable_clip_lora": method.enable_lora,
"enable_text_lora": method.enable_lora,
"seed": seed,
"batch_size": batch_size,
"max_length": max_length,
}
)
model = model_cls(cfg).to(device)
if hasattr(model, "vision_lora"):
freeze_non_lora(model.vision_lora)
if hasattr(model, "text"):
freeze_non_lora(model.text)
apply_method_trainability(model, method)
load_result = model.load_state_dict(ckpt.get("model_state", {}), strict=False)
if load_result.unexpected_keys:
print(f" unexpected keys: {len(load_result.unexpected_keys)}", flush=True)
criterion = nn.CrossEntropyLoss()
metrics, rows, val_logits, val_labels = evaluate(
model=model,
loader=val_loader,
samples=val_samples,
method=method,
device=device,
criterion=criterion,
num_classes=num_classes,
dataset_key=dataset,
seed=seed,
)
val_df = prediction_rows_to_frame(rows, num_classes)
val_df.to_csv(validation_path, index=False)
temperature, val_before, val_after = fit_temperature(val_logits, val_labels)
test_df = pd.read_csv(test_pred_path)
test_logits, test_labels = logits_from_prediction_df(test_df)
raw_probs = softmax_np(test_logits)
ts_probs = softmax_np(test_logits / max(temperature, 1e-8))
payload = {
"dataset": dataset,
"method": method_key,
"seed": seed,
"fit_split": "validation",
"temperature": temperature,
"validation_nll_before": val_before,
"validation_nll_after": val_after,
"validation_accuracy": metrics.get("accuracy"),
"validation_macro_f1": metrics.get("macro_f1"),
"validation_weighted_f1": metrics.get("weighted_f1"),
"test_raw_ece": ece_equal_width(raw_probs, test_labels),
"test_ts_ece": ece_equal_width(ts_probs, test_labels),
"test_raw_nll": nll_score(raw_probs, test_labels),
"test_ts_nll": nll_score(ts_probs, test_labels),
}
temp_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
cal_path = run_dir / "calibration.json"
if cal_path.exists():
try:
cal = json.loads(cal_path.read_text(encoding="utf-8"))
except Exception:
cal = {}
else:
cal = {}
cal.update(
{
"raw_ece": payload["test_raw_ece"],
"temperature_scaled_ece": payload["test_ts_ece"],
"negative_log_likelihood": payload["test_raw_nll"],
"negative_log_likelihood_temperature_scaled": payload["test_ts_nll"],
"temperature": payload["temperature"],
"note": "Temperature fit on validation_predictions.csv.",
}
)
cal_path.write_text(json.dumps(cal, indent=2), encoding="utf-8")
if torch.cuda.is_available():
torch.cuda.empty_cache()
print(
f"DONE {dataset} {method_key} seed={seed} "
f"T={temperature:.4f} rawECE={payload['test_raw_ece']:.4f} tsECE={payload['test_ts_ece']:.4f}",
flush=True,
)
return True
def main() -> None:
args = parse_args()
runs = discover_runs(args)
print(f"Planned validation runs: {len(runs)}", flush=True)
completed = 0
for dataset, method, seed, run_dir in runs:
ok = process_run(args, dataset, method, seed, run_dir)
completed += int(ok)
print(f"Validation temperature artifacts complete: {completed}/{len(runs)}", flush=True)
if __name__ == "__main__":
main()