#!/usr/bin/env python3 """Materialize controlled full-data ChartQA OPDVR experiment recipes.""" from __future__ import annotations import argparse import os import sys from pathlib import Path from typing import Any import yaml PROJECT_ROOT = Path(__file__).resolve().parents[1] if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) from config.loader import load_config, validate_config ARMS = ("reward_gate", "reward_gate_trajectory") def _require_student_checkpoint(path: Path) -> Path: resolved = path.expanduser().resolve(strict=False) if not (resolved / "config.json").is_file(): raise FileNotFoundError(f"student checkpoint lacks config.json: {resolved}") if not any(resolved.glob("*.safetensors")) and not any(resolved.glob("*.bin")): raise FileNotFoundError(f"student checkpoint lacks model weights: {resolved}") return resolved def materialize(args: argparse.Namespace) -> dict[str, Any]: config = load_config(str(args.base_config)) student_checkpoint = _require_student_checkpoint(args.student_checkpoint) output_dir = args.output_dir.expanduser().resolve(strict=False) config["model"]["pretrained_model_path"] = str(student_checkpoint) config["model"]["teacher_device_map"] = "local" training = config["training"] training["num_gpus"] = 4 training["num_client"] = 4 train_args = training["dyme_args"] train_args.update( output_dir=str(output_dir), num_train_epochs=float(args.num_train_epochs), learning_rate=float(args.learning_rate), per_device_train_batch_size=2, gradient_accumulation_steps=8, num_generations=4, max_completion_length=96, temperature=0.5, repetition_penalty=1.0, logging_steps=1, save_total_limit=3, seed=42, ) train_args.pop("resume_from_checkpoint", None) loss = config["opsd"]["loss"] loss.update( loss_type="opdvr", opsd_weight=1.0, grpo_weight=0.0, sft_weight=0.0, acc_gate=False, reward_gate={ "enabled": True, "correct_threshold": 0.5, "token_scope": "answer_span", # Full-data OPDVR uses the verifier sign constraint rather than # dropping source rows where the independent teacher was wrong. "require_teacher_correct": False, "require_teacher_gate_metadata": True, # Zero means exact, unclipped log-probability ratios. "max_abs_log_ratio": 0.0, }, ) config["opsd"]["teacher_probe"]["enabled"] = False config["opsd"]["teacher_trajectory"]["enabled"] = ( args.arm == "reward_gate_trajectory" ) config["opsd"]["teacher_trajectory"]["context_providers"] = [ "visual_facts_deplot", "format_only", ] # Generate the four unique prompts in each rollout micro-batch together. # This is the controlled trajectory-only difference between the arms; it # avoids four serial 7B generate calls without changing any samples. config["opsd"]["teacher_trajectory"]["batch_size"] = 4 config["opsd"]["teacher_trajectory"]["max_loss_rows_per_rank"] = 2 config["opsd"]["teacher_trajectory"]["loss_type"] = "fkl" config["opsd"]["teacher_trajectory"]["weight"] = 0.5 config["opsd"]["visual_supervision"]["enabled"] = False config["opsd"]["visual_supervision"]["checker"]["enabled"] = False config["opsd"]["visual_supervision"]["refiner"]["enabled"] = False teacher_gate = config["dataset"]["teacher_gate"] teacher_gate.update( enabled=True, selection_mode="all", expected_selected_rows=4576, ) config["dataset"]["max_train_samples"] = None config["data_validation"]["expected_samples"] = 4576 config["data_validation"]["require_qwen_rewrite"] = False checkpoint_eval = config["checkpoint_eval"] checkpoint_eval.update( enabled=not args.smoke, initial_eval=not args.smoke, gate_steps=[10, 25, 50, 75], regular_steps=100, patience=3, patience_start_step=100, tie_policy="count", ) if args.smoke: train_args["max_steps"] = int(args.max_steps) train_args["num_train_epochs"] = 1.0 train_args["save_strategy"] = "no" config["opsd"]["debug"]["detail_every"] = 1 config["opsd"]["debug"]["distributed_trace"]["cuda_sync"] = True else: train_args.pop("max_steps", None) train_args["save_strategy"] = "steps" train_args["save_steps"] = 100 config["opsd"]["debug"]["detail_every"] = 10 config["opsd"]["debug"]["distributed_trace"]["cuda_sync"] = False validated = validate_config(config, source=str(args.output_config)) target = args.output_config.expanduser().resolve(strict=False) target.parent.mkdir(parents=True, exist_ok=True) temporary = target.with_name(f".{target.name}.tmp-{os.getpid()}") with temporary.open("w", encoding="utf-8") as handle: yaml.safe_dump(validated, handle, allow_unicode=True, sort_keys=False) os.replace(temporary, target) return validated def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--base-config", type=Path, required=True) parser.add_argument("--student-checkpoint", type=Path, required=True) parser.add_argument("--output-config", type=Path, required=True) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--arm", choices=ARMS, required=True) parser.add_argument("--num-train-epochs", type=float, default=4.0) parser.add_argument("--learning-rate", type=float, default=1.0e-6) parser.add_argument("--smoke", action="store_true") parser.add_argument("--max-steps", type=int, default=4) return parser.parse_args() def main() -> None: args = parse_args() if args.num_train_epochs <= 0: raise ValueError("--num-train-epochs must be positive") if args.learning_rate <= 0: raise ValueError("--learning-rate must be positive") if args.smoke and args.max_steps <= 0: raise ValueError("--max-steps must be positive for smoke runs") config = materialize(args) print( "Materialized ChartQA OPDVR recipe: " f"arm={args.arm} smoke={args.smoke} " f"trajectory={config['opsd']['teacher_trajectory']['enabled']} " f"rows={config['dataset']['teacher_gate']['expected_selected_rows']} " f"student={config['model']['pretrained_model_path']} " f"output={config['training']['dyme_args']['output_dir']}", flush=True, ) if __name__ == "__main__": main()