iapo-gsm8k-eval-job / evaluate_iapo_gsm8k.py
Srishti280992's picture
Fix decoder-only padding and completion slicing
e5f2a36 verified
Raw History Blame Contribute Delete
11.3 kB
#!/usr/bin/env python3
# /// script
# requires-python = ">=3.10"
# dependencies = [
# "accelerate>=0.33",
# "datasets>=2.20",
# "huggingface_hub>=0.24",
# "safetensors>=0.4",
# "sentencepiece>=0.2",
# "torch>=2.4",
# "transformers>=4.44",
# ]
# ///
"""Scaled GSM8K checkpoint evaluation for IAPO.
This evaluates released merged checkpoints against the base Qwen models using the
IAPO paper's chat prompt style and reports Pass@k, generated-token length, ratio,
and wall-clock generation time.
"""
from __future__ import annotations
import argparse
import csv
import json
import math
import os
import random
import re
import time
from pathlib import Path
from typing import Any
import torch
from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer
SYSTEM_PROMPT = """A conversation between User and Assistant. The user asks a question, and the Assistant solves it.
The assistant first thinks about the reasoning process in the mind and then provides the user
with the answer. The reasoning process and answer are enclosed within <think> and
<answer> tags, respectively, i.e., <think> reasoning process here </think>
<answer> answer here </answer>.
The answer must be a single integer."""
FEWSHOT_USER = "What is 2+2?"
FEWSHOT_ASSISTANT = "<think>To calculate 2+2, we simply add the numbers together: 2 + 2 = 4.</think>\n<answer>4</answer>"
MODEL_SPECS = {
"base-0.5b": {"repo": "Qwen/Qwen2.5-0.5B-Instruct", "subfolder": None},
"iapo-0.5b-gsm8k": {"repo": "jonathanhe123/iapo", "subfolder": "Qwen2.5-0.5B-Instruct_GSM8K"},
"base-7b": {"repo": "Qwen/Qwen2.5-7B-Instruct", "subfolder": None},
"iapo-7b-gsm8k": {"repo": "jonathanhe123/iapo", "subfolder": "Qwen2.5-7B-Instruct_GSM8K"},
}
def normalize_answer(s: str | None) -> str:
if s is None:
return ""
s = str(s).strip()
s = s.split("=")[-1]
s = s.replace(",", "").replace("$", "").strip()
s = re.sub(r"\\boxed\{([^{}]+)\}", r"\1", s)
return s.strip().rstrip(".")
def extract_gold(answer: str) -> str:
return normalize_answer(answer.split("####")[-1])
def extract_pred(text: str) -> str:
xml = re.search(r"<answer>(.*?)</answer>", text, flags=re.DOTALL | re.IGNORECASE)
if xml:
return normalize_answer(xml.group(1))
answer_line = re.findall(r"Answer:\s*([-+]?\d[\d,]*(?:\.\d+)?)", text, flags=re.IGNORECASE)
if answer_line:
return normalize_answer(answer_line[-1])
nums = re.findall(r"[-+]?\d[\d,]*(?:\.\d+)?", text)
return normalize_answer(nums[-1]) if nums else ""
def make_prompt(tokenizer: Any, question: str) -> str:
messages = [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": FEWSHOT_USER},
{"role": "assistant", "content": FEWSHOT_ASSISTANT},
{"role": "user", "content": question},
]
return tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
def load_model(spec: dict[str, str | None], dtype: str):
torch_dtype = {
"auto": "auto",
"float16": torch.float16,
"bfloat16": torch.bfloat16,
"float32": torch.float32,
}[dtype]
kwargs = {"trust_remote_code": True}
if spec["subfolder"]:
kwargs["subfolder"] = spec["subfolder"]
tokenizer = AutoTokenizer.from_pretrained(spec["repo"], **kwargs)
if tokenizer.pad_token_id is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"
model = AutoModelForCausalLM.from_pretrained(
spec["repo"],
torch_dtype=torch_dtype,
device_map="auto",
low_cpu_mem_usage=True,
**kwargs,
)
model.eval()
return tokenizer, model
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--models", nargs="+", default=["base-0.5b", "iapo-0.5b-gsm8k"])
parser.add_argument("--sample-size", type=int, default=32)
parser.add_argument("--num-return-sequences", type=int, default=4)
parser.add_argument("--max-new-tokens", type=int, default=384)
parser.add_argument("--batch-size", type=int, default=4)
parser.add_argument("--temperature", type=float, default=1.0)
parser.add_argument("--top-p", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--dtype", choices=["auto", "float16", "bfloat16", "float32"], default="bfloat16")
parser.add_argument("--output-dir", default="outputs/hf_gsm8k_eval")
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
random.seed(args.seed)
torch.manual_seed(args.seed)
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
print(json.dumps({
"event": "env",
"torch": torch.__version__,
"cuda_available": torch.cuda.is_available(),
"cuda_device_count": torch.cuda.device_count(),
"models": args.models,
"sample_size": args.sample_size,
"num_return_sequences": args.num_return_sequences,
}))
if args.dry_run:
print(json.dumps({"event": "dry_run_ok"}))
return
dataset = load_dataset("openai/gsm8k", "main", split="test")
indices = list(range(len(dataset)))
random.Random(args.seed).shuffle(indices)
indices = sorted(indices[: args.sample_size])
examples = [dataset[i] for i in indices]
golds = [extract_gold(ex["answer"]) for ex in examples]
all_summaries = []
details_path = out_dir / "generations.jsonl"
with details_path.open("w") as details:
for model_key in args.models:
if model_key not in MODEL_SPECS:
raise ValueError(f"Unknown model key: {model_key}")
spec = MODEL_SPECS[model_key]
tokenizer, model = load_model(spec, args.dtype)
prompts = [make_prompt(tokenizer, ex["question"]) for ex in examples]
per_question_correct: list[list[int]] = []
per_question_lengths: list[list[int]] = []
generated_count = 0
gen_seconds = 0.0
for start in range(0, len(prompts), args.batch_size):
batch_prompts = prompts[start : start + args.batch_size]
encoded = tokenizer(batch_prompts, return_tensors="pt", padding=True).to(model.device)
do_sample = args.num_return_sequences > 1 or args.temperature > 0
t0 = time.perf_counter()
with torch.inference_mode():
output_ids = model.generate(
**encoded,
max_new_tokens=args.max_new_tokens,
do_sample=do_sample,
temperature=args.temperature if do_sample else None,
top_p=args.top_p if do_sample else None,
num_return_sequences=args.num_return_sequences,
pad_token_id=tokenizer.pad_token_id,
eos_token_id=tokenizer.eos_token_id,
)
gen_seconds += time.perf_counter() - t0
prompt_lens = encoded["attention_mask"].sum(dim=1).tolist()
input_width = int(encoded["input_ids"].shape[1])
expanded_prompt_lens = [
input_width
for i in range(len(prompt_lens))
for _ in range(args.num_return_sequences)
]
texts = []
lengths = []
for ids, plen in zip(output_ids, expanded_prompt_lens):
new_ids = ids[int(plen) :]
if tokenizer.eos_token_id in new_ids:
eos_pos = (new_ids == tokenizer.eos_token_id).nonzero(as_tuple=True)[0]
if len(eos_pos):
new_ids = new_ids[: int(eos_pos[0]) + 1]
lengths.append(int(len(new_ids)))
texts.append(tokenizer.decode(new_ids, skip_special_tokens=True))
for local_idx in range(len(batch_prompts)):
global_idx = start + local_idx
group_texts = texts[
local_idx * args.num_return_sequences : (local_idx + 1) * args.num_return_sequences
]
group_lengths = lengths[
local_idx * args.num_return_sequences : (local_idx + 1) * args.num_return_sequences
]
preds = [extract_pred(t) for t in group_texts]
correct = [int(pred == golds[global_idx]) for pred in preds]
per_question_correct.append(correct)
per_question_lengths.append(group_lengths)
generated_count += len(group_texts)
for j, (text, pred, ok, length) in enumerate(zip(group_texts, preds, correct, group_lengths)):
details.write(json.dumps({
"model": model_key,
"dataset_index": indices[global_idx],
"sample_in_group": j,
"gold": golds[global_idx],
"prediction": pred,
"correct": ok,
"completion_tokens": length,
"completion": text,
}) + "\n")
details.flush()
k_values = [k for k in [1, 2, 4, 8, 16, 32] if k <= args.num_return_sequences]
summary = {
"model": model_key,
"repo": spec["repo"],
"subfolder": spec["subfolder"],
"sample_size": len(examples),
"num_return_sequences": args.num_return_sequences,
"max_new_tokens": args.max_new_tokens,
"generated_count": generated_count,
"generation_seconds": gen_seconds,
"avg_seconds_per_completion": gen_seconds / max(1, generated_count),
}
for k in k_values:
pass_at_k = sum(int(any(row[:k])) for row in per_question_correct) / len(per_question_correct)
length_at_k = sum(sum(row[:k]) for row in per_question_lengths) / (len(per_question_lengths) * k)
summary[f"pass@{k}"] = pass_at_k
summary[f"length@{k}"] = length_at_k
summary[f"ratio@{k}"] = pass_at_k / length_at_k if length_at_k else math.nan
all_summaries.append(summary)
print(json.dumps({"event": "summary", **summary}))
del model
if torch.cuda.is_available():
torch.cuda.empty_cache()
summary_path = out_dir / "summary.json"
csv_path = out_dir / "summary.csv"
summary_path.write_text(json.dumps(all_summaries, indent=2))
keys = sorted({k for row in all_summaries for k in row})
with csv_path.open("w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=keys)
writer.writeheader()
writer.writerows(all_summaries)
print(json.dumps({"event": "wrote_outputs", "summary": str(summary_path), "csv": str(csv_path), "details": str(details_path)}))
if __name__ == "__main__":
main()