acdir-llada-math500 / eval_math500.py
NYCU-MLLab's picture
Upload folder using huggingface_hub
4a28d4d verified
Raw
History Blame Contribute Delete
12.7 kB
#!/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 '<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())