Eval: inline helpers (no cross-script import)
Browse files- eval_securecoder.py +91 -7
eval_securecoder.py
CHANGED
|
@@ -76,7 +76,7 @@ def parse_args() -> argparse.Namespace:
|
|
| 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 |
|
|
@@ -87,13 +87,97 @@ def _fetch_first_rows(repo: str, config: str | None, split: str, n: int) -> list
|
|
| 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 |
-
|
| 97 |
|
| 98 |
prompts = []
|
| 99 |
for row in rows:
|
|
@@ -105,16 +189,16 @@ def _build_tool_prompts(rows: list[dict]) -> list[dict]:
|
|
| 105 |
tools = []
|
| 106 |
if isinstance(tools_raw, list):
|
| 107 |
for t in tools_raw:
|
| 108 |
-
n =
|
| 109 |
if n:
|
| 110 |
tools.append(n)
|
| 111 |
elif isinstance(tools_raw, dict):
|
| 112 |
-
n =
|
| 113 |
if n:
|
| 114 |
tools.append(n)
|
| 115 |
if not tools:
|
| 116 |
continue
|
| 117 |
-
messages, _ =
|
| 118 |
if not messages:
|
| 119 |
continue
|
| 120 |
user = next((m["content"] for m in messages if m["role"] == "user"), None)
|
|
@@ -161,7 +245,7 @@ def _score_call(call: dict, tools: list[dict]) -> dict:
|
|
| 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 |
|
|
@@ -260,7 +344,7 @@ def eval_code_sanity(tokenizer, model, args) -> dict:
|
|
| 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 |
|
|
|
|
| 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 |
|
|
|
|
| 87 |
out = []
|
| 88 |
for row in ds:
|
| 89 |
out.append(dict(row))
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
# --------------------------------------------------------------------------
|
| 93 |
+
# Local helpers (inlined so the script runs standalone on HF Jobs without
|
| 94 |
+
# needing train_securecoder.py to also be uploaded).
|
| 95 |
+
# --------------------------------------------------------------------------
|
| 96 |
+
_TYPE_ALIASES = {
|
| 97 |
+
"str": "string", "string": "string", "text": "string",
|
| 98 |
+
"int": "integer", "integer": "integer", "long": "integer",
|
| 99 |
+
"float": "number", "double": "number", "number": "number",
|
| 100 |
+
"bool": "boolean", "boolean": "boolean",
|
| 101 |
+
"list": "array", "array": "array", "dict": "object", "object": "object",
|
| 102 |
+
}
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def _normalise_tool_schema_local(tool: Any) -> dict | None:
|
| 106 |
+
if not isinstance(tool, dict):
|
| 107 |
+
return None
|
| 108 |
+
fn = tool.get("function") if isinstance(tool.get("function"), dict) else tool
|
| 109 |
+
if not isinstance(fn, dict) or not fn.get("name"):
|
| 110 |
+
return None
|
| 111 |
+
params = fn.get("parameters") or {}
|
| 112 |
+
if isinstance(params, dict) and params and "properties" not in params:
|
| 113 |
+
properties: dict[str, Any] = {}
|
| 114 |
+
required: list[str] = []
|
| 115 |
+
for name, spec in params.items():
|
| 116 |
+
if isinstance(spec, dict):
|
| 117 |
+
cleaned = {k: v for k, v in spec.items()
|
| 118 |
+
if k in ("type", "description", "enum", "default", "title", "items")}
|
| 119 |
+
cleaned["type"] = _TYPE_ALIASES.get(
|
| 120 |
+
str(cleaned.get("type", "")).lower(), "string")
|
| 121 |
+
properties[name] = cleaned
|
| 122 |
+
if "default" not in spec:
|
| 123 |
+
required.append(name)
|
| 124 |
+
else:
|
| 125 |
+
properties[name] = {"type": "string"}
|
| 126 |
+
required.append(name)
|
| 127 |
+
params = {"type": "object", "properties": properties}
|
| 128 |
+
if required:
|
| 129 |
+
params["required"] = required
|
| 130 |
+
return {"name": fn["name"], "description": fn.get("description", ""), "parameters": params}
|
| 131 |
+
|
| 132 |
+
|
| 133 |
+
def _messages_from_any_local(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
|
| 134 |
+
tools: list[dict] = []
|
| 135 |
+
raw = row.get("tools")
|
| 136 |
+
candidates = [raw] if isinstance(raw, dict) else (raw if isinstance(raw, list) else [])
|
| 137 |
+
for c in candidates:
|
| 138 |
+
n = _normalise_tool_schema_local(c)
|
| 139 |
+
if n:
|
| 140 |
+
tools.append(n)
|
| 141 |
+
|
| 142 |
+
messages: list[dict] = []
|
| 143 |
+
if isinstance(row.get("messages"), list):
|
| 144 |
+
for turn in row["messages"]:
|
| 145 |
+
if not isinstance(turn, dict):
|
| 146 |
+
continue
|
| 147 |
+
role = turn.get("role")
|
| 148 |
+
content = turn.get("content", "")
|
| 149 |
+
if role in {"user", "assistant", "system"} and content:
|
| 150 |
+
messages.append({"role": role, "content": str(content)})
|
| 151 |
+
return messages, tools
|
| 152 |
+
|
| 153 |
+
if row.get("query") and row.get("answers"):
|
| 154 |
+
messages.append({"role": "user", "content": str(row["query"])})
|
| 155 |
+
messages.append({"role": "assistant", "content": str(row["answers"])})
|
| 156 |
+
return messages, tools
|
| 157 |
+
|
| 158 |
+
if row.get("user") and row.get("assistant"):
|
| 159 |
+
if row.get("system"):
|
| 160 |
+
messages.append({"role": "system", "content": str(row["system"])})
|
| 161 |
+
messages.append({"role": "user", "content": str(row["user"])})
|
| 162 |
+
messages.append({"role": "assistant", "content": str(row["assistant"])})
|
| 163 |
+
return messages, tools
|
| 164 |
+
|
| 165 |
+
if row.get("instruction") and (row.get("output") or row.get("response")):
|
| 166 |
+
messages.append({"role": "user", "content": str(row["instruction"])})
|
| 167 |
+
messages.append({"role": "assistant", "content": str(row.get("output") or row.get("response"))})
|
| 168 |
+
return messages, tools
|
| 169 |
+
|
| 170 |
+
return [], tools
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
|
| 174 |
if len(out) >= n:
|
| 175 |
break
|
| 176 |
return out
|
| 177 |
|
| 178 |
|
| 179 |
def _build_tool_prompts(rows: list[dict]) -> list[dict]:
|
| 180 |
+
|
| 181 |
|
| 182 |
prompts = []
|
| 183 |
for row in rows:
|
|
|
|
| 189 |
tools = []
|
| 190 |
if isinstance(tools_raw, list):
|
| 191 |
for t in tools_raw:
|
| 192 |
+
n = _normalise_tool_schema_local(t)
|
| 193 |
if n:
|
| 194 |
tools.append(n)
|
| 195 |
elif isinstance(tools_raw, dict):
|
| 196 |
+
n = _normalise_tool_schema_local(tools_raw)
|
| 197 |
if n:
|
| 198 |
tools.append(n)
|
| 199 |
if not tools:
|
| 200 |
continue
|
| 201 |
+
messages, _ = _messages_from_any_local(row, "auto")
|
| 202 |
if not messages:
|
| 203 |
continue
|
| 204 |
user = next((m["content"] for m in messages if m["role"] == "user"), None)
|
|
|
|
| 245 |
return {"parse": True, "name_ok": True,
|
| 246 |
"schema_ok": expected.issubset(given) if expected else True,
|
| 247 |
"expected_keys": sorted(expected), "given_keys": sorted(given)}
|
| 248 |
+
|
| 249 |
def eval_tool_calls(tokenizer, model, args) -> dict:
|
| 250 |
import torch
|
| 251 |
|
|
|
|
| 344 |
return {"section": "code_sanity", "n_prompts": n,
|
| 345 |
"ast_rate": ast_ok / max(n, 1), "compile_rate": compile_ok / max(n, 1),
|
| 346 |
"details": out}
|
| 347 |
+
|
| 348 |
def eval_security_mcq(tokenizer, model, n_questions: int = 25) -> dict:
|
| 349 |
import torch
|
| 350 |
|