agentic-rl-main / scripts /materialize_chartqa_opdvr_config.py
Jack04810's picture
Add files using upload-large-folder tool
df529cc verified
Raw History Blame Contribute Delete
6.74 kB
#!/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()