| |
| """Evaluate the released ACDiR-LLaDA MATH500 checkpoint.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import hashlib |
| import json |
| import os |
| import subprocess |
| import sys |
| from pathlib import Path |
|
|
|
|
| ROOT = Path(__file__).resolve().parent |
| DEFAULT_CONFIG = ROOT / "configs" / "math500_44.json" |
| HASH_CHUNK_SIZE = 1024 * 1024 |
|
|
|
|
| def _bool_arg(value: bool) -> str: |
| return "True" if bool(value) else "False" |
|
|
|
|
| def _bool_override(value: str, default: bool) -> bool: |
| text = str(value or "").strip().lower() |
| if not text: |
| return bool(default) |
| if text in {"1", "true", "yes", "on"}: |
| return True |
| if text in {"0", "false", "no", "off"}: |
| return False |
| raise ValueError(f"Invalid boolean override: {value!r}") |
|
|
|
|
| def _int_override(value: str, default: int) -> int: |
| text = str(value or "").strip() |
| if not text: |
| return int(default) |
| return int(text) |
|
|
|
|
| def _float_override(value: str, default: float) -> float: |
| text = str(value or "").strip() |
| if not text: |
| return float(default) |
| return float(text) |
|
|
|
|
| def _sha256_file(path: Path) -> str: |
| digest = hashlib.sha256() |
| with path.open("rb") as handle: |
| for chunk in iter(lambda: handle.read(HASH_CHUNK_SIZE), b""): |
| digest.update(chunk) |
| return digest.hexdigest() |
|
|
|
|
| def resolve_base_model(model: str, revision: str = "", cache_dir: Path | None = None) -> str: |
| """Return a local model directory, downloading an HF repo when necessary.""" |
| candidate = Path(model).expanduser() |
| if candidate.exists(): |
| return str(candidate.resolve()) |
| if not model or "/" not in model: |
| raise FileNotFoundError( |
| f"Base model is neither a local path nor a Hugging Face repo id: {model!r}" |
| ) |
| from huggingface_hub import snapshot_download |
|
|
| resolved = snapshot_download( |
| repo_id=model, |
| repo_type="model", |
| revision=revision or None, |
| cache_dir=str(cache_dir) if cache_dir is not None else None, |
| ) |
| return str(Path(resolved).resolve()) |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--config", default=str(DEFAULT_CONFIG)) |
| parser.add_argument("--base_model", default="") |
| parser.add_argument("--critic_ckpt", default=os.environ.get("EVAL_CRITIC_CKPT", "")) |
| parser.add_argument("--dataset", default=os.environ.get("EVAL_DATASET", "")) |
| parser.add_argument("--batch_size", type=int, default=0) |
| parser.add_argument("--nproc_per_node", type=int, default=1) |
| parser.add_argument("--master_port", type=int, default=29517) |
| parser.add_argument("--max_eval_samples", type=int, default=0) |
| parser.add_argument("--debug_samples", type=int, default=0) |
| parser.add_argument("--compare_with_baseline", default=os.environ.get("EVAL_COMPARE_WITH_BASELINE", "")) |
| parser.add_argument("--lookback_blocks", default=os.environ.get("EVAL_LOOKBACK_BLOCKS", "")) |
| parser.add_argument("--remask_min_age_current", default=os.environ.get("EVAL_REMASK_MIN_AGE_CURRENT", "")) |
| parser.add_argument("--remask_max_age_lookback", default=os.environ.get("EVAL_REMASK_MAX_AGE_LOOKBACK", "")) |
| parser.add_argument("--max_total_remask_per_sample", default=os.environ.get("EVAL_MAX_TOTAL_REMASK_PER_SAMPLE", "")) |
| parser.add_argument("--force_remask_window", default=os.environ.get("EVAL_FORCE_REMASK_WINDOW", "")) |
| parser.add_argument("--reforward_after_remask", default=os.environ.get("EVAL_REFORWARD_AFTER_REMASK", "")) |
| parser.add_argument("--deterministic_joint_argmax", default=os.environ.get("EVAL_DETERMINISTIC_JOINT_ARGMAX", "")) |
| parser.add_argument("--sample_remask", default=os.environ.get("EVAL_SAMPLE_REMASK", "")) |
| parser.add_argument("--remask_temperature", default=os.environ.get("EVAL_REMASK_TEMPERATURE", "")) |
| parser.add_argument("--remask_timing", default=os.environ.get("EVAL_REMASK_TIMING", "")) |
| parser.add_argument("--count_logit_bias", default=os.environ.get("EVAL_COUNT_LOGIT_BIAS", "")) |
| parser.add_argument("--clean_output", default=os.environ.get("EVAL_CLEAN_OUTPUT", "True")) |
| parser.add_argument("--progress_every", type=int, default=int(os.environ.get("EVAL_PROGRESS_EVERY", "50") or 50)) |
| parser.add_argument("--result_dir", default="outputs/math500_eval") |
| return parser.parse_args() |
|
|
|
|
| def main() -> int: |
| args = parse_args() |
| config_path = Path(args.config) |
| if not config_path.is_absolute(): |
| config_path = ROOT / config_path |
| with config_path.open("r", encoding="utf-8") as f: |
| cfg = json.load(f) |
|
|
| eval_cfg = cfg["eval"] |
| decode_cfg = cfg["decode"] |
| runtime_cfg = cfg["runtime"] |
| weights = cfg["released_weights"] |
|
|
| base_model_source = args.base_model or cfg["base_model"] |
| base_model_revision = "" if args.base_model else str(cfg.get("base_model_revision", "")) |
| critic_ckpt = Path(args.critic_ckpt or weights["critic"]) |
| dataset = Path(args.dataset or "datasets/MATH500") |
| result_dir = Path(args.result_dir) |
|
|
| if not critic_ckpt.is_absolute(): |
| critic_ckpt = ROOT / critic_ckpt |
| if not dataset.is_absolute(): |
| dataset = ROOT / dataset |
| if not result_dir.is_absolute(): |
| result_dir = ROOT / result_dir |
| result_dir.mkdir(parents=True, exist_ok=True) |
|
|
| base_model = resolve_base_model( |
| str(base_model_source), |
| revision=base_model_revision, |
| cache_dir=ROOT / ".cache" / "huggingface" / "hub", |
| ) |
|
|
| if not critic_ckpt.exists(): |
| raise FileNotFoundError(f"Missing critic checkpoint: {critic_ckpt}") |
| if not dataset.exists(): |
| raise FileNotFoundError(f"Missing dataset: {dataset}") |
|
|
| critic_sha256 = _sha256_file(critic_ckpt) |
| batch_size = int(args.batch_size or eval_cfg["batch_size"]) |
| compare_with_baseline = _bool_override(args.compare_with_baseline, eval_cfg["compare_with_baseline"]) |
| lookback_blocks = _int_override(args.lookback_blocks, decode_cfg["lookback_blocks"]) |
| remask_min_age_current = _int_override(args.remask_min_age_current, decode_cfg["remask_min_age_current"]) |
| remask_max_age_lookback = _int_override(args.remask_max_age_lookback, decode_cfg["remask_max_age_lookback"]) |
| max_total_remask_per_sample = _int_override(args.max_total_remask_per_sample, decode_cfg["max_total_remask_per_sample"]) |
| force_remask_window = _int_override(args.force_remask_window, decode_cfg["force_remask_window"]) |
| reforward_after_remask = _bool_override(args.reforward_after_remask, decode_cfg["reforward_after_remask"]) |
| deterministic_joint_argmax = _bool_override( |
| args.deterministic_joint_argmax, |
| decode_cfg["deterministic_joint_argmax"], |
| ) |
| sample_remask = _bool_override(args.sample_remask, decode_cfg["sample_remask"]) |
| remask_temperature = _float_override(args.remask_temperature, decode_cfg["remask_temperature"]) |
| remask_timing = str(args.remask_timing or decode_cfg.get("remask_timing", "step")).strip().lower().replace("-", "_") |
| if remask_timing in {"blockend", "block_final", "end_of_block"}: |
| remask_timing = "block_end" |
| if remask_timing not in {"step", "block_end"}: |
| raise ValueError("remask_timing must be one of: step, block_end.") |
| count_logit_bias = str(args.count_logit_bias or decode_cfg.get("count_logit_bias", "")).strip() |
| clean_output = _bool_override(args.clean_output, True) |
| nproc = max(1, int(args.nproc_per_node)) |
|
|
| env = os.environ.copy() |
| env.setdefault("LLADA_EXACT_BACKEND", runtime_cfg["llada_exact_backend"]) |
| env.setdefault("LLADA_LMDEPLOY_FAST_MODE", runtime_cfg["llada_fast_mode"]) |
| env.setdefault("LLADA_LMDEPLOY_CUDAGRAPH", "1" if runtime_cfg["lmdeploy_cuda_graph"] else "0") |
| env.setdefault("LLADA_LMDEPLOY_VARLEN_FLASH", "1" if runtime_cfg.get("varlen_flash", False) else "0") |
| env.setdefault("ACDIR_DIST_TIMEOUT_MIN", "120") |
| env.setdefault("HF_HOME", str(ROOT / ".cache" / "huggingface")) |
|
|
| eval_script = ROOT / "metrics" / "phase2_critic_guided_math.py" |
| eval_cmd = [ |
| str(eval_script), |
| "--ckpt_path", |
| str(base_model), |
| "--critic_ckpt_path", |
| str(critic_ckpt), |
| "--local_data_path", |
| str(dataset), |
| "--batch_size", |
| str(batch_size), |
| "--num_workers", |
| str(eval_cfg["num_workers"]), |
| "--seed", |
| str(eval_cfg["seed"]), |
| "--steps", |
| str(eval_cfg["steps"]), |
| "--gen_length", |
| str(eval_cfg["gen_length"]), |
| "--block_length", |
| str(eval_cfg["block_length"]), |
| "--block_steps", |
| str(eval_cfg["block_steps"]), |
| "--no_sample", |
| _bool_arg(eval_cfg["no_sample"]), |
| "--temperature", |
| str(eval_cfg["temperature"]), |
| "--cfg_scale", |
| str(eval_cfg["cfg_scale"]), |
| "--actor_type", |
| "llada", |
| "--mask_id", |
| str(runtime_cfg["mask_id"]), |
| "--eos_id", |
| str(runtime_cfg["eos_id"]), |
| "--unmask_policy", |
| "confidence", |
| "--remask_method", |
| decode_cfg["remask_method"], |
| "--ablation_remask_probability", |
| str(decode_cfg["ablation_remask_probability"]), |
| "--remask_candidate_disagree_only", |
| _bool_arg(decode_cfg["remask_candidate_disagree_only"]), |
| "--remask_candidate_max_confidence", |
| str(decode_cfg["remask_candidate_max_confidence"]), |
| "--sample_remask", |
| _bool_arg(sample_remask), |
| "--remask_temperature", |
| str(remask_temperature), |
| "--remask_timing", |
| remask_timing, |
| f"--count_logit_bias={count_logit_bias}", |
| "--lookback_blocks", |
| str(lookback_blocks), |
| "--remask_min_age_current", |
| str(remask_min_age_current), |
| "--remask_max_age_lookback", |
| str(remask_max_age_lookback), |
| "--deterministic_joint_argmax", |
| _bool_arg(deterministic_joint_argmax), |
| "--force_remask_window", |
| str(force_remask_window), |
| "--max_total_remask_per_sample", |
| str(max_total_remask_per_sample), |
| "--reforward_after_remask", |
| _bool_arg(reforward_after_remask), |
| "--oracle_rollouts", |
| "1", |
| "--oracle_rollout_batch_size", |
| "1", |
| "--oracle_seed_stride", |
| "1009", |
| "--max_eval_samples", |
| str(args.max_eval_samples), |
| "--prediction_dir", |
| str(result_dir / "predictions"), |
| "--debug_samples", |
| str(args.debug_samples), |
| "--eval_style", |
| eval_cfg["eval_style"], |
| "--use_chat_template", |
| _bool_arg(eval_cfg["use_chat_template"]), |
| "--prompt_style", |
| eval_cfg["prompt_style"], |
| "--compare_with_baseline", |
| _bool_arg(compare_with_baseline), |
| "--normalize_no_remask_to_baseline", |
| "False", |
| "--no_lmdeploy_cuda_graph", |
| "--llada_fast_mode", |
| runtime_cfg["llada_fast_mode"], |
| "--sdar_confidence_threshold", |
| "0.85", |
| "--actor_forward_backend", |
| runtime_cfg["actor_forward_backend"], |
| "--actor_forward_dtype", |
| runtime_cfg["actor_forward_dtype"], |
| "--clean_output", |
| _bool_arg(clean_output), |
| "--progress_every", |
| str(max(1, int(args.progress_every))), |
| ] |
|
|
| cmd = [ |
| sys.executable, |
| "-m", |
| "torch.distributed.run", |
| "--standalone", |
| f"--nproc-per-node={nproc}", |
| f"--master-port={int(args.master_port)}", |
| *eval_cmd, |
| ] |
| log_path = result_dir / "eval_command.txt" |
| log_path.write_text(" ".join(cmd) + "\n", encoding="utf-8") |
| print(f"[acdir] critic checkpoint: {critic_ckpt}", flush=True) |
| print(f"[acdir] critic sha256: {critic_sha256}", flush=True) |
| print( |
| f"[acdir] base model: {base_model_source}" |
| + (f" @ {base_model_revision}" if base_model_revision else "") |
| + f" -> {base_model}", |
| flush=True, |
| ) |
| print(f"[acdir] compare_with_baseline: {compare_with_baseline}", flush=True) |
| print( |
| "[acdir] remask: " |
| f"timing={remask_timing} " |
| f"count_bias={count_logit_bias or '<none>'} " |
| f"lookback_blocks={lookback_blocks} " |
| f"age={remask_min_age_current}/{remask_max_age_lookback} " |
| f"k_cap={max_total_remask_per_sample} " |
| f"force_window={force_remask_window} " |
| f"reforward={reforward_after_remask}", |
| flush=True, |
| ) |
| print(f"[acdir] clean_output: {clean_output} progress_every={max(1, int(args.progress_every))}", flush=True) |
| print(f"[acdir] running MATH500 eval; command saved to {log_path}", flush=True) |
| return subprocess.call(cmd, cwd=str(ROOT), env=env) |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|