488 lines
15 KiB
Python
488 lines
15 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
"""
|
|
Evaluate a HF-format model on AIME 2025 using vLLM.
|
|
|
|
Example:
|
|
|
|
CUDA_VISIBLE_DEVICES=0,1,2,3 python tools/eval_aime_vllm.py \
|
|
--model /workspace/Lightning-OPD/checkpoints/qwen3-4b-lightning-opd-hf \
|
|
--num-gpus 4 \
|
|
--output outputs/eval_aime_2025_qwen3_4b_lightning_opd.jsonl
|
|
|
|
Paper-style AIME setting:
|
|
temperature = 0.6
|
|
top_p = 0.95
|
|
max_tokens = 32768
|
|
n_samples = 32
|
|
metric = average pass@1
|
|
"""
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
import re
|
|
from pathlib import Path
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
|
|
def parse_args():
|
|
parser = argparse.ArgumentParser()
|
|
|
|
# Model / hardware
|
|
parser.add_argument(
|
|
"--model",
|
|
type=str,
|
|
required=True,
|
|
help="Path to HF-format model checkpoint.",
|
|
)
|
|
parser.add_argument(
|
|
"--num-gpus",
|
|
type=int,
|
|
default=1,
|
|
help="Tensor parallel size for vLLM.",
|
|
)
|
|
parser.add_argument(
|
|
"--gpu-ids",
|
|
type=str,
|
|
default=None,
|
|
help='Optional CUDA_VISIBLE_DEVICES, e.g. "0,1,2,3". Must match --num-gpus.',
|
|
)
|
|
parser.add_argument(
|
|
"--dtype",
|
|
type=str,
|
|
default="bfloat16",
|
|
choices=["auto", "float16", "bfloat16", "float32"],
|
|
)
|
|
parser.add_argument(
|
|
"--gpu-memory-utilization",
|
|
type=float,
|
|
default=0.90,
|
|
)
|
|
parser.add_argument(
|
|
"--trust-remote-code",
|
|
action="store_true",
|
|
default=True,
|
|
)
|
|
|
|
# Dataset
|
|
parser.add_argument(
|
|
"--dataset",
|
|
type=str,
|
|
default="MathArena/aime_2025",
|
|
help="HF dataset name for AIME 2025.",
|
|
)
|
|
parser.add_argument(
|
|
"--split",
|
|
type=str,
|
|
default=None,
|
|
help="Dataset split. If not set, use the first available split.",
|
|
)
|
|
parser.add_argument(
|
|
"--hf-cache",
|
|
type=str,
|
|
default=None,
|
|
help="Optional HuggingFace cache dir.",
|
|
)
|
|
|
|
# Generation settings from paper
|
|
parser.add_argument("--n-samples", type=int, default=32)
|
|
parser.add_argument("--temperature", type=float, default=0.6)
|
|
parser.add_argument("--top-p", type=float, default=0.95)
|
|
parser.add_argument("--max-tokens", type=int, default=32768)
|
|
parser.add_argument("--seed", type=int, default=42)
|
|
|
|
# Prompt / tokenizer behavior
|
|
parser.add_argument(
|
|
"--no-chat-template",
|
|
action="store_true",
|
|
help="Do not apply tokenizer chat template; use raw prompt.",
|
|
)
|
|
parser.add_argument(
|
|
"--enable-thinking",
|
|
action="store_true",
|
|
default=False,
|
|
help="Pass enable_thinking=True to Qwen3 chat template if supported.",
|
|
)
|
|
parser.add_argument(
|
|
"--disable-thinking",
|
|
action="store_true",
|
|
default=False,
|
|
help="Pass enable_thinking=False to Qwen3 chat template if supported.",
|
|
)
|
|
parser.add_argument(
|
|
"--assistant-prefix",
|
|
type=str,
|
|
default=None,
|
|
help=(
|
|
"Optional text appended after the assistant generation prompt. "
|
|
"Example for no-thinking style: '<think>\\n\\n</think>\\n\\n'"
|
|
),
|
|
)
|
|
parser.add_argument(
|
|
"--prompt-template",
|
|
type=str,
|
|
default="train",
|
|
choices=["train", "paper", "chat"],
|
|
help=(
|
|
"Prompt template to use. "
|
|
"'train' uses the exact training ChatML prompt; "
|
|
"'paper' uses the paper-style ChatML prompt; "
|
|
"'chat' uses build_problem_prompt plus tokenizer.apply_chat_template."
|
|
),
|
|
)
|
|
|
|
# Output / debug
|
|
parser.add_argument(
|
|
"--output",
|
|
type=str,
|
|
default="outputs/eval_aime_2025_vllm.jsonl",
|
|
help="Path to save per-sample generations and scores.",
|
|
)
|
|
parser.add_argument(
|
|
"--limit",
|
|
type=int,
|
|
default=None,
|
|
help="Limit number of problems for quick debugging.",
|
|
)
|
|
parser.add_argument(
|
|
"--print-examples",
|
|
type=int,
|
|
default=3,
|
|
help="Print first N examples with predictions.",
|
|
)
|
|
|
|
return parser.parse_args()
|
|
|
|
|
|
def normalize_answer(x: Any) -> Optional[str]:
|
|
if x is None:
|
|
return None
|
|
|
|
s = str(x).strip()
|
|
s = s.replace("$", "")
|
|
s = s.replace(",", "")
|
|
s = s.replace("\\,", "")
|
|
s = s.strip()
|
|
|
|
# Remove simple latex wrappers.
|
|
s = s.replace("\\text", "")
|
|
s = s.replace("\\mathrm", "")
|
|
s = s.strip("{}").strip()
|
|
|
|
# AIME answers are integers from 0 to 999.
|
|
m = re.search(r"-?\d+", s)
|
|
if m is None:
|
|
return None
|
|
|
|
try:
|
|
return str(int(m.group(0)))
|
|
except ValueError:
|
|
return None
|
|
|
|
|
|
def extract_last_boxed(text: str) -> Optional[str]:
|
|
"""
|
|
Extract the last \\boxed{...}. Handles simple nested braces better than regex.
|
|
"""
|
|
marker = r"\boxed{"
|
|
positions = [m.start() for m in re.finditer(re.escape(marker), text)]
|
|
if not positions:
|
|
return None
|
|
|
|
for start in reversed(positions):
|
|
i = start + len(marker)
|
|
depth = 1
|
|
chars = []
|
|
|
|
while i < len(text):
|
|
ch = text[i]
|
|
if ch == "{":
|
|
depth += 1
|
|
chars.append(ch)
|
|
elif ch == "}":
|
|
depth -= 1
|
|
if depth == 0:
|
|
return "".join(chars)
|
|
chars.append(ch)
|
|
else:
|
|
chars.append(ch)
|
|
i += 1
|
|
|
|
return None
|
|
|
|
|
|
def extract_answer(text: str) -> Optional[str]:
|
|
# Prefer boxed answer, because the actual training prompt asks for \boxed{$Answer}.
|
|
boxed = extract_last_boxed(text)
|
|
if boxed is not None:
|
|
ans = normalize_answer(boxed)
|
|
if ans is not None:
|
|
return ans
|
|
|
|
# Fallback: final Answer: line.
|
|
matches = re.findall(r"Answer:\s*([^\n]+)", text, flags=re.IGNORECASE)
|
|
if matches:
|
|
ans = normalize_answer(matches[-1])
|
|
if ans is not None:
|
|
return ans
|
|
|
|
# No valid final answer.
|
|
return None
|
|
|
|
def build_exact_training_prompt(problem: str) -> str:
|
|
return (
|
|
"<|im_start|>user\n"
|
|
"Solve the following math problem step by step. "
|
|
"The last line of your response should be of the form "
|
|
"Answer: \\boxed{$Answer} where $Answer is the answer to the problem.\n\n"
|
|
f"{problem}\n\n"
|
|
"Remember to put your answer on its own line after \"Answer:\"."
|
|
"<|im_end|>\n"
|
|
"<|im_start|>assistant\n\n\n\n"
|
|
)
|
|
|
|
def build_problem_prompt(problem: str) -> str:
|
|
return (
|
|
f"{problem}\n\n"
|
|
"Please reason step by step, and put your final answer within \\boxed{}."
|
|
)
|
|
|
|
def build_paper_math_eval_prompt(problem: str) -> str:
|
|
return (
|
|
"<|im_start|>user\n"
|
|
f"Question: {problem}\n"
|
|
"Please reason step by step, and put your final answer within \\boxed{}.\n"
|
|
"<|im_end|>\n"
|
|
"<|im_start|>assistant\n"
|
|
)
|
|
|
|
def get_field(ex: Dict[str, Any], candidates: List[str]) -> Any:
|
|
for key in candidates:
|
|
if key in ex and ex[key] is not None:
|
|
return ex[key]
|
|
raise KeyError(f"Cannot find any of {candidates}. Example keys: {list(ex.keys())}")
|
|
|
|
# def build_problem_prompt(problem: str) -> str:
|
|
# return (
|
|
# "Solve the following math problem step by step. "
|
|
# "The last line of your response should be of the form "
|
|
# "Answer: $Answer (without quotes) where $Answer is the answer to the problem.\n\n"
|
|
# f"{problem}\n\n"
|
|
# 'Remember to put your answer on its own line after "Answer:".'
|
|
# )
|
|
|
|
def apply_chat_template(tokenizer, prompt: str, args) -> str:
|
|
if args.no_chat_template:
|
|
return prompt
|
|
|
|
messages = [{"role": "user", "content": prompt}]
|
|
|
|
kwargs = {
|
|
"tokenize": False,
|
|
"add_generation_prompt": True,
|
|
}
|
|
|
|
# Qwen3 tokenizer may support enable_thinking.
|
|
if args.enable_thinking and args.disable_thinking:
|
|
raise ValueError("Do not set both --enable-thinking and --disable-thinking.")
|
|
|
|
if args.enable_thinking:
|
|
kwargs["enable_thinking"] = True
|
|
elif args.disable_thinking:
|
|
kwargs["enable_thinking"] = False
|
|
|
|
try:
|
|
text = tokenizer.apply_chat_template(messages, **kwargs)
|
|
except TypeError:
|
|
# Some tokenizers do not accept enable_thinking.
|
|
kwargs.pop("enable_thinking", None)
|
|
text = tokenizer.apply_chat_template(messages, **kwargs)
|
|
|
|
if args.assistant_prefix is not None:
|
|
text += args.assistant_prefix.encode("utf-8").decode("unicode_escape")
|
|
|
|
return text
|
|
|
|
|
|
def main():
|
|
args = parse_args()
|
|
|
|
if args.gpu_ids is not None:
|
|
os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu_ids
|
|
|
|
if args.hf_cache is not None:
|
|
os.environ["HF_HOME"] = args.hf_cache
|
|
os.environ["HF_DATASETS_CACHE"] = str(Path(args.hf_cache) / "datasets")
|
|
os.environ["HF_HUB_CACHE"] = str(Path(args.hf_cache) / "hub")
|
|
|
|
# Import after CUDA_VISIBLE_DEVICES is set.
|
|
from datasets import load_dataset
|
|
from transformers import AutoTokenizer
|
|
from vllm import LLM, SamplingParams
|
|
|
|
output_path = Path(args.output)
|
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
|
|
print("=" * 80)
|
|
print("AIME 2025 vLLM Evaluation")
|
|
print("=" * 80)
|
|
print(f"model: {args.model}")
|
|
print(f"dataset: {args.dataset}")
|
|
print(f"num_gpus / TP size: {args.num_gpus}")
|
|
print(f"CUDA_VISIBLE_DEVICES: {os.environ.get('CUDA_VISIBLE_DEVICES')}")
|
|
print(f"n_samples/problem: {args.n_samples}")
|
|
print(f"temperature: {args.temperature}")
|
|
print(f"top_p: {args.top_p}")
|
|
print(f"max_tokens: {args.max_tokens}")
|
|
print(f"prompt_template: {args.prompt_template}")
|
|
print(f"assistant_prefix: {repr(args.assistant_prefix)}")
|
|
print(f"output: {args.output}")
|
|
print("=" * 80)
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(
|
|
args.model,
|
|
trust_remote_code=args.trust_remote_code,
|
|
)
|
|
|
|
# Load HF dataset or local JSON/JSONL.
|
|
if args.dataset.endswith(".json") or args.dataset.endswith(".jsonl"):
|
|
dataset_dict = load_dataset("json", data_files=args.dataset)
|
|
else:
|
|
dataset_dict = load_dataset(args.dataset)
|
|
|
|
split = args.split or list(dataset_dict.keys())[0]
|
|
data = dataset_dict[split]
|
|
|
|
assert len(data) == 30, f"Expected 30 AIME 2025 problems, got {len(data)}"
|
|
|
|
print(f"Loaded AIME 2025 examples: {len(data)}")
|
|
for i in range(min(3, len(data))):
|
|
print(i, {k: data[i].get(k, None) for k in data.column_names[:6]})
|
|
|
|
if args.limit is not None:
|
|
data = data.select(range(min(args.limit, len(data))))
|
|
|
|
print(f"Loaded split: {split}")
|
|
print(f"Number of problems: {len(data)}")
|
|
print(f"First example keys: {list(data[0].keys())}")
|
|
|
|
prompts = []
|
|
examples = []
|
|
|
|
for idx, ex in enumerate(data):
|
|
problem = get_field(ex, ["problem", "question", "prompt"])
|
|
gold_raw = get_field(ex, ["answer", "final_answer", "target", "solution"])
|
|
gold = normalize_answer(gold_raw)
|
|
|
|
if args.prompt_template == "train":
|
|
# Raw ChatML prompt matching the OPD training parquet.
|
|
# Do NOT call apply_chat_template again, otherwise ChatML will be nested.
|
|
full_prompt = build_exact_training_prompt(problem)
|
|
elif args.prompt_template == "paper":
|
|
# Raw ChatML prompt matching the paper-style math evaluation prompt.
|
|
# Do NOT call apply_chat_template again.
|
|
full_prompt = build_paper_math_eval_prompt(problem)
|
|
elif args.prompt_template == "chat":
|
|
# Normal user-content prompt; tokenizer will add ChatML.
|
|
raw_prompt = build_problem_prompt(problem)
|
|
full_prompt = apply_chat_template(tokenizer, raw_prompt, args)
|
|
else:
|
|
raise ValueError(f"Unknown prompt_template: {args.prompt_template}")
|
|
|
|
# Important: when using raw ChatML prompts, apply_chat_template() is bypassed.
|
|
# Therefore --assistant-prefix must be appended here, not only inside apply_chat_template().
|
|
# This is useful for Qwen3 no-thinking style, e.g.:
|
|
# --assistant-prefix '<think>\\n\\n</think>\\n\\n'
|
|
if args.prompt_template in {"train", "paper"} and args.assistant_prefix is not None:
|
|
full_prompt += args.assistant_prefix.encode("utf-8").decode("unicode_escape")
|
|
|
|
prompts.append(full_prompt)
|
|
examples.append(
|
|
{
|
|
"idx": idx,
|
|
"problem": problem,
|
|
"gold_raw": gold_raw,
|
|
"gold": gold,
|
|
"prompt": full_prompt,
|
|
}
|
|
)
|
|
|
|
llm = LLM(
|
|
model=args.model,
|
|
tensor_parallel_size=args.num_gpus,
|
|
dtype=args.dtype,
|
|
trust_remote_code=args.trust_remote_code,
|
|
gpu_memory_utilization=args.gpu_memory_utilization,
|
|
seed=args.seed,
|
|
)
|
|
|
|
sampling_params = SamplingParams(
|
|
n=args.n_samples,
|
|
temperature=args.temperature,
|
|
top_p=args.top_p,
|
|
max_tokens=args.max_tokens,
|
|
stop=["<|im_end|>"],
|
|
)
|
|
|
|
outputs = llm.generate(prompts, sampling_params)
|
|
|
|
total = 0
|
|
correct = 0
|
|
per_problem_records = []
|
|
|
|
with output_path.open("w", encoding="utf-8") as f:
|
|
for ex, out in zip(examples, outputs):
|
|
sample_records = []
|
|
problem_correct = 0
|
|
|
|
for sample_id, completion in enumerate(out.outputs):
|
|
text = completion.text
|
|
pred = extract_answer(text)
|
|
is_correct = pred == ex["gold"]
|
|
|
|
total += 1
|
|
correct += int(is_correct)
|
|
problem_correct += int(is_correct)
|
|
|
|
record = {
|
|
"idx": ex["idx"],
|
|
"sample_id": sample_id,
|
|
"gold": ex["gold"],
|
|
"gold_raw": str(ex["gold_raw"]),
|
|
"pred": pred,
|
|
"correct": is_correct,
|
|
"completion": text,
|
|
"finish_reason": completion.finish_reason,
|
|
"stop_reason": getattr(completion, "stop_reason", None),
|
|
}
|
|
sample_records.append(record)
|
|
f.write(json.dumps(record, ensure_ascii=False) + "\n")
|
|
|
|
problem_acc = problem_correct / max(1, len(out.outputs))
|
|
per_problem_records.append(problem_acc)
|
|
|
|
if ex["idx"] < args.print_examples:
|
|
print("-" * 80)
|
|
print(f"Problem {ex['idx']}")
|
|
print(f"Gold: {ex['gold']} | Correct samples: {problem_correct}/{len(out.outputs)}")
|
|
print(f"First pred: {sample_records[0]['pred']}")
|
|
print(f"First completion preview:\n{sample_records[0]['completion'][:1000]}")
|
|
|
|
avg_pass1_micro = correct / total if total > 0 else 0.0
|
|
avg_pass1_macro = sum(per_problem_records) / len(per_problem_records) if per_problem_records else 0.0
|
|
|
|
print("=" * 80)
|
|
print("Final results")
|
|
print("=" * 80)
|
|
print(f"Problems: {len(examples)}")
|
|
print(f"Samples per problem: {args.n_samples}")
|
|
print(f"Total samples: {total}")
|
|
print(f"Correct samples: {correct}")
|
|
print(f"Average pass@1 micro: {avg_pass1_micro:.6f} ({100 * avg_pass1_micro:.2f}%)")
|
|
print(f"Average pass@1 macro: {avg_pass1_macro:.6f} ({100 * avg_pass1_macro:.2f}%)")
|
|
print(f"Saved generations to: {output_path}")
|
|
print("=" * 80)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main() |