Download scripts/materialize_chartqa_opdvr_config.py from Jack04810/agentic-rl-main: direct link, hf CLI and curl.
- Browser
- Download file 6.74 kB
-
https://huggingface.co/Jack04810/agentic-rl-main/resolve/main/scripts/materialize_chartqa_opdvr_config.py
- Command line
-
hf download hf://Jack04810/agentic-rl-main/scripts/materialize_chartqa_opdvr_config.py
-
curl -L -o materialize_chartqa_opdvr_config.py https://huggingface.co/Jack04810/agentic-rl-main/resolve/main/scripts/materialize_chartqa_opdvr_config.py
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() | |