#!/usr/bin/env python3 """Train CLARA on HFM using code converted from notebooks/CLARA_HFM.ipynb.""" from __future__ import annotations import argparse import json from pathlib import Path from typing import Any, Dict import torch from transformers import CLIPProcessor, DebertaV2Tokenizer import sys PROJECT_ROOT = Path(__file__).resolve().parents[1] if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) from src.hfm_pipeline import ( CLARAModel, DEFAULT_HFM_CONFIG, HFMLoader, Trainer, create_dataloaders, estimate_max_length, resolve_device, set_seed, summarize_splits, ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Train CLARA on HFM") parser.add_argument("--data-root", default="data/HFM", help="Path to HFM root folder") parser.add_argument( "--text-dir", default=None, help="Path to HFM text folder (default: /text)", ) parser.add_argument( "--split-manifest", default=None, help="Optional HFM manifest CSV, e.g. ravel_revision_results/data_audit/hfm_split_manifest_deleaked.csv", ) parser.add_argument("--output-dir", default="outputs/hfm", help="Training output directory") parser.add_argument( "--checkpoint-name", default="clara_hfm.pt", help="Checkpoint filename inside output-dir", ) parser.add_argument("--batch-size", type=int, default=32) parser.add_argument("--max-epochs", type=int, default=50) parser.add_argument("--learning-rate", type=float, default=1e-4) parser.add_argument("--weight-decay", type=float, default=0.01) parser.add_argument("--warmup-ratio", type=float, default=0.1) parser.add_argument("--scheduler-type", choices=["linear", "cosine"], default=None) parser.add_argument("--early-stopping-patience", type=int, default=10) parser.add_argument( "--architecture", choices=["token", "legacy_global"], default=None, help="token: revised token-level RAVEL; legacy_global: audited global-vector baseline", ) parser.add_argument("--num-workers", type=int, default=8) parser.add_argument("--max-length", type=int, default=None) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--device", default="auto", help="auto|cuda|cpu") parser.add_argument("--hidden-dim", type=int, default=512) parser.add_argument("--lora-rank", type=int, default=8) parser.add_argument("--lora-alpha", type=int, default=None) parser.add_argument("--lora-dropout", type=float, default=None) parser.add_argument("--vision-model-id", default=None) parser.add_argument("--text-model-id", default=None) parser.add_argument("--label-smoothing", type=float, default=None) parser.add_argument("--paper-loss-mode", action="store_true") parser.add_argument("--feedback-loss-weight", type=float, default=0.5) parser.add_argument("--disagreement-tau", type=float, default=1.0) parser.add_argument("--lambda-primary", type=float, default=0.5) parser.add_argument("--lambda-unimodal", type=float, default=0.25) parser.add_argument("--loss-verify-weight", type=float, default=0.35) parser.add_argument("--loss-consistency-weight", type=float, default=0.25) parser.add_argument("--contrastive-weight", type=float, default=None) parser.add_argument( "--text-unfreeze-mode", choices=["freeze_all", "freeze_lower", "unfreeze_all"], default="freeze_all", ) parser.add_argument( "--disable-weighted-sampler", action="store_true", help="Disable weighted sampler on train split", ) return parser.parse_args() def build_config(args: argparse.Namespace) -> Dict[str, Any]: cfg = dict(DEFAULT_HFM_CONFIG) text_dir = args.text_dir or str(Path(args.data_root) / "text") cfg.update( { "text_dir": text_dir, "image_root": args.data_root, "split_manifest": args.split_manifest, "architecture": args.architecture or cfg.get("architecture", "token"), "num_classes": 2, "hidden_dim": args.hidden_dim, "batch_size": args.batch_size, "max_epochs": args.max_epochs, "learning_rate": args.learning_rate, "weight_decay": args.weight_decay, "warmup_ratio": args.warmup_ratio, "scheduler_type": args.scheduler_type or cfg.get("scheduler_type", "linear"), "early_stopping_patience": args.early_stopping_patience, "num_workers": args.num_workers, "max_length": args.max_length, "seed": args.seed, "text_unfreeze_mode": args.text_unfreeze_mode, "paper_loss_mode": args.paper_loss_mode, "feedback_loss_weight": args.feedback_loss_weight, "disagreement_tau": args.disagreement_tau, "lambda_primary": args.lambda_primary, "lambda_unimodal": args.lambda_unimodal, "loss_verify_weight": args.loss_verify_weight, "loss_consistency_weight": args.loss_consistency_weight, } ) if args.vision_model_id: cfg["vision_model_id"] = args.vision_model_id if args.text_model_id: cfg["text_model_id"] = args.text_model_id if args.label_smoothing is not None: cfg["label_smoothing"] = float(args.label_smoothing) if args.contrastive_weight is not None: cfg["contrastive_weight"] = float(args.contrastive_weight) rank = int(args.lora_rank) alpha = int(args.lora_alpha) if args.lora_alpha is not None else int(rank * 2) cfg["lora_clip"] = dict(cfg.get("lora_clip", {})) cfg["lora_deb"] = dict(cfg.get("lora_deb", {})) cfg["lora_clip"]["r"] = rank cfg["lora_deb"]["r"] = rank cfg["lora_clip"]["alpha"] = alpha cfg["lora_deb"]["alpha"] = alpha if args.lora_dropout is not None: cfg["lora_clip"]["dropout"] = float(args.lora_dropout) cfg["lora_deb"]["dropout"] = float(args.lora_dropout) return cfg def main() -> None: args = parse_args() cfg = build_config(args) set_seed(cfg["seed"]) device = resolve_device(args.device) print(f"Device: {device}") print(f"Data root: {cfg['image_root']}") print(f"Text dir: {cfg['text_dir']}") loader = HFMLoader(cfg["text_dir"], cfg["image_root"]) if cfg.get("split_manifest"): all_samples = loader.load_from_manifest(str(cfg["split_manifest"])) else: all_samples = loader.load() train_samples = loader.get_split("train") val_samples = loader.get_split("val") test_samples = loader.get_split("test") if not train_samples or not val_samples or not test_samples: raise RuntimeError("Missing train/val/test samples. Please check HFM data extraction.") split_stats = summarize_splits(train_samples, val_samples, test_samples) print("Split stats:") print(json.dumps(split_stats, indent=2)) clip_processor = CLIPProcessor.from_pretrained(cfg["vision_model_id"]) tokenizer = DebertaV2Tokenizer.from_pretrained(cfg["text_model_id"]) max_length = cfg["max_length"] if not max_length: max_length = estimate_max_length( all_samples, percentile=cfg["max_length_percentile"], sample_size=cfg["max_length_sample_size"], ) cfg["max_length"] = int(max_length) pin_memory = bool(cfg["pin_memory"] and device.type == "cuda") train_loader, val_loader, _ = create_dataloaders( train_samples=train_samples, val_samples=val_samples, test_samples=test_samples, clip_processor=clip_processor, tokenizer=tokenizer, batch_size=cfg["batch_size"], max_length=cfg["max_length"], num_workers=cfg["num_workers"], pin_memory=pin_memory, weighted_train_sampler=not args.disable_weighted_sampler, ) model = CLARAModel(cfg).to(device) stats = model.parameter_stats() pct = 100.0 * stats["trainable"] / max(1, stats["total"]) print( f"Model params: total={stats['total']:,}, trainable={stats['trainable']:,} ({pct:.2f}%)" ) use_bf16 = bool( cfg["amp_prefer_bf16"] and device.type == "cuda" and torch.cuda.is_bf16_supported() ) trainer = Trainer( model=model, train_loader=train_loader, val_loader=val_loader, cfg=cfg, device=device, output_dir=args.output_dir, checkpoint_name=args.checkpoint_name, use_bf16=use_bf16, ) result = trainer.train() cfg_path = Path(args.output_dir) / "train_config_used.json" cfg_path.parent.mkdir(parents=True, exist_ok=True) with cfg_path.open("w", encoding="utf-8") as f: json.dump(cfg, f, indent=2) print("Training complete") print(f"Best Val F1-Macro: {result['best_val_f1']:.4f}") print(f"Checkpoint: {result['checkpoint_path']}") print(f"History: {result['history_path']}") print(f"Config: {cfg_path}") if __name__ == "__main__": main()