ravel / scripts /run_correctness_pilot.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
12.8 kB
#!/usr/bin/env python3
"""Small correctness pilot for revised RAVEL.
This is intentionally not a full experiment. It runs a few training/evaluation
batches on real local data to verify:
- data loader works;
- forward/backward/optimizer step works;
- loss decreases or at least remains finite;
- token-level shapes and disagreement tensors are produced.
"""
from __future__ import annotations
import argparse
import csv
import json
import sys
import time
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any, Callable, Dict, Iterable, List, Tuple
import torch
import torch.nn as nn
from torch.optim import AdamW
from transformers import CLIPProcessor, DebertaV2Tokenizer
PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
from src.revised_ravel_model import token_loss
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run a tiny correctness pilot.")
parser.add_argument("--dataset", choices=["mvsa_multiple", "hfm"], default="mvsa_multiple")
parser.add_argument("--architecture", choices=["token", "legacy_global"], default="token")
parser.add_argument("--device", default="cuda")
parser.add_argument("--batch-size", type=int, default=2)
parser.add_argument("--max-length", type=int, default=64)
parser.add_argument("--per-class-train", type=int, default=8)
parser.add_argument("--per-class-val", type=int, default=4)
parser.add_argument("--train-batches", type=int, default=3)
parser.add_argument("--val-batches", type=int, default=2)
parser.add_argument("--learning-rate", type=float, default=1e-4)
parser.add_argument("--seed", type=int, default=7)
parser.add_argument("--output-dir", default="ravel_revision_results/pilot")
parser.add_argument(
"--hfm-manifest",
default="ravel_revision_results/data_audit/hfm_split_manifest_deleaked.csv",
)
return parser.parse_args()
def write_csv(path: Path, rows: Iterable[Dict[str, Any]], fieldnames: List[str]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames, extrasaction="ignore")
writer.writeheader()
for row in rows:
writer.writerow(row)
def shape_summary(value: Any) -> Any:
if hasattr(value, "shape"):
return list(value.shape)
if isinstance(value, dict):
return {key: shape_summary(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [shape_summary(item) for item in value]
return type(value).__name__
def balanced_subset(
samples: List[Any],
per_class: int,
label_fn: Callable[[Any], Any],
) -> List[Any]:
buckets: Dict[Any, List[Any]] = defaultdict(list)
for sample in samples:
buckets[label_fn(sample)].append(sample)
subset: List[Any] = []
for label in sorted(buckets):
subset.extend(buckets[label][:per_class])
return subset
def load_mvsa_multiple(args: argparse.Namespace) -> Tuple[Any, Any, Any, Dict[str, Any], int]:
from src.mvsa_multiple_pipeline import (
CLARAModel,
DEFAULT_MVSA_MULTIPLE_CONFIG,
MVSALoader,
create_dataloaders,
)
cfg = dict(DEFAULT_MVSA_MULTIPLE_CONFIG)
cfg.update(
{
"architecture": args.architecture,
"batch_size": args.batch_size,
"max_length": args.max_length,
"num_workers": 0,
"pin_memory": False,
"persistent_workers": False,
"prefetch_factor": 2,
"learning_rate": args.learning_rate,
"seed": args.seed,
"use_mixup_negative": False,
"use_weighted_sampler": False,
"text_unfreeze_mode": "freeze_all",
"unfreeze_epoch": 0,
}
)
loader = MVSALoader(cfg["text_dir"], cfg["label_file"])
loader.load(
preprocessing_mode=str(cfg.get("preprocessing_mode", "paper")),
require_unanimous=bool(cfg["require_unanimous"]),
require_cross_agree=bool(cfg["require_cross_agree"]),
paper_exact_counts=bool(cfg.get("paper_exact_counts", True)),
)
train_samples, val_samples, test_samples = loader.split(
train_ratio=float(cfg["train_ratio"]),
val_ratio=float(cfg["val_ratio"]),
seed=int(cfg["seed"]),
paper_811=True,
)
train_subset = balanced_subset(train_samples, args.per_class_train, lambda sample: sample.combined_majority)
val_subset = balanced_subset(val_samples, args.per_class_val, lambda sample: sample.combined_majority)
clip_processor = CLIPProcessor.from_pretrained(cfg["vision_model_id"])
tokenizer = DebertaV2Tokenizer.from_pretrained(cfg["text_model_id"])
train_loader, val_loader, _ = create_dataloaders(
train_subset,
val_subset,
test_samples[: max(1, args.batch_size)],
clip_processor=clip_processor,
tokenizer=tokenizer,
batch_size=args.batch_size,
max_length=args.max_length,
num_workers=0,
pin_memory=False,
persistent_workers=False,
prefetch_factor=2,
use_mixup_negative=False,
mixup_alpha=0.0,
negative_class_boost=1.0,
min_ratio_negative=0.0,
weighted_train_sampler=False,
)
return CLARAModel, train_loader, val_loader, cfg, int(cfg["num_classes"])
def load_hfm(args: argparse.Namespace) -> Tuple[Any, Any, Any, Dict[str, Any], int]:
from src.hfm_pipeline import (
CLARAModel,
DEFAULT_HFM_CONFIG,
HFMLoader,
create_dataloaders,
)
cfg = dict(DEFAULT_HFM_CONFIG)
cfg.update(
{
"architecture": args.architecture,
"batch_size": args.batch_size,
"max_length": args.max_length,
"num_workers": 0,
"pin_memory": False,
"learning_rate": args.learning_rate,
"seed": args.seed,
"split_manifest": args.hfm_manifest,
"text_unfreeze_mode": "freeze_all",
}
)
loader = HFMLoader(cfg["text_dir"], cfg["image_root"])
loader.load_from_manifest(str(cfg["split_manifest"]))
train_samples = loader.get_split("train")
val_samples = loader.get_split("val")
test_samples = loader.get_split("test")
train_subset = balanced_subset(train_samples, args.per_class_train, lambda sample: sample.label)
val_subset = balanced_subset(val_samples, args.per_class_val, lambda sample: sample.label)
clip_processor = CLIPProcessor.from_pretrained(cfg["vision_model_id"])
tokenizer = DebertaV2Tokenizer.from_pretrained(cfg["text_model_id"])
train_loader, val_loader, _ = create_dataloaders(
train_samples=train_subset,
val_samples=val_subset,
test_samples=test_samples[: max(1, args.batch_size)],
clip_processor=clip_processor,
tokenizer=tokenizer,
batch_size=args.batch_size,
max_length=args.max_length,
num_workers=0,
pin_memory=False,
weighted_train_sampler=False,
)
return CLARAModel, train_loader, val_loader, cfg, int(cfg["num_classes"])
def trainable_summary(model: nn.Module) -> Dict[str, Any]:
total = sum(parameter.numel() for parameter in model.parameters())
trainable = sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad)
return {
"total_params": total,
"trainable_params": trainable,
"trainable_pct": 100.0 * trainable / max(1, total),
}
def run_pilot(args: argparse.Namespace) -> Dict[str, Any]:
torch.manual_seed(args.seed)
device = torch.device(args.device if torch.cuda.is_available() or args.device == "cpu" else "cpu")
if args.dataset == "mvsa_multiple":
model_cls, train_loader, val_loader, cfg, num_classes = load_mvsa_multiple(args)
else:
model_cls, train_loader, val_loader, cfg, num_classes = load_hfm(args)
model = model_cls(cfg).to(device)
model.train()
optimizer = AdamW(
[parameter for parameter in model.parameters() if parameter.requires_grad],
lr=float(args.learning_rate),
)
criterion = nn.CrossEntropyLoss()
metrics: List[Dict[str, Any]] = []
first_shapes: Dict[str, Any] = {}
start = time.time()
for batch_idx, batch in enumerate(train_loader, start=1):
if batch_idx > args.train_batches:
break
pixel_values = batch["pixel_values"].to(device)
input_ids = batch["input_ids"].to(device)
attention_mask = batch["attention_mask"].to(device)
labels = batch["labels"].long().to(device)
optimizer.zero_grad(set_to_none=True)
try:
outputs = model(
pixel_values=pixel_values,
input_ids=input_ids,
attention_mask=attention_mask,
return_attention=(batch_idx == 1),
)
except TypeError:
outputs = model(
pixel_values=pixel_values,
input_ids=input_ids,
attention_mask=attention_mask,
)
if batch_idx == 1:
first_shapes = shape_summary(outputs)
if {"visual_logits", "text_logits", "pred_logits", "logits"}.issubset(outputs):
loss, parts = token_loss(outputs, labels, criterion)
else:
loss = criterion(outputs["logits"], labels)
parts = {}
loss.backward()
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
preds = outputs["logits"].argmax(dim=-1)
acc = (preds == labels).float().mean().item()
metrics.append(
{
"dataset": args.dataset,
"architecture": args.architecture,
"phase": "train",
"batch": batch_idx,
"loss": float(loss.detach().item()),
"accuracy": float(acc),
"grad_norm": float(grad_norm),
"labels": json.dumps(Counter(labels.detach().cpu().tolist()), sort_keys=True),
"num_classes": num_classes,
**{key: float(value.item()) for key, value in parts.items()},
}
)
model.eval()
correct = 0
total = 0
val_losses: List[float] = []
with torch.no_grad():
for batch_idx, batch in enumerate(val_loader, start=1):
if batch_idx > args.val_batches:
break
pixel_values = batch["pixel_values"].to(device)
input_ids = batch["input_ids"].to(device)
attention_mask = batch["attention_mask"].to(device)
labels = batch["labels"].long().to(device)
outputs = model(pixel_values=pixel_values, input_ids=input_ids, attention_mask=attention_mask)
loss = criterion(outputs["logits"], labels)
val_losses.append(float(loss.item()))
preds = outputs["logits"].argmax(dim=-1)
correct += int((preds == labels).sum().item())
total += int(labels.numel())
elapsed = time.time() - start
summary = {
"dataset": args.dataset,
"architecture": args.architecture,
"train_batches": min(args.train_batches, len(train_loader)),
"val_batches": min(args.val_batches, len(val_loader)),
"val_accuracy": correct / max(1, total),
"val_loss_mean": sum(val_losses) / max(1, len(val_losses)),
"elapsed_seconds": elapsed,
**trainable_summary(model),
}
return {"metrics": metrics, "summary": summary, "shapes": first_shapes}
def main() -> None:
args = parse_args()
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
result = run_pilot(args)
suffix = f"{args.dataset}_{args.architecture}"
metric_fields = [
"dataset",
"architecture",
"phase",
"batch",
"loss",
"accuracy",
"grad_norm",
"labels",
"num_classes",
"loss_refined",
"loss_primary",
"loss_visual",
"loss_text",
]
write_csv(out_dir / f"pilot_metrics_{suffix}.csv", result["metrics"], metric_fields)
(out_dir / f"pilot_summary_{suffix}.json").write_text(
json.dumps(result["summary"], indent=2),
encoding="utf-8",
)
(out_dir / f"pilot_tensor_shapes_{suffix}.json").write_text(
json.dumps(result["shapes"], indent=2),
encoding="utf-8",
)
print(json.dumps(result["summary"], indent=2))
print(f"Wrote pilot outputs to {out_dir}")
if __name__ == "__main__":
main()