ravel / scripts /run_stress_tests.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
23 kB
#!/usr/bin/env python3
"""Run controlled conflict/corruption stress tests from trained RAVEL checkpoints."""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import sys
from dataclasses import replace
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 PIL import Image
from sklearn.metrics import accuracy_score, f1_score
from torch.utils.data import DataLoader, Dataset
from transformers import CLIPProcessor, DebertaV2Tokenizer
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,
aurc_score,
brier_score,
evaluate,
expected_calibration_error,
freeze_non_lora,
load_dataset,
nll_score,
set_seed,
)
hf_logging.set_verbosity_error()
DATASETS = ["mvsa_multiple", "hfm_deleak"]
METHODS_E08 = ["legacy_global", "token_aux", "param_mlp", "full_revised"]
SEEDS = [1, 3, 5, 7, 11]
CONDITIONS = [
"original",
"within_class_image_shuffle",
"cross_class_image_shuffle",
"within_class_text_shuffle",
"cross_class_text_shuffle",
"blank_image",
"empty_text",
]
IMAGE_CACHE: Dict[str, torch.Tensor] = {}
CLIP_MEAN = np.array([0.48145466, 0.4578275, 0.40821073], dtype=np.float32)
CLIP_STD = np.array([0.26862954, 0.26130258, 0.27577711], dtype=np.float32)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run E08 stress tests from checkpoints.")
parser.add_argument("--output-root", default="ravel_revision_results")
parser.add_argument("--datasets", nargs="+", default=DATASETS, choices=DATASETS)
parser.add_argument("--methods", nargs="+", default=METHODS_E08, choices=METHODS_E08)
parser.add_argument("--seeds", nargs="+", type=int, default=SEEDS)
parser.add_argument("--conditions", nargs="+", default=CONDITIONS, choices=CONDITIONS)
parser.add_argument("--device", default="cuda")
parser.add_argument("--batch-size", type=int, default=32)
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)
return parser.parse_args()
def stable_seed(dataset: str, seed: int, condition: str) -> int:
digest = hashlib.sha256(f"{dataset}:{seed}:{condition}".encode("utf-8")).hexdigest()
return (int(digest[:12], 16) + int(seed)) % (2**32 - 1)
def sample_id(sample: Any, index: int) -> str:
return str(getattr(sample, "image_id", getattr(sample, "sample_id", index)))
def sample_label(dataset: str, sample: Any) -> int:
if dataset.startswith("mvsa"):
mapping = {"positive": 0, "neutral": 1, "negative": 2}
return int(mapping[str(sample.combined_majority).lower()])
return int(sample.label)
def sample_text(sample: Any) -> str:
return str(getattr(sample, "text", ""))
def image_tensor(image_path: str) -> torch.Tensor:
key = str(image_path)
cached = IMAGE_CACHE.get(key)
if cached is not None:
return cached
try:
image = Image.open(image_path).convert("RGB").resize((224, 224))
except Exception:
image = Image.new("RGB", (224, 224), (0, 0, 0))
arr = np.asarray(image, dtype=np.float32) / 255.0
arr = (arr - CLIP_MEAN) / CLIP_STD
tensor = torch.from_numpy(np.transpose(arr, (2, 0, 1))).float()
IMAGE_CACHE[key] = tensor
return tensor
class StressPairDataset(Dataset):
def __init__(self, dataset: str, samples: Sequence[Any]):
self.dataset = dataset
self.samples = list(samples)
def __len__(self) -> int:
return len(self.samples)
def __getitem__(self, idx: int) -> Dict[str, Any]:
sample = self.samples[idx]
return {
"pixel_values": image_tensor(sample.image_path),
"text": sample_text(sample),
"labels": sample_label(self.dataset, sample),
}
def set_mvsa_text(sample: Any, text: str) -> Any:
new_sample = replace(sample)
setattr(new_sample, "_text_cache", text)
return new_sample
def clone_with_image(dataset: str, sample: Any, image_path: str) -> Any:
return replace(sample, image_path=image_path)
def clone_with_text(dataset: str, sample: Any, text: str) -> Any:
if dataset.startswith("mvsa"):
return set_mvsa_text(sample, text)
return replace(sample, text=text)
def choose_donors(dataset: str, samples: Sequence[Any], seed: int, condition: str) -> List[int]:
rng = np.random.default_rng(stable_seed(dataset, seed, condition))
labels = np.array([sample_label(dataset, sample) for sample in samples])
donor_indices: List[int] = []
for idx, label in enumerate(labels):
if condition.startswith("within_class"):
candidates = np.flatnonzero(labels == label)
else:
candidates = np.flatnonzero(labels != label)
candidates = candidates[candidates != idx]
if candidates.size == 0:
donor_indices.append(idx)
else:
donor_indices.append(int(rng.choice(candidates)))
return donor_indices
def make_stress_samples(
dataset: str,
samples: Sequence[Any],
seed: int,
condition: str,
) -> Tuple[List[Any], List[str]]:
if condition == "original":
return list(samples), ["" for _ in samples]
if condition == "blank_image":
missing_path = str(PROJECT_ROOT / "ravel_revision_results" / "_stress_blank_missing_image.jpg")
return [clone_with_image(dataset, sample, missing_path) for sample in samples], ["" for _ in samples]
if condition == "empty_text":
return [clone_with_text(dataset, sample, "") for sample in samples], ["" for _ in samples]
donor_indices = choose_donors(dataset, samples, seed, condition)
out: List[Any] = []
donor_ids: List[str] = []
for idx, donor_idx in enumerate(donor_indices):
sample = samples[idx]
donor = samples[donor_idx]
donor_ids.append(sample_id(donor, donor_idx))
if condition.endswith("image_shuffle"):
out.append(clone_with_image(dataset, sample, donor.image_path))
elif condition.endswith("text_shuffle"):
out.append(clone_with_text(dataset, sample, sample_text(donor)))
else:
raise ValueError(condition)
return out, donor_ids
def build_loader(
dataset: str,
samples: List[Any],
cfg: Dict[str, Any],
batch_size: int,
max_length: int,
num_workers: int,
) -> DataLoader:
tokenizer = DebertaV2Tokenizer.from_pretrained(cfg["text_model_id"])
def collate_fn(batch: List[Dict[str, Any]]) -> Dict[str, torch.Tensor]:
texts = [item["text"] for item in batch]
labels = torch.tensor([item["labels"] for item in batch], dtype=torch.long)
text_inputs = tokenizer(
texts,
max_length=max_length,
padding=True,
truncation=True,
return_tensors="pt",
)
return {
"pixel_values": torch.stack([item["pixel_values"] for item in batch], dim=0),
"input_ids": text_inputs["input_ids"],
"attention_mask": text_inputs["attention_mask"],
"labels": labels,
}
return DataLoader(
StressPairDataset(dataset, samples),
batch_size=batch_size,
shuffle=False,
num_workers=num_workers,
pin_memory=True,
collate_fn=collate_fn,
)
def prob_columns(df: pd.DataFrame) -> List[str]:
cols = [col for col in df.columns if col.startswith("prob_class_")]
if not cols:
cols = [col for col in df.columns if col.startswith("prob_") and col[5:].isdigit()]
return sorted(cols, key=lambda name: int(name.rsplit("_", 1)[-1]))
def class_vector(row: pd.Series, prefix: str, num_classes: int) -> List[Optional[float]]:
out: List[Optional[float]] = []
for c in range(num_classes):
for name in (f"{prefix}{c}", f"{prefix}_class_{c}"):
if name in row and pd.notna(row[name]):
out.append(float(row[name]))
break
else:
out.append(None)
return out
def to_optional_float(value: Any) -> float:
try:
if value == "":
return float("nan")
return float(value)
except Exception:
return float("nan")
def rows_to_stress_frame(
rows: List[Dict[str, Any]],
dataset: str,
method: str,
seed: int,
condition: str,
donor_ids: Sequence[str],
num_classes: int,
) -> pd.DataFrame:
out_rows: List[Dict[str, Any]] = []
for idx, row in enumerate(rows):
probs = [float(row.get(f"prob_class_{c}", np.nan)) for c in range(num_classes)]
visual_probs = [to_optional_float(row.get(f"visual_prob_class_{c}", np.nan)) for c in range(num_classes)]
text_probs = [to_optional_float(row.get(f"text_prob_class_{c}", np.nan)) for c in range(num_classes)]
true_label = int(row["true_label"])
pred = int(row["predicted_label"])
confidence = float(np.nanmax(probs))
nll = -math.log(max(float(probs[true_label]), 1e-12))
if all(np.isfinite(float(x)) for x in visual_probs + text_probs):
tv = 0.5 * float(np.sum(np.abs(np.array(visual_probs, dtype=float) - np.array(text_probs, dtype=float))))
else:
tv = float("nan")
payload: Dict[str, Any] = {
"sample_id": row["sample_id"],
"donor_sample_id": donor_ids[idx] if idx < len(donor_ids) else "",
"dataset": dataset,
"condition": condition,
"method": method,
"seed": seed,
"true_label": true_label,
"predicted_label": pred,
"confidence": confidence,
"entropy": float(row.get("prediction_entropy", np.nan)),
"raw_ece_component": abs(float(pred == true_label) - confidence),
"nll": nll,
"visual_probs": json.dumps([None if pd.isna(x) else float(x) for x in visual_probs]),
"text_probs": json.dumps([None if pd.isna(x) else float(x) for x in text_probs]),
"tv_disagreement": tv,
}
for c in range(num_classes):
payload[f"logit_{c}"] = row.get(f"logit_class_{c}", np.nan)
payload[f"prob_{c}"] = probs[c]
payload[f"visual_prob_{c}"] = visual_probs[c]
payload[f"text_prob_{c}"] = text_probs[c]
out_rows.append(payload)
return pd.DataFrame(out_rows)
def original_to_stress_frame(path: Path, dataset: str, method: str, seed: int) -> pd.DataFrame:
df = pd.read_csv(path)
pcols = prob_columns(df)
num_classes = len(pcols)
rows: List[Dict[str, Any]] = []
for _, row in df.iterrows():
probs = [float(row[col]) for col in pcols]
true_label = int(row["true_label"])
pred = int(row["predicted_label"])
confidence = float(np.max(probs))
visual_probs = class_vector(row, "visual_prob", num_classes)
text_probs = class_vector(row, "text_prob", num_classes)
if all(x is not None and not pd.isna(x) for x in visual_probs + text_probs):
tv = 0.5 * float(np.sum(np.abs(np.array(visual_probs, dtype=float) - np.array(text_probs, dtype=float))))
else:
tv = float("nan")
payload: Dict[str, Any] = {
"sample_id": row["sample_id"],
"donor_sample_id": "",
"dataset": dataset,
"condition": "original",
"method": method,
"seed": seed,
"true_label": true_label,
"predicted_label": pred,
"confidence": confidence,
"entropy": float(row.get("prediction_entropy", -(np.array(probs) * np.log(np.clip(probs, 1e-12, 1))).sum())),
"raw_ece_component": abs(float(pred == true_label) - confidence),
"nll": -math.log(max(float(probs[true_label]), 1e-12)),
"visual_probs": json.dumps([None if x is None or pd.isna(x) else float(x) for x in visual_probs]),
"text_probs": json.dumps([None if x is None or pd.isna(x) else float(x) for x in text_probs]),
"tv_disagreement": tv,
}
for c in range(num_classes):
payload[f"logit_{c}"] = row.get(f"logit_class_{c}", np.nan)
payload[f"prob_{c}"] = probs[c]
payload[f"visual_prob_{c}"] = visual_probs[c]
payload[f"text_prob_{c}"] = text_probs[c]
rows.append(payload)
return pd.DataFrame(rows)
def load_model_and_data(
args: argparse.Namespace,
dataset: str,
method_key: str,
seed: int,
) -> Tuple[nn.Module, Dict[str, Any], List[Any], int, int, torch.device]:
root = Path(args.output_root)
run_dir = root / "runs" / dataset / method_key / f"seed_{seed}"
ckpt_path = run_dir / "checkpoint.pt"
if not ckpt_path.exists():
raise FileNotFoundError(ckpt_path)
ckpt = torch.load(ckpt_path, map_location="cpu")
ckpt_cfg = ckpt.get("cfg", {}) if isinstance(ckpt, dict) else {}
max_length = int(args.max_length or ckpt_cfg.get("max_length") or 96)
method = METHODS[method_key]
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=args.batch_size,
max_length=max_length,
num_workers=0,
method=method,
hfm_deleak_manifest=args.hfm_deleak_manifest,
limits=(None, None, None),
)
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": args.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)
model.load_state_dict(ckpt.get("model_state", {}), strict=False)
model.eval()
return model, cfg, test_samples, num_classes, max_length, device
def run_condition(
args: argparse.Namespace,
model: nn.Module,
cfg: Dict[str, Any],
test_samples: List[Any],
num_classes: int,
max_length: int,
device: torch.device,
dataset: str,
method: str,
seed: int,
condition: str,
) -> pd.DataFrame:
stress_samples, donor_ids = make_stress_samples(dataset, test_samples, seed, condition)
loader = build_loader(dataset, stress_samples, cfg, args.batch_size, max_length, args.num_workers)
criterion = nn.CrossEntropyLoss()
metrics, rows, _logits, _labels = evaluate(
model=model,
loader=loader,
samples=test_samples,
method=METHODS[method],
device=device,
criterion=criterion,
num_classes=num_classes,
dataset_key=dataset,
seed=seed,
)
df = rows_to_stress_frame(rows, dataset, method, seed, condition, donor_ids, num_classes)
df.attrs["metrics"] = metrics
return df
def metrics_from_frame(df: pd.DataFrame) -> Dict[str, float]:
pcols = prob_columns(df)
probs = df[pcols].astype(float).to_numpy()
y = df["true_label"].astype(int).to_numpy()
pred = df["predicted_label"].astype(int).to_numpy()
num_classes = probs.shape[1]
aurc, *_ = aurc_score(probs, y)
return {
"accuracy": float(accuracy_score(y, pred)),
"macro_f1": float(f1_score(y, pred, average="macro", zero_division=0)),
"weighted_f1": float(f1_score(y, pred, average="weighted", zero_division=0)),
"raw_ece": expected_calibration_error(probs, y),
"nll": nll_score(probs, y),
"brier": brier_score(probs, y, num_classes),
"aurc": float(aurc),
"mean_confidence": float(df["confidence"].astype(float).mean()),
"mean_entropy": float(df["entropy"].astype(float).mean()),
"mean_tv_disagreement": float(df["tv_disagreement"].astype(float).mean()) if df["tv_disagreement"].notna().any() else float("nan"),
}
def write_original_if_needed(args: argparse.Namespace, dataset: str, method: str, seed: int) -> None:
root = Path(args.output_root)
out_path = root / "stress_predictions" / dataset / "original" / f"{method}_seed_{seed}.csv"
if out_path.exists() and not args.overwrite:
return
src = root / "predictions" / dataset / f"{method}_seed_{seed}.csv"
if not src.exists():
raise FileNotFoundError(src)
out_path.parent.mkdir(parents=True, exist_ok=True)
original_to_stress_frame(src, dataset, method, seed).to_csv(out_path, index=False)
def aggregate_outputs(output_root: Path, datasets: Sequence[str], methods: Sequence[str], seeds: Sequence[int], conditions: Sequence[str]) -> None:
rows: List[Dict[str, Any]] = []
for dataset in datasets:
for condition in conditions:
for method in methods:
for seed in seeds:
path = output_root / "stress_predictions" / dataset / condition / f"{method}_seed_{seed}.csv"
if not path.exists():
continue
df = pd.read_csv(path)
if df.empty:
continue
metrics = metrics_from_frame(df)
rows.append(
{
"dataset": dataset,
"condition": condition,
"method": method,
"seed": seed,
"accuracy": metrics["accuracy"],
"macro_f1": metrics["macro_f1"],
"weighted_f1": metrics["weighted_f1"],
"f1": metrics["weighted_f1"] if dataset.startswith("mvsa") else metrics["macro_f1"],
"raw_ece": metrics["raw_ece"],
"nll": metrics["nll"],
"brier": metrics["brier"],
"aurc": metrics["aurc"],
"mean_confidence": metrics["mean_confidence"],
"mean_entropy": metrics["mean_entropy"],
"prediction_entropy": metrics["mean_entropy"],
"mean_tv_disagreement": metrics["mean_tv_disagreement"],
"status": "COMPLETE",
}
)
agg_dir = output_root / "aggregate_results"
agg_dir.mkdir(parents=True, exist_ok=True)
result = pd.DataFrame(rows)
result.to_csv(agg_dir / "stress_test_results.csv", index=False)
if result.empty:
pd.DataFrame().to_csv(agg_dir / "stress_test_summary.csv", index=False)
return
summary_rows: List[Dict[str, Any]] = []
for (dataset, condition, method), group in result.groupby(["dataset", "condition", "method"], dropna=False):
out: Dict[str, Any] = {
"dataset": dataset,
"condition": condition,
"method": method,
"num_seeds": int(group["seed"].nunique()),
"status": "COMPLETE" if int(group["seed"].nunique()) >= 5 else "PARTIAL",
}
for metric in ["accuracy", "macro_f1", "weighted_f1", "f1", "raw_ece", "nll", "brier", "aurc", "mean_confidence", "mean_entropy", "mean_tv_disagreement"]:
vals = group[metric].astype(float)
out[f"{metric}_mean"] = float(vals.mean())
out[f"{metric}_std"] = float(vals.std(ddof=1)) if len(vals) > 1 else 0.0
out["prediction_entropy_mean"] = out["mean_entropy_mean"]
out["prediction_entropy_std"] = out["mean_entropy_std"]
summary_rows.append(out)
pd.DataFrame(summary_rows).to_csv(agg_dir / "stress_test_summary.csv", index=False)
def main() -> None:
args = parse_args()
root = Path(args.output_root)
planned = [(d, m, s) for d in args.datasets for m in args.methods for s in args.seeds]
if args.max_runs is not None:
planned = planned[: args.max_runs]
print(f"Planned E08 model runs: {len(planned)}", flush=True)
for dataset, method, seed in planned:
for condition in args.conditions:
if condition == "original":
write_original_if_needed(args, dataset, method, seed)
pending = [
condition
for condition in args.conditions
if condition != "original"
and (
args.overwrite
or not (root / "stress_predictions" / dataset / condition / f"{method}_seed_{seed}.csv").exists()
)
]
if not pending:
print(f"SKIP E08 {dataset} {method} seed={seed}", flush=True)
continue
print(f"RUN E08 dataset={dataset} method={method} seed={seed} conditions={','.join(pending)}", flush=True)
model, cfg, test_samples, num_classes, max_length, device = load_model_and_data(args, dataset, method, seed)
for condition in pending:
out_path = root / "stress_predictions" / dataset / condition / f"{method}_seed_{seed}.csv"
df = run_condition(args, model, cfg, test_samples, num_classes, max_length, device, dataset, method, seed, condition)
out_path.parent.mkdir(parents=True, exist_ok=True)
df.to_csv(out_path, index=False)
metrics = metrics_from_frame(df)
print(
f"DONE E08 {dataset} {condition} {method} seed={seed} "
f"F1={metrics['weighted_f1' if dataset.startswith('mvsa') else 'macro_f1']:.4f} "
f"ECE={metrics['raw_ece']:.4f} conf={metrics['mean_confidence']:.4f}",
flush=True,
)
del model
if torch.cuda.is_available():
torch.cuda.empty_cache()
aggregate_outputs(root, args.datasets, args.methods, args.seeds, args.conditions)
rows = pd.read_csv(root / "aggregate_results" / "stress_test_results.csv")
print(f"E08 stress rows: {len(rows)}", flush=True)
if __name__ == "__main__":
main()