ravel / scripts /train_mvsa_multiple.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
12.1 kB
#!/usr/bin/env python3
"""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()