securecoder-scripts / eval_securecoder.py
Taimwe's picture
Eval v6: restore _fetch_first_rows
f00a4df verified
Raw History Blame
21.1 kB
# /// script
# requires-python = ">=3.10"
# dependencies = [
# "datasets",
# "huggingface_hub",
# "torch",
# "transformers",
# "bitsandbytes",
# "accelerate",
# ]
# ///
"""Cheap evaluations for SecureCoder.
Three things are scored:
1. Tool-call validity - sample prompts with real schemas, render the model's
reply, parse the emitted <function=...><parameter=...> block back into
JSON, score parse OK / correct function name / schema-conformant args.
2. Code sanity - self-contained Python coding prompts; the model generates
code, we AST-parse + compile it (no execution).
3. Security knowledge - CyberSecurityEval MCQ when available, otherwise skip.
Writes results to --out-dir/report.json, prints a summary table, and uploads
the report to --upload-repo if set (default: the adapter repo itself).
"""
from __future__ import annotations
import argparse
import ast
import json
import logging
import os
import random
import re
import sys
import time
from typing import Any
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("eval")
CODE_PROMPTS: list[str] = [
"Write a Python function `def is_palindrome(s: str) -> bool:` that returns True iff s reads the same backwards ignoring case and non-alphanumeric characters.",
"Write a Python function `def merge_intervals(intervals: list[list[int]]) -> list[list[int]]:` that merges overlapping intervals.",
"Write a Python function `def two_sum(nums: list[int], target: int) -> list[int]:` returning the indices of two numbers that sum to target.",
"Write a Python function `def flatten(nested: list) -> list:` that flattens arbitrarily nested lists without recursion-limit errors.",
"Write a Python function `def parse_csv_line(line: str) -> list[str]:` that handles quoted fields and escaped quotes per RFC 4180.",
"Write a Python function `def lru_cache(k: int):` returning a decorator that keeps at most k most-recently-used call results.",
"Write a Python function `def is_anagram(a: str, b: str) -> bool:` ignoring spaces, punctuation, and case.",
"Write a Python function `def topological_order(graph: dict[str, list[str]]) -> list[str] | None:` returning a valid order or None on cycle.",
"Write a Python function `def tokenise(s: str) -> list[str]:` for a simple expression language with integers, +, -, *, / and parentheses.",
"Write a Python function `def slugify(text: str) -> str:` producing a URL-safe ASCII slug from any unicode text.",
"Write a Python function `def read_jsonl(path: str) -> list[dict]:` streaming a JSONL file one record at a time without loading the whole file.",
"Write a Python function `def binary_search(arr: list[int], target: int) -> int:` returning the index of target or -1.",
"Write a Python function `def unique_in_order(s: str) -> list[str]:` preserving order while deduplicating adjacent equal characters.",
"Write a Python function `def safe_eval(expr: str) -> int:` evaluating an integer expression with + - * / parens, no eval(), no imports.",
"Write a Python function `def dedupe_preserve_order(items: list) -> list:` returning the input without duplicates, original order.",
"Write a Python function `def word_frequency(text: str) -> dict[str, int]:` counting occurrences after normalising case and stripping punctuation.",
"Write a Python function `def matrix_multiply(a: list[list[float]], b: list[list[float]]) -> list[list[float]]:` for any compatible shapes.",
"Write a Python function `def roman_to_int(s: str) -> int:` handling subtractive notation (IV, IX, XL, XC, CD, CM) up to 3999.",
"Write a Python function `def fibonacci(n: int) -> int:` returning the n-th Fibonacci number with O(n) time and O(1) space.",
"Write a Python function `def validate_ipv4(s: str) -> bool:` accepting only dotted-quad strings with each octet in 0..255.",
]
TOOL_FN_PAT = re.compile(r"<function=([A-Za-z0-9_\.]+)>", re.S)
TOOL_PARAM_PAT = re.compile(r"<parameter=([A-Za-z0-9_]+)>\s*(.*?)\s*</parameter>", re.S)
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="SecureCoder evaluations")
p.add_argument("--adapter", default="Taimwe/securecoder-30b-pro", help="(unused when --merged-repo is set)")
p.add_argument("--base", default="unsloth/Qwen3-Coder-30B-A3B-Instruct", help="(unused when --merged-repo is set)")
p.add_argument("--merged-repo", default="Taimwe/securecoder-30b-pro-merged", help="merged safetensors repo")
p.add_argument("--out-dir", default="/data/eval-out")
p.add_argument("--tool-prompts", type=int, default=80)
p.add_argument("--code-prompts", type=int, default=15)
p.add_argument("--upload-repo", default="Taimwe/securecoder-30b-pro")
p.add_argument("--max-new-tokens", type=int, default=384)
p.add_argument("--seed", type=int, default=3407)
return p.parse_args()
_TYPE_ALIASES = {
"str": "string", "string": "string", "text": "string",
"int": "integer", "integer": "integer", "long": "integer",
"float": "number", "double": "number", "number": "number",
"bool": "boolean", "boolean": "boolean",
"list": "array", "array": "array", "dict": "object", "object": "object",
}
def _normalise_tool_schema(tool):
"""Accept tool specs as a dict, a list of dicts, or a JSON string. Hermes
ships ``tools`` as a JSON string; xLAM and most others use a list/dict."""
if tool is None:
return None
if isinstance(tool, str):
try:
tool = json.loads(tool)
except (json.JSONDecodeError, TypeError):
return None
if isinstance(tool, list):
for entry in tool:
norm = _normalise_tool_schema(entry)
if norm:
return norm
return None
if not isinstance(tool, dict):
return None
fn = tool.get("function") if isinstance(tool.get("function"), dict) else tool
if not isinstance(fn, dict) or not fn.get("name"):
return None
params = fn.get("parameters") or {}
if isinstance(params, str):
try:
params = json.loads(params)
except (json.JSONDecodeError, TypeError):
params = {}
if isinstance(params, dict) and params and "properties" not in params:
properties: dict[str, Any] = {}
required: list[str] = []
for name, spec in params.items():
if isinstance(spec, dict):
cleaned = {k: v for k, v in spec.items()
if k in ("type", "description", "enum", "default", "title", "items")}
cleaned["type"] = _TYPE_ALIASES.get(str(cleaned.get("type", "")).lower(), "string")
properties[name] = cleaned
if "default" not in spec:
required.append(name)
else:
properties[name] = {"type": "string"}
required.append(name)
params = {"type": "object", "properties": properties}
if required:
params["required"] = required
return {"name": fn["name"], "description": fn.get("description", ""), "parameters": params}
def _messages_from_any(row, kind):
"""Extract a single user turn from a Hermes ``conversations`` row (from/value)
or an OpenAI-style ``messages`` row (role/content). Tools can be a list, dict,
or JSON string - we keep the first valid schema."""
tools_raw = row.get("tools")
if isinstance(tools_raw, str):
try:
tools_raw = json.loads(tools_raw)
except (json.JSONDecodeError, TypeError):
tools_raw = None
tools: list[dict] = []
candidates = tools_raw if isinstance(tools_raw, list) else ([tools_raw] if isinstance(tools_raw, dict) else [])
for c in candidates:
n = _normalise_tool_schema(c)
if n:
tools.append(n)
messages: list[dict] = []
convo = row.get("messages") or row.get("conversations") or []
if isinstance(convo, str):
try:
convo = json.loads(convo)
except (json.JSONDecodeError, TypeError):
convo = []
ROLE_MAP = {"human": "user", "system": "system", "gpt": "assistant",
"user": "user", "assistant": "assistant"}
if isinstance(convo, list):
for turn in convo:
if not isinstance(turn, dict):
continue
role_raw = str(turn.get("from") or turn.get("role") or "").lower()
role = ROLE_MAP.get(role_raw)
if role == "user":
content = turn.get("value", turn.get("content", ""))
messages.append({"role": "user", "content": str(content)})
break
return messages, tools
def _fetch_first_rows(repo: str, config: str | None, split: str, n: int) -> list[dict]:
"""Stream up to ``n`` rows from a Hub dataset. ``config`` is the dataset config
name (e.g. ``func_calling`` for Hermes FC); pass ``None`` for datasets that
have only the default config."""
from datasets import load_dataset
kwargs: dict[str, Any] = {"split": split, "streaming": True}
if config:
kwargs["name"] = config
ds = load_dataset(repo, token=os.environ.get("HF_TOKEN"), **kwargs)
out = []
for row in ds:
out.append(dict(row))
if len(out) >= n:
break
return out
def _build_tool_prompts(rows: list[dict]) -> list[dict]:
"""Build (prompt, tools) tuples via the shared _messages_from_any helper.
Hermes rows use ``conversations`` (from/value); OpenAI-style use ``messages``
(role/content). Both are accepted transparently."""
prompts = []
for row in rows:
messages, tools = _messages_from_any(row, "auto")
if not tools or not messages:
continue
user = next((m["content"] for m in messages if m["role"] == "user"), None)
if not user:
continue
prompts.append({"prompt": str(user)[:1200], "tools": tools,
"expected_call": True})
if len(prompts) >= 200:
break
return prompts
def _render_prompt(tokenizer, prompt: str, tools: list[dict]) -> str:
return tokenizer.apply_chat_template(
[{"role": "user", "content": prompt}],
tools=tools, tokenize=False, add_generation_prompt=True,
)
def _parse_emitted_calls(text: str) -> list[dict]:
fns = TOOL_FN_PAT.findall(text)
if not fns:
return []
calls = []
for fn in fns:
start = text.find(f"<function={fn}>")
if start < 0:
continue
end = text.find("</function>", start)
block = text[start:end if end > 0 else start + 4000]
params = TOOL_PARAM_PAT.findall(block)
calls.append({"name": fn, "arguments": dict(params)})
return calls
def _score_call(call: dict, tools: list[dict]) -> dict:
name = call.get("name")
fn = next((t for t in tools if t.get("name") == name), None)
if not fn:
return {"parse": True, "name_ok": False, "schema_ok": False}
props = fn.get("parameters", {}).get("properties", {}) or {}
expected = set(props.keys())
given = set((call.get("arguments") or {}).keys())
return {"parse": True, "name_ok": True,
"schema_ok": expected.issubset(given) if expected else True,
"expected_keys": sorted(expected), "given_keys": sorted(given)}
def eval_tool_calls(tokenizer, model, args) -> dict:
import torch
log.info("tool-call eval: streaming candidates from hermes FC ...")
rows = _fetch_first_rows("NousResearch/hermes-function-calling-v1", "func_calling", "train",
args.tool_prompts * 4)
prompts = _build_tool_prompts(rows)
if args.tool_prompts:
prompts = prompts[: args.tool_prompts]
log.info("tool-call eval: %d usable prompts", len(prompts))
out = []
for i, p in enumerate(prompts):
try:
text = _render_prompt(tokenizer, p["prompt"], p["tools"])
ids = tokenizer(text, return_tensors="pt", add_special_tokens=False).input_ids.to(model.device)
with torch.no_grad():
generated = model.generate(ids, max_new_tokens=args.max_new_tokens, do_sample=False, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.eos_token_id)
reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True)
except Exception as exc:
out.append({"prompt": p["prompt"][:60], "error": repr(exc)[:120]})
continue
calls = _parse_emitted_calls(reply)
scored = [_score_call(c, p["tools"]) for c in calls]
out.append({"prompt": p["prompt"][:80], "reply_first_160": reply[:160],
"expected_call": p["expected_call"], "n_calls": len(calls),
"calls": calls, "scores": scored})
if (i + 1) % 25 == 0:
log.info(" tool-call progress: %d/%d", i + 1, len(prompts))
n = len(out)
parse_ok = sum(1 for r in out if r.get("scores") and any(s["parse"] for s in r["scores"]))
name_ok = sum(1 for r in out if r.get("scores") and any(s["name_ok"] for s in r["scores"]))
schema_ok = sum(1 for r in out if r.get("scores") and any(s["schema_ok"] for s in r["scores"]))
return {"section": "tool_calls", "n_prompts": n,
"parse_rate": parse_ok / max(n, 1), "name_rate": name_ok / max(n, 1),
"schema_rate": schema_ok / max(n, 1), "details": out}
def eval_code_sanity(tokenizer, model, args) -> dict:
import torch
out = []
prompts = CODE_PROMPTS[: args.code_prompts]
for i, prompt in enumerate(prompts):
text = tokenizer.apply_chat_template(
[{"role": "user", "content": prompt}],
tokenize=False, add_generation_prompt=True,
)
ids = tokenizer(text, return_tensors="pt", add_special_tokens=False).input_ids.to(model.device)
try:
with torch.no_grad():
generated = model.generate(ids, max_new_tokens=384, do_sample=False, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.eos_token_id)
reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True)
except Exception as exc:
out.append({"prompt": prompt[:60], "error": repr(exc)[:120]})
continue
block = None
# Truncation inside a code block is common (max_new_tokens cuts
# mid-line), so we keep the LONGEST ```python ... ``` block in the
# reply rather than the first one. Most failures in the first eval
# run were valid code cut mid-comment, not bad code.
blocks = re.findall(r"```(?:python)?\s*\n(.*?)```", reply, re.S)
if blocks:
block = max(blocks, key=len)
else:
start = reply.find("def ")
if start >= 0:
block = reply[start:]
if block:
# Strip anything past a closing fence on the same line
block = re.split(r"\n```", block, maxsplit=1)[0]
parsed = compiles = None
if block:
try:
ast.parse(block)
parsed = True
except SyntaxError:
parsed = False
if block and parsed:
try:
compile(block, "<eval>", "exec")
compiles = True
except Exception:
compiles = False
if (i + 1) % 5 == 0:
log.info(" code sanity: %d/%d", i + 1, len(prompts))
n = len(out)
ast_ok = sum(1 for r in out if r.get("ast_ok"))
compile_ok = sum(1 for r in out if r.get("compile_ok"))
return {"section": "code_sanity", "n_prompts": n,
"ast_rate": ast_ok / max(n, 1), "compile_rate": compile_ok / max(n, 1),
"details": out}
def eval_security_mcq(tokenizer, model, n_questions: int = 25) -> dict:
import torch
try:
rows = _fetch_first_rows("CyberNative/CyberSecurityEval", None, "train", n_questions * 2)
except Exception as exc:
return {"section": "security_mcq", "error": repr(exc)[:200], "skipped": True}
rows = rows[:n_questions]
if not rows:
return {"section": "security_mcq", "skipped": True, "reason": "no rows"}
correct = 0; details = []
for r in rows:
question = r.get("question") or r.get("prompt") or r.get("input")
options = r.get("options") or r.get("choices") or r.get("answers")
answer = r.get("answer") or r.get("label")
if not question or not options or answer is None:
continue
if isinstance(options, dict):
opts = "\n".join(f"{k}. {v}" for k, v in options.items())
key_map = {str(k): v for k, v in options.items()}
else:
opts = "\n".join(f"{i}. {o}" for i, o in enumerate(options))
key_map = {str(i): options[i]}
user = f"Question: {question}\n\n{opts}\n\nRespond with the letter of the correct answer only."
text = tokenizer.apply_chat_template(
[{"role": "user", "content": user}], tokenize=False, add_generation_prompt=True)
ids = tokenizer(text, return_tensors="pt", add_special_tokens=False).input_ids.to(model.device)
try:
with torch.no_grad():
generated = model.generate(ids, max_new_tokens=8, do_sample=False, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.eos_token_id)
reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True).strip()
except Exception:
continue
first_letter = reply[:1].upper()
predicted = key_map.get(first_letter)
is_correct = predicted == answer
correct += int(is_correct)
details.append({"question": str(question)[:80], "reply": reply[:10], "ok": is_correct})
return {"section": "security_mcq", "n_questions": len(details),
"accuracy": correct / max(len(details), 1), "details": details}
def main() -> int:
args = parse_args()
token = os.environ.get("HF_TOKEN")
if not token:
log.error("HF_TOKEN not set"); return 1
os.makedirs(args.out_dir, exist_ok=True)
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
random.seed(args.seed)
log.info("loading merged model %s in 4-bit ...", args.merged_repo)
tokenizer = AutoTokenizer.from_pretrained(args.merged_repo, token=token)
bnb = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
model = AutoModelForCausalLM.from_pretrained(
args.merged_repo,
token=token,
quantization_config=bnb,
dtype=torch.bfloat16,
device_map="auto",
)
log.info("model loaded")
sections = [
eval_tool_calls(tokenizer, model, args),
eval_code_sanity(tokenizer, model, args),
eval_security_mcq(tokenizer, model),
]
summary = {
"adapter": args.adapter,
"base": args.base,
"when": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
"sections": [
{"section": s["section"], **{k: v for k, v in s.items() if k not in {"section", "details"}}}
for s in sections
],
"raw": sections,
}
out_json = os.path.join(args.out_dir, "report.json")
with open(out_json, "w", encoding="utf-8") as fh:
json.dump(summary, fh, indent=2, default=str)
log.info("report written: %s", out_json)
print("\n" + "=" * 70)
print(f"{'section':<18}{'metric':<22}{'value':>10}")
print("-" * 70)
for s in sections:
for k, v in s.items():
if isinstance(v, float) and k.endswith(("rate", "accuracy")):
print(f"{s['section']:<18}{k:<22}{v*100:>9.1f}%")
if "n_prompts" in s:
print(f"{s['section']:<18}{'n_prompts':<22}{s['n_prompts']:>10}")
print("=" * 70)
if args.upload_repo:
from huggingface_hub import HfApi
api = HfApi(token=token)
api.create_repo(args.upload_repo, repo_type="model", exist_ok=True, private=True)
api.upload_file(path_or_fileobj=out_json, path_in_repo="eval-report.json",
repo_id=args.upload_repo, repo_type="model",
commit_message="Add evaluation report")
log.info("report pushed to https://huggingface.co/%s", args.upload_repo)
return 0
if __name__ == "__main__":
raise SystemExit(main())