dna-diskchat-2b-peer-v25 / scripts /v25_eval /generate_evalplus.py
jaivial's picture
Upload folder using huggingface_hub
580cb69 verified
Raw
History Blame Contribute Delete
6.77 kB
#!/usr/bin/env python3
"""Generate EvalPlus-compatible HumanEval(+)/MBPP(+) samples."""
from __future__ import annotations
import argparse
import importlib
import json
import re
import sys
from pathlib import Path
from typing import Any
def arguments() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--dataset", choices=("humaneval", "mbpp"), required=True)
parser.add_argument("--checkpoint", type=Path, required=True)
parser.add_argument("--tokenizer", type=Path, required=True)
parser.add_argument("--model-path", type=Path, required=True)
parser.add_argument("--model-module", required=True)
parser.add_argument("--model-class", default="DnaPeerV21")
parser.add_argument("--adapter-module", help="Optional module whose load(args) returns an object with generate(prompt, max_new_tokens)")
parser.add_argument("--out", type=Path, required=True)
parser.add_argument("--device", default="cuda")
parser.add_argument("--max-new-tokens", type=int, default=512)
parser.add_argument("--limit", type=int, help="Smoke testing only; never report limited runs as full benchmark")
return parser.parse_args()
def strip_fence(text: str) -> str:
text = text.strip()
if "```" not in text:
return text
parts = text.split("```")
fenced = parts[1] if len(parts) >= 3 else parts[-1]
if fenced.lstrip().startswith("python"):
fenced = fenced.lstrip()[6:]
return fenced.strip()
class BuiltinRecurrentAdapter:
"""Adapter for v23/v24-style DnaPeer recurrent blocks.
v25 model changes should use --adapter-module instead of editing evaluator code.
"""
def __init__(self, args: argparse.Namespace):
import torch
from tokenizers import Tokenizer, decoders
sys.path.insert(0, str(args.model_path))
module = importlib.import_module(args.model_module)
model_type = getattr(module, args.model_class)
checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
self.model = model_type(**checkpoint["config"]).to(args.device)
self.model.load_state_dict({key.replace("_orig_mod.", ""): value for key, value in checkpoint["model"].items()}, strict=True)
self.model.eval()
self.tokenizer = Tokenizer.from_file(str(args.tokenizer))
self.tokenizer.decoder = decoders.ByteLevel()
self.bos = self.tokenizer.token_to_id("<s>")
self.eos = self.tokenizer.token_to_id("</s>")
self.user = self.tokenizer.token_to_id("<|user|>")
self.assistant = self.tokenizer.token_to_id("<|assistant|>")
self.device, self.torch = args.device, torch
if None in (self.bos, self.eos, self.user, self.assistant):
raise ValueError("tokenizer must contain <s>, </s>, <|user|>, and <|assistant|>")
def generate(self, prompt: str, max_new_tokens: int) -> str:
torch = self.torch
model = self.model
states = [torch.zeros(1, model.d, device=self.device) for _ in model.blocks]
def step(token: int):
x = model.embed(torch.tensor([token], device=self.device))
for index, block in enumerate(model.blocks):
key, value, receptance, gate = block.proj(block.n1(x)).chunk(4, -1)
gate = torch.sigmoid(gate + block.decay)
states[index] = gate * states[index] + (1 - gate) * torch.tanh(key)
x = x + block.o(torch.sigmoid(receptance) * states[index] * torch.sigmoid(value))
x = x + block.ffn(block.n2(x))
memory, _, _ = model.route(x)
return model.norm(x + memory)
prefix = [self.bos, self.user] + self.tokenizer.encode("\n" + prompt).ids + [self.eos, self.assistant]
with torch.no_grad():
feature = None
for token in prefix:
feature = step(token)
generated = []
for _ in range(max_new_tokens):
token = int(torch.nn.functional.linear(feature, model.embed.weight)[0].argmax())
if token == self.eos:
break
generated.append(token)
text = self.tokenizer.decode(generated)
markers = ("\n<|user|>", "\nif __name__ ==", "\n# End")
hits = [text.find(marker) for marker in markers if marker in text]
if hits:
generated = self.tokenizer.encode(text[: min(hits)]).ids
break
feature = step(token)
return self.tokenizer.decode(generated)
def main() -> None:
args = arguments()
try:
from evalplus.data import get_human_eval_plus, get_mbpp_plus
except ImportError as exc:
raise SystemExit("install pinned EvalPlus first: pip install 'evalplus==0.3.1'") from exc
problems: dict[str, dict[str, Any]] = (get_human_eval_plus() if args.dataset == "humaneval" else get_mbpp_plus())
if args.adapter_module:
sys.path.insert(0, str(args.model_path))
adapter = importlib.import_module(args.adapter_module).load(args)
else:
adapter = BuiltinRecurrentAdapter(args)
args.out.parent.mkdir(parents=True, exist_ok=True)
selected = list(problems.items())[: args.limit]
with args.out.open("w", encoding="utf-8", buffering=1) as output:
for index, (task_id, problem) in enumerate(selected, 1):
prompt = problem["prompt"]
instruction = "Complete this Python program. Return only valid Python code, without Markdown fences.\n\n" + prompt
completion = strip_fence(adapter.generate(instruction, args.max_new_tokens))
entry_point = problem.get("entry_point")
definition = rf"\bdef\s+{re.escape(entry_point)}\s*\(" if entry_point else None
if definition and re.search(definition, completion):
solution = completion
elif definition and re.search(definition, prompt):
solution = prompt + completion
else:
solution = completion
output.write(json.dumps({"task_id": task_id, "solution": solution}) + "\n")
print(f"GENERATED {index}/{len(selected)} {task_id}", flush=True)
metadata = {"dataset": args.dataset, "samples": len(selected), "full_dataset": args.limit is None,
"generation": "greedy", "max_new_tokens": args.max_new_tokens,
"checkpoint": str(args.checkpoint), "model_module": args.model_module,
"adapter_module": args.adapter_module}
args.out.with_suffix(args.out.suffix + ".meta.json").write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8")
if __name__ == "__main__":
main()