| |
| """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() |
|
|