Taimwe commited on
Commit
3d3d53f
·
verified ·
1 Parent(s): 9fa3ddd

Eval: inline helpers (no cross-script import)

Browse files
Files changed (1) hide show
  1. 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
- from train_securecoder import _normalise_tool_schema, _messages_from_any
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 = _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)
@@ -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