Taimwe commited on
Commit
a2ba1ae
·
verified ·
1 Parent(s): d9dfe3c

Add eval script (tool-call validity, code sanity, security MCQ)

Browse files
Files changed (1) hide show
  1. eval_securecoder.py +384 -0
eval_securecoder.py ADDED
@@ -0,0 +1,384 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # requires-python = ">=3.10"
3
+ # dependencies = [
4
+ # "datasets",
5
+ # "huggingface_hub",
6
+ # ]
7
+ # ///
8
+ """Cheap evaluations for SecureCoder.
9
+
10
+ Three things are scored:
11
+
12
+ 1. Tool-call validity - sample prompts with real schemas, render the model's
13
+ reply, parse the emitted <function=...><parameter=...> block back into
14
+ JSON, score parse OK / correct function name / schema-conformant args.
15
+
16
+ 2. Code sanity - self-contained Python coding prompts; the model generates
17
+ code, we AST-parse + compile it (no execution).
18
+
19
+ 3. Security knowledge - CyberSecurityEval MCQ when available, otherwise skip.
20
+
21
+ The script writes results to --out-dir/report.json, prints a summary table, and
22
+ uploads the report to --upload-repo if set (default: the adapter repo itself).
23
+ """
24
+
25
+ from __future__ import annotations
26
+
27
+ import argparse
28
+ import ast
29
+ import json
30
+ import logging
31
+ import os
32
+ import random
33
+ import re
34
+ import sys
35
+ import time
36
+ from typing import Any
37
+
38
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
39
+ log = logging.getLogger("eval")
40
+
41
+ CODE_PROMPTS: list[str] = [
42
+ "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.",
43
+ "Write a Python function `def merge_intervals(intervals: list[list[int]]) -> list[list[int]]:` that merges overlapping intervals.",
44
+ "Write a Python function `def two_sum(nums: list[int], target: int) -> list[int]:` returning the indices of two numbers that sum to target.",
45
+ "Write a Python function `def flatten(nested: list) -> list:` that flattens arbitrarily nested lists without recursion-limit errors.",
46
+ "Write a Python function `def parse_csv_line(line: str) -> list[str]:` that handles quoted fields and escaped quotes per RFC 4180.",
47
+ "Write a Python function `def lru_cache(k: int):` returning a decorator that keeps at most k most-recently-used call results.",
48
+ "Write a Python function `def is_anagram(a: str, b: str) -> bool:` ignoring spaces, punctuation, and case.",
49
+ "Write a Python function `def topological_order(graph: dict[str, list[str]]) -> list[str] | None:` returning a valid order or None on cycle.",
50
+ "Write a Python function `def tokenise(s: str) -> list[str]:` for a simple expression language with integers, +, -, *, / and parentheses.",
51
+ "Write a Python function `def slugify(text: str) -> str:` producing a URL-safe ASCII slug from any unicode text.",
52
+ "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.",
53
+ "Write a Python function `def binary_search(arr: list[int], target: int) -> int:` returning the index of target or -1.",
54
+ "Write a Python function `def unique_in_order(s: str) -> list[str]:` preserving order while deduplicating adjacent equal characters.",
55
+ "Write a Python function `def safe_eval(expr: str) -> int:` evaluating an integer expression with + - * / parens, no eval(), no imports.",
56
+ "Write a Python function `def dedupe_preserve_order(items: list) -> list:` returning the input without duplicates, original order.",
57
+ "Write a Python function `def word_frequency(text: str) -> dict[str, int]:` counting occurrences after normalising case and stripping punctuation.",
58
+ "Write a Python function `def matrix_multiply(a: list[list[float]], b: list[list[float]]) -> list[list[float]]:` for any compatible shapes.",
59
+ "Write a Python function `def roman_to_int(s: str) -> int:` handling subtractive notation (IV, IX, XL, XC, CD, CM) up to 3999.",
60
+ "Write a Python function `def fibonacci(n: int) -> int:` returning the n-th Fibonacci number with O(n) time and O(1) space.",
61
+ "Write a Python function `def validate_ipv4(s: str) -> bool:` accepting only dotted-quad strings with each octet in 0..255.",
62
+ ]
63
+
64
+ TOOL_FN_PAT = re.compile(r"<function=([A-Za-z0-9_\.]+)>", re.S)
65
+ TOOL_PARAM_PAT = re.compile(r"<parameter=([A-Za-z0-9_]+)>\s*(.*?)\s*</parameter>", re.S)
66
+
67
+
68
+ def parse_args() -> argparse.Namespace:
69
+ p = argparse.ArgumentParser(description="SecureCoder evaluations")
70
+ p.add_argument("--adapter", default="Taimwe/securecoder-30b-pro")
71
+ p.add_argument("--base", default="unsloth/Qwen3-Coder-30B-A3B-Instruct")
72
+ p.add_argument("--out-dir", default="/data/eval-out")
73
+ p.add_argument("--tool-prompts", type=int, default=80)
74
+ p.add_argument("--code-prompts", type=int, default=15)
75
+ p.add_argument("--upload-repo", default="Taimwe/securecoder-30b-pro")
76
+ p.add_argument("--max-new-tokens", type=int, default=384)
77
+ p.add_argument("--seed", type=int, default=3407)
78
+ return p.parse_args()
79
+
80
+ def _fetch_first_rows(repo: str, config: str | None, split: str, n: int) -> list[dict]:
81
+ from datasets import load_dataset
82
+
83
+ kwargs: dict[str, Any] = {"split": split, "streaming": True}
84
+ if config:
85
+ kwargs["name"] = config
86
+ ds = load_dataset(repo, token=os.environ.get("HF_TOKEN"), **kwargs)
87
+ out = []
88
+ for row in ds:
89
+ out.append(dict(row))
90
+ if len(out) >= n:
91
+ break
92
+ return out
93
+
94
+
95
+ def _build_tool_prompts(rows: list[dict]) -> list[dict]:
96
+ from train_securecoder import _normalise_tool_schema, _messages_from_any
97
+
98
+ prompts = []
99
+ for row in rows:
100
+ if not isinstance(row.get("messages"), list):
101
+ continue
102
+ tools_raw = row.get("tools")
103
+ if not tools_raw:
104
+ continue
105
+ tools = []
106
+ if isinstance(tools_raw, list):
107
+ for t in tools_raw:
108
+ n = _normalise_tool_schema(t)
109
+ if n:
110
+ tools.append(n)
111
+ elif isinstance(tools_raw, dict):
112
+ n = _normalise_tool_schema(tools_raw)
113
+ if n:
114
+ tools.append(n)
115
+ if not tools:
116
+ continue
117
+ messages, _ = _messages_from_any(row, "auto")
118
+ if not messages:
119
+ continue
120
+ user = next((m["content"] for m in messages if m["role"] == "user"), None)
121
+ if not user:
122
+ continue
123
+ prompts.append({"prompt": str(user)[:1200], "tools": tools,
124
+ "expected_call": any(m.get("tool_calls") for m in messages)})
125
+ if len(prompts) >= 200:
126
+ break
127
+ return prompts
128
+
129
+
130
+ def _render_prompt(tokenizer, prompt: str, tools: list[dict]) -> str:
131
+ return tokenizer.apply_chat_template(
132
+ [{"role": "user", "content": prompt}],
133
+ tools=tools, tokenize=False, add_generation_prompt=True,
134
+ )
135
+
136
+
137
+ def _parse_emitted_calls(text: str) -> list[dict]:
138
+ fns = TOOL_FN_PAT.findall(text)
139
+ if not fns:
140
+ return []
141
+ calls = []
142
+ for fn in fns:
143
+ start = text.find(f"<function={fn}>")
144
+ if start < 0:
145
+ continue
146
+ end = text.find("</function>", start)
147
+ block = text[start:end if end > 0 else start + 4000]
148
+ params = TOOL_PARAM_PAT.findall(block)
149
+ calls.append({"name": fn, "arguments": dict(params)})
150
+ return calls
151
+
152
+
153
+ def _score_call(call: dict, tools: list[dict]) -> dict:
154
+ name = call.get("name")
155
+ fn = next((t for t in tools if t.get("name") == name), None)
156
+ if not fn:
157
+ return {"parse": True, "name_ok": False, "schema_ok": False}
158
+ props = fn.get("parameters", {}).get("properties", {}) or {}
159
+ expected = set(props.keys())
160
+ given = set((call.get("arguments") or {}).keys())
161
+ return {"parse": True, "name_ok": True,
162
+ "schema_ok": expected.issubset(given) if expected else True,
163
+ "expected_keys": sorted(expected), "given_keys": sorted(given)}
164
+
165
+ def eval_tool_calls(tokenizer, model, args) -> dict:
166
+ import torch
167
+
168
+ log.info("tool-call eval: streaming candidates from hermes FC ...")
169
+ rows = _fetch_first_rows("NousResearch/hermes-function-calling-v1", "func_calling", "train",
170
+ args.tool_prompts * 4)
171
+ prompts = _build_tool_prompts(rows)
172
+ if args.tool_prompts:
173
+ prompts = prompts[: args.tool_prompts]
174
+ log.info("tool-call eval: %d usable prompts", len(prompts))
175
+
176
+ out = []
177
+ for i, p in enumerate(prompts):
178
+ try:
179
+ text = _render_prompt(tokenizer, p["prompt"], p["tools"])
180
+ ids = tokenizer(text, return_tensors="pt", add_special_tokens=False).input_ids.to(model.device)
181
+ with torch.no_grad():
182
+ generated = model.generate(ids, max_new_tokens=args.max_new_tokens, do_sample=False)
183
+ reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True)
184
+ except Exception as exc: # noqa: BLE001
185
+ out.append({"prompt": p["prompt"][:60], "error": repr(exc)[:120]})
186
+ continue
187
+
188
+ calls = _parse_emitted_calls(reply)
189
+ scored = [_score_call(c, p["tools"]) for c in calls]
190
+ out.append({
191
+ "prompt": p["prompt"][:80],
192
+ "reply_first_160": reply[:160],
193
+ "expected_call": p["expected_call"],
194
+ "n_calls": len(calls),
195
+ "calls": calls,
196
+ "scores": scored,
197
+ })
198
+ if (i + 1) % 25 == 0:
199
+ log.info(" tool-call progress: %d/%d", i + 1, len(prompts))
200
+
201
+ n = len(out)
202
+ parse_ok = sum(1 for r in out if r.get("scores") and any(s["parse"] for s in r["scores"]))
203
+ name_ok = sum(1 for r in out if r.get("scores") and any(s["name_ok"] for s in r["scores"]))
204
+ schema_ok = sum(1 for r in out if r.get("scores") and any(s["schema_ok"] for s in r["scores"]))
205
+ return {"section": "tool_calls", "n_prompts": n,
206
+ "parse_rate": parse_ok / max(n, 1), "name_rate": name_ok / max(n, 1),
207
+ "schema_rate": schema_ok / max(n, 1), "details": out}
208
+
209
+
210
+ def eval_code_sanity(tokenizer, model, args) -> dict:
211
+ import torch
212
+
213
+ out = []
214
+ prompts = CODE_PROMPTS[: args.code_prompts]
215
+ for i, prompt in enumerate(prompts):
216
+ text = tokenizer.apply_chat_template(
217
+ [{"role": "user", "content": prompt}],
218
+ tokenize=False, add_generation_prompt=True,
219
+ )
220
+ ids = tokenizer(text, return_tensors="pt", add_special_tokens=False).input_ids.to(model.device)
221
+ try:
222
+ with torch.no_grad():
223
+ generated = model.generate(ids, max_new_tokens=384, do_sample=False)
224
+ reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True)
225
+ except Exception as exc: # noqa: BLE001
226
+ out.append({"prompt": prompt[:60], "error": repr(exc)[:120]})
227
+ continue
228
+
229
+ block = None
230
+ match = re.search(r"```(?:python)?\s*\n(.*?)```", reply, re.S)
231
+ if match:
232
+ block = match.group(1)
233
+ else:
234
+ start = reply.find("def ")
235
+ if start >= 0:
236
+ block = reply[start:]
237
+
238
+ parsed = compiles = None
239
+ if block:
240
+ try:
241
+ ast.parse(block)
242
+ parsed = True
243
+ except SyntaxError:
244
+ parsed = False
245
+ block = None
246
+ if block:
247
+ try:
248
+ compile(block, "<eval>", "exec")
249
+ compiles = True
250
+ except Exception: # noqa: BLE001
251
+ compiles = False
252
+ out.append({"prompt": prompt[:60], "ast_ok": parsed, "compile_ok": compiles,
253
+ "reply_first_160": reply[:160]})
254
+ if (i + 1) % 5 == 0:
255
+ log.info(" code sanity: %d/%d", i + 1, len(prompts))
256
+
257
+ n = len(out)
258
+ ast_ok = sum(1 for r in out if r.get("ast_ok"))
259
+ compile_ok = sum(1 for r in out if r.get("compile_ok"))
260
+ return {"section": "code_sanity", "n_prompts": n,
261
+ "ast_rate": ast_ok / max(n, 1), "compile_rate": compile_ok / max(n, 1),
262
+ "details": out}
263
+
264
+ def eval_security_mcq(tokenizer, model, n_questions: int = 25) -> dict:
265
+ import torch
266
+
267
+ try:
268
+ rows = _fetch_first_rows("CyberNative/CyberSecurityEval", None, "train", n_questions * 2)
269
+ except Exception as exc: # noqa: BLE001
270
+ return {"section": "security_mcq", "error": repr(exc)[:200], "skipped": True}
271
+
272
+ rows = rows[:n_questions]
273
+ if not rows:
274
+ return {"section": "security_mcq", "skipped": True, "reason": "no rows"}
275
+
276
+ correct = 0
277
+ details = []
278
+ for r in rows:
279
+ question = r.get("question") or r.get("prompt") or r.get("input")
280
+ options = r.get("options") or r.get("choices") or r.get("answers")
281
+ answer = r.get("answer") or r.get("label")
282
+ if not question or not options or answer is None:
283
+ continue
284
+ if isinstance(options, dict):
285
+ opts = "\n".join(f"{k}. {v}" for k, v in options.items())
286
+ key_map = {str(k): v for k, v in options.items()}
287
+ else:
288
+ opts = "\n".join(f"{i}. {o}" for i, o in enumerate(options))
289
+ key_map = {str(i): options[i]}
290
+
291
+ user = f"Question: {question}\n\n{opts}\n\nRespond with the letter of the correct answer only."
292
+ text = tokenizer.apply_chat_template(
293
+ [{"role": "user", "content": user}], tokenize=False, add_generation_prompt=True,
294
+ )
295
+ ids = tokenizer(text, return_tensors="pt", add_special_tokens=False).input_ids.to(model.device)
296
+ try:
297
+ with torch.no_grad():
298
+ generated = model.generate(ids, max_new_tokens=8, do_sample=False)
299
+ reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True).strip()
300
+ except Exception: # noqa: BLE001
301
+ continue
302
+
303
+ first_letter = reply[:1].upper()
304
+ predicted = key_map.get(first_letter)
305
+ is_correct = predicted == answer
306
+ correct += int(is_correct)
307
+ details.append({"question": str(question)[:80], "reply": reply[:10], "ok": is_correct})
308
+
309
+ return {
310
+ "section": "security_mcq",
311
+ "n_questions": len(details),
312
+ "accuracy": correct / max(len(details), 1),
313
+ "details": details,
314
+ }
315
+
316
+
317
+ def main() -> int:
318
+ args = parse_args()
319
+ token = os.environ.get("HF_TOKEN")
320
+ if not token:
321
+ log.error("HF_TOKEN not set")
322
+ return 1
323
+
324
+ os.makedirs(args.out_dir, exist_ok=True)
325
+ import torch
326
+ from transformers import AutoTokenizer
327
+ from peft import PeftModel
328
+ from unsloth import FastLanguageModel
329
+
330
+ random.seed(args.seed)
331
+ log.info("loading adapter %s on top of %s ...", args.adapter, args.base)
332
+ model, tokenizer = FastLanguageModel.from_pretrained(
333
+ model_name=args.base, max_seq_length=2048, dtype=torch.bfloat16, load_in_4bit=True,
334
+ )
335
+ model = PeftModel.from_pretrained(model, args.adapter, token=token)
336
+ log.info("adapter loaded")
337
+
338
+ sections = [
339
+ eval_tool_calls(tokenizer, model, args),
340
+ eval_code_sanity(tokenizer, model, args),
341
+ eval_security_mcq(tokenizer, model),
342
+ ]
343
+
344
+ summary = {
345
+ "adapter": args.adapter,
346
+ "base": args.base,
347
+ "when": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
348
+ "sections": [
349
+ {"section": s["section"], **{k: v for k, v in s.items() if k not in {"section", "details"}}}
350
+ for s in sections
351
+ ],
352
+ "raw": sections,
353
+ }
354
+
355
+ out_json = os.path.join(args.out_dir, "report.json")
356
+ with open(out_json, "w", encoding="utf-8") as fh:
357
+ json.dump(summary, fh, indent=2, default=str)
358
+ log.info("report written: %s", out_json)
359
+
360
+ print("\n" + "=" * 70)
361
+ print(f"{'section':<18}{'metric':<22}{'value':>10}")
362
+ print("-" * 70)
363
+ for s in sections:
364
+ for k, v in s.items():
365
+ if isinstance(v, float) and k.endswith(("rate", "accuracy")):
366
+ print(f"{s['section']:<18}{k:<22}{v*100:>9.1f}%")
367
+ if "n_prompts" in s:
368
+ print(f"{s['section']:<18}{'n_prompts':<22}{s['n_prompts']:>10}")
369
+ print("=" * 70)
370
+
371
+ if args.upload_repo:
372
+ from huggingface_hub import HfApi
373
+ api = HfApi(token=token)
374
+ api.create_repo(args.upload_repo, repo_type="model", exist_ok=True, private=True)
375
+ api.upload_file(path_or_fileobj=out_json, path_in_repo="eval-report.json",
376
+ repo_id=args.upload_repo, repo_type="model",
377
+ commit_message="Add evaluation report")
378
+ log.info("report pushed to https://huggingface.co/%s", args.upload_repo)
379
+
380
+ return 0
381
+
382
+
383
+ if __name__ == "__main__":
384
+ raise SystemExit(main())