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