#!/usr/bin/env python """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 ''} " 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())