| |
| """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() |
|
|