#!/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()