| |
| """Train CLARA on MVSA-Multiple using script workflow converted from notebook.""" |
|
|
| 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.mvsa_multiple_pipeline import ( |
| CLARAModel, |
| DEFAULT_MVSA_MULTIPLE_CONFIG, |
| MVSALoader, |
| Trainer, |
| create_dataloaders, |
| resolve_device, |
| set_seed, |
| summarize_splits, |
| ) |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description="Train CLARA on MVSA-Multiple") |
| parser.add_argument("--data-root", default="data/MVSA-Multiple") |
| parser.add_argument("--text-dir", default=None, help="Default: <data-root>/data") |
| parser.add_argument("--label-file", default=None, help="Default: <data-root>/labelResultAll.txt") |
| 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("--output-dir", default="outputs/mvsa_multiple") |
| parser.add_argument("--checkpoint-name", default="clara_mvsa_multiple.pt") |
| parser.add_argument( |
| "--init-checkpoint", |
| default=None, |
| help="Optional checkpoint path to initialize model weights before training.", |
| ) |
|
|
| parser.add_argument("--batch-size", type=int, default=48) |
| parser.add_argument("--max-epochs", type=int, default=50) |
| parser.add_argument("--learning-rate", type=float, default=5e-5) |
| parser.add_argument("--weight-decay", type=float, default=0.01) |
| parser.add_argument("--warmup-ratio", type=float, default=0.16) |
| parser.add_argument("--scheduler-type", choices=["linear", "cosine"], default=None) |
| parser.add_argument("--early-stopping-patience", type=int, default=12) |
| parser.add_argument("--grad-accum-steps", type=int, default=2) |
|
|
| parser.add_argument("--max-length", type=int, default=128) |
| parser.add_argument("--num-workers", type=int, default=6) |
| parser.add_argument("--seed", type=int, default=42) |
| parser.add_argument("--device", default="auto", help="auto|cuda|cpu") |
|
|
| parser.add_argument("--train-ratio", type=float, default=0.8) |
| parser.add_argument("--val-ratio", type=float, default=0.1) |
| parser.add_argument( |
| "--preprocessing-mode", |
| choices=["paper", "strict"], |
| default="paper", |
| help="paper: paper-style MVSA preprocessing; strict: unanimous + cross-agree filtering", |
| ) |
| parser.add_argument("--disable-paper-exact-counts", action="store_true") |
| parser.add_argument("--allow-non-unanimous", action="store_true") |
| parser.add_argument("--allow-cross-disagree", action="store_true") |
|
|
| parser.add_argument("--negative-class-boost", type=float, default=12.0) |
| parser.add_argument("--negative-focal-gamma", type=float, default=4.0) |
| parser.add_argument("--min-ratio-negative", type=float, default=0.30) |
| parser.add_argument("--label-smoothing", type=float, default=0.02) |
| parser.add_argument("--consistency-weight", type=float, default=0.10) |
| parser.add_argument("--pred-veri-weight", type=float, default=0.30) |
| parser.add_argument("--feedback-tau", type=float, default=0.5) |
| parser.add_argument("--disagreement-tau", type=float, default=1.0) |
|
|
| parser.add_argument("--disable-mixup-negative", action="store_true") |
| parser.add_argument("--disable-weighted-sampler", action="store_true") |
| parser.add_argument("--mixup-alpha", type=float, default=0.4) |
| parser.add_argument("--paper-loss-mode", action="store_true") |
| parser.add_argument("--feedback-loss-weight", type=float, default=0.5) |
| parser.add_argument("--lambda-primary", type=float, default=0.5) |
| parser.add_argument("--lambda-unimodal", type=float, default=0.25) |
|
|
| parser.add_argument("--disable-clip-lora", action="store_true") |
| parser.add_argument("--hidden-dim", type=int, default=512) |
| parser.add_argument( |
| "--unfreeze-epoch", |
| type=int, |
| default=0, |
| help="0 disables non-LoRA backbone unfreezing; set >0 only for full/unfrozen ablations.", |
| ) |
| parser.add_argument("--unfreeze-vision-backbone", action="store_true") |
| 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( |
| "--text-unfreeze-mode", |
| choices=["freeze_all", "freeze_lower", "unfreeze_all"], |
| default="freeze_all", |
| ) |
| return parser.parse_args() |
|
|
|
|
| def build_config(args: argparse.Namespace) -> Dict[str, Any]: |
| cfg = dict(DEFAULT_MVSA_MULTIPLE_CONFIG) |
|
|
| text_dir = args.text_dir or str(Path(args.data_root) / "data") |
| label_file = args.label_file or str(Path(args.data_root) / "labelResultAll.txt") |
|
|
| cfg.update( |
| { |
| "text_dir": text_dir, |
| "label_file": label_file, |
| "architecture": args.architecture or cfg.get("architecture", "token"), |
| "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, |
| "grad_accum_steps": args.grad_accum_steps, |
| "max_length": args.max_length, |
| "num_workers": args.num_workers, |
| "seed": args.seed, |
| "train_ratio": args.train_ratio, |
| "val_ratio": args.val_ratio, |
| "preprocessing_mode": args.preprocessing_mode, |
| "paper_exact_counts": not args.disable_paper_exact_counts, |
| "require_unanimous": not args.allow_non_unanimous, |
| "require_cross_agree": not args.allow_cross_disagree, |
| "negative_class_boost": args.negative_class_boost, |
| "negative_focal_gamma": args.negative_focal_gamma, |
| "min_ratio_negative": args.min_ratio_negative, |
| "label_smoothing": args.label_smoothing, |
| "consistency_weight": args.consistency_weight, |
| "pred_veri_weight": args.pred_veri_weight, |
| "feedback_tau": args.feedback_tau, |
| "disagreement_tau": args.disagreement_tau, |
| "use_mixup_negative": not args.disable_mixup_negative, |
| "use_weighted_sampler": not args.disable_weighted_sampler, |
| "mixup_alpha": args.mixup_alpha, |
| "paper_loss_mode": args.paper_loss_mode, |
| "feedback_loss_weight": args.feedback_loss_weight, |
| "lambda_primary": args.lambda_primary, |
| "lambda_unimodal": args.lambda_unimodal, |
| "enable_clip_lora": not args.disable_clip_lora, |
| "unfreeze_epoch": int(args.unfreeze_epoch), |
| "unfreeze_vision_backbone": bool(args.unfreeze_vision_backbone), |
| "text_unfreeze_mode": args.text_unfreeze_mode, |
| } |
| ) |
| 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 |
| 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(int(cfg["seed"])) |
| device = resolve_device(args.device) |
|
|
| print(f"Device: {device}") |
| print(f"Text dir: {cfg['text_dir']}") |
| print(f"Label file: {cfg['label_file']}") |
|
|
| 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", False)), |
| ) |
| 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=bool(str(cfg.get("preprocessing_mode", "paper")).lower() == "paper"), |
| ) |
|
|
| if not train_samples or not val_samples or not test_samples: |
| raise RuntimeError("Missing train/val/test split after filtering. Please check dataset/filter setup.") |
|
|
| 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"]) |
|
|
| 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=int(cfg["batch_size"]), |
| max_length=int(cfg["max_length"]), |
| num_workers=int(cfg["num_workers"]), |
| pin_memory=pin_memory, |
| persistent_workers=bool(cfg["persistent_workers"]), |
| prefetch_factor=int(cfg["prefetch_factor"]), |
| use_mixup_negative=bool(cfg["use_mixup_negative"]), |
| mixup_alpha=float(cfg["mixup_alpha"]), |
| negative_class_boost=float(cfg["negative_class_boost"]), |
| min_ratio_negative=float(cfg["min_ratio_negative"]), |
| weighted_train_sampler=bool(cfg.get("use_weighted_sampler", True)), |
| ) |
|
|
| model = CLARAModel(cfg).to(device) |
| if args.init_checkpoint: |
| init_path = Path(args.init_checkpoint) |
| if not init_path.exists(): |
| raise FileNotFoundError(f"Init checkpoint not found: {init_path}") |
| init_payload = torch.load(str(init_path), map_location="cpu") |
| init_state = init_payload.get("model_state", init_payload) |
| missing, unexpected = model.load_state_dict(init_state, strict=False) |
| print(f"Initialized from checkpoint: {init_path}") |
| print(f"- Missing keys: {len(missing)}") |
| print(f"- Unexpected keys: {len(unexpected)}") |
|
|
| 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(device.type == "cuda" and torch.cuda.is_bf16_supported()) |
|
|
| trainer = Trainer( |
| model=model, |
| train_loader=train_loader, |
| val_loader=val_loader, |
| train_samples=train_samples, |
| 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) |
| cfg_path.write_text(json.dumps(cfg, indent=2), encoding="utf-8") |
|
|
| print("Training complete") |
| print(f"Best Val F1-Weighted: {result['best_val_f1_weighted']:.4f}") |
| print(f"Checkpoint: {result['checkpoint_path']}") |
| print(f"History: {result['history_path']}") |
| print(f"Config: {cfg_path}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|