Download evaluate_iapo_gsm8k.py from Srishti280992/iapo-gsm8k-eval-job: direct link, hf CLI and curl.
- Browser
- Download file 11.3 kB
-
https://huggingface.co/Srishti280992/iapo-gsm8k-eval-job/resolve/main/evaluate_iapo_gsm8k.py
- Command line
-
hf download hf://Srishti280992/iapo-gsm8k-eval-job/evaluate_iapo_gsm8k.py
-
curl -L -o evaluate_iapo_gsm8k.py https://huggingface.co/Srishti280992/iapo-gsm8k-eval-job/resolve/main/evaluate_iapo_gsm8k.py
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() | |