File size: 16,616 Bytes
a2ba1ae
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
# /// script
# requires-python = ">=3.10"
# dependencies = [
#     "datasets",
#     "huggingface_hub",
# ]
# ///
"""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.



The script 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")
    p.add_argument("--base", default="unsloth/Qwen3-Coder-30B-A3B-Instruct")
    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()

def _fetch_first_rows(repo: str, config: str | None, split: str, n: int) -> list[dict]:
    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]:
    from train_securecoder import _normalise_tool_schema, _messages_from_any

    prompts = []
    for row in rows:
        if not isinstance(row.get("messages"), list):
            continue
        tools_raw = row.get("tools")
        if not tools_raw:
            continue
        tools = []
        if isinstance(tools_raw, list):
            for t in tools_raw:
                n = _normalise_tool_schema(t)
                if n:
                    tools.append(n)
        elif isinstance(tools_raw, dict):
            n = _normalise_tool_schema(tools_raw)
            if n:
                tools.append(n)
        if not tools:
            continue
        messages, _ = _messages_from_any(row, "auto")
        if 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": any(m.get("tool_calls") for m in messages)})
        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)
            reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True)
        except Exception as exc:  # noqa: BLE001
            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)
            reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True)
        except Exception as exc:  # noqa: BLE001
            out.append({"prompt": prompt[:60], "error": repr(exc)[:120]})
            continue

        block = None
        match = re.search(r"```(?:python)?\s*\n(.*?)```", reply, re.S)
        if match:
            block = match.group(1)
        else:
            start = reply.find("def ")
            if start >= 0:
                block = reply[start:]

        parsed = compiles = None
        if block:
            try:
                ast.parse(block)
                parsed = True
            except SyntaxError:
                parsed = False
                block = None
        if block:
            try:
                compile(block, "<eval>", "exec")
                compiles = True
            except Exception:  # noqa: BLE001
                compiles = False
        out.append({"prompt": prompt[:60], "ast_ok": parsed, "compile_ok": compiles,
                    "reply_first_160": reply[:160]})
        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:  # noqa: BLE001
        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)
            reply = tokenizer.decode(generated[0, ids.shape[-1]:], skip_special_tokens=True).strip()
        except Exception:  # noqa: BLE001
            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 AutoTokenizer
    from peft import PeftModel
    from unsloth import FastLanguageModel

    random.seed(args.seed)
    log.info("loading adapter %s on top of %s ...", args.adapter, args.base)
    model, tokenizer = FastLanguageModel.from_pretrained(
        model_name=args.base, max_seq_length=2048, dtype=torch.bfloat16, load_in_4bit=True,
    )
    model = PeftModel.from_pretrained(model, args.adapter, token=token)
    log.info("adapter 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())