Format-adaptive tool rendering + resilient dataset loading
Browse files- train_securecoder.py +192 -45
train_securecoder.py
CHANGED
|
@@ -170,26 +170,41 @@ def _normalise_parameters(params: Any) -> dict:
|
|
| 170 |
|
| 171 |
|
| 172 |
def _normalise_tool_schema(tool: Any) -> dict | None:
|
| 173 |
-
"""
|
| 174 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 175 |
if not isinstance(tool, dict):
|
| 176 |
return None
|
| 177 |
-
fn = tool.get("function") if "function"
|
| 178 |
if not isinstance(fn, dict) or not fn.get("name"):
|
| 179 |
return None
|
| 180 |
return {
|
| 181 |
-
"
|
| 182 |
-
"
|
| 183 |
-
|
| 184 |
-
"description": fn.get("description", ""),
|
| 185 |
-
"parameters": _normalise_parameters(fn.get("parameters")),
|
| 186 |
-
},
|
| 187 |
}
|
| 188 |
|
| 189 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 190 |
|
| 191 |
def _parse_calls(value: Any) -> list[dict] | None:
|
| 192 |
-
"""Return OpenAI
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 193 |
if isinstance(value, str):
|
| 194 |
text = value.strip()
|
| 195 |
if not text.startswith(("{", "[")):
|
|
@@ -204,12 +219,21 @@ def _parse_calls(value: Any) -> list[dict] | None:
|
|
| 204 |
calls = []
|
| 205 |
for i, item in enumerate(items):
|
| 206 |
arguments = item.get("arguments", item.get("parameters", {}))
|
| 207 |
-
if
|
| 208 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 209 |
calls.append({
|
| 210 |
"id": f"call_{i}",
|
| 211 |
"type": "function",
|
| 212 |
-
"function": {
|
|
|
|
|
|
|
|
|
|
|
|
|
| 213 |
})
|
| 214 |
return calls
|
| 215 |
|
|
@@ -305,16 +329,55 @@ def _messages_from_any(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
|
|
| 305 |
messages.append({"role": "assistant", "content": str(row["solution"])})
|
| 306 |
return messages, tools
|
| 307 |
|
| 308 |
-
# --- SecOps reasoning: goal/command ->
|
| 309 |
-
|
| 310 |
-
|
| 311 |
-
|
| 312 |
-
|
| 313 |
-
|
| 314 |
-
|
| 315 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 316 |
messages.append({"role": "user", "content": "\n".join(str(a) for a in ask)})
|
| 317 |
-
messages.append({"role": "assistant", "content":
|
| 318 |
return messages, tools
|
| 319 |
|
| 320 |
# --- generic prompt/completion fallback -------------------------------
|
|
@@ -331,34 +394,118 @@ def _messages_from_any(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
|
|
| 331 |
# --------------------------------------------------------------------------
|
| 332 |
def load_source(src: Source, token: str | None, progress: bool = False) -> list[dict]:
|
| 333 |
"""Pull up to ``limit`` rows from one Hub dataset, streaming so we never
|
| 334 |
-
download more than we need.
|
| 335 |
-
from datasets import load_dataset
|
| 336 |
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
|
| 340 |
-
|
| 341 |
-
|
| 342 |
|
| 343 |
-
|
| 344 |
-
|
| 345 |
-
|
| 346 |
-
|
| 347 |
-
|
| 348 |
-
|
| 349 |
-
|
| 350 |
-
|
| 351 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 352 |
|
| 353 |
|
| 354 |
def render_record(tokenizer, messages: list[dict], tools: list[dict] | None = None) -> str:
|
| 355 |
-
"""Render
|
| 356 |
-
|
| 357 |
-
|
| 358 |
-
|
| 359 |
-
|
| 360 |
-
|
| 361 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 362 |
|
| 363 |
|
| 364 |
def build_dataset(tokenizer, sources: list[Source], token: str | None, validate: bool):
|
|
|
|
| 170 |
|
| 171 |
|
| 172 |
def _normalise_tool_schema(tool: Any) -> dict | None:
|
| 173 |
+
"""Canonical **flat** tool schema: {"name", "description", "parameters"}.
|
| 174 |
+
|
| 175 |
+
Qwen3-Coder's chat template walks ``tool.parameters.properties`` directly, so
|
| 176 |
+
the flat form is what we store; ``_tools_for_style`` re-wraps it into the
|
| 177 |
+
OpenAI ``{"type": "function", "function": {...}}`` shape for templates that
|
| 178 |
+
want that instead.
|
| 179 |
+
"""
|
| 180 |
if not isinstance(tool, dict):
|
| 181 |
return None
|
| 182 |
+
fn = tool.get("function") if isinstance(tool.get("function"), dict) else tool
|
| 183 |
if not isinstance(fn, dict) or not fn.get("name"):
|
| 184 |
return None
|
| 185 |
return {
|
| 186 |
+
"name": fn["name"],
|
| 187 |
+
"description": fn.get("description", ""),
|
| 188 |
+
"parameters": _normalise_parameters(fn.get("parameters")),
|
|
|
|
|
|
|
|
|
|
| 189 |
}
|
| 190 |
|
| 191 |
|
| 192 |
+
def _tools_for_style(tools: list[dict], style: str) -> list[dict] | None:
|
| 193 |
+
if not tools:
|
| 194 |
+
return None
|
| 195 |
+
if style == "nested":
|
| 196 |
+
return [{"type": "function", "function": t} for t in tools]
|
| 197 |
+
return tools
|
| 198 |
+
|
| 199 |
+
|
| 200 |
|
| 201 |
def _parse_calls(value: Any) -> list[dict] | None:
|
| 202 |
+
"""Return OpenAI-style tool calls if ``value`` is one or more function calls.
|
| 203 |
+
|
| 204 |
+
``arguments`` is kept as a **dict** plus a JSON-string copy, because templates
|
| 205 |
+
disagree: Qwen3-Coder iterates ``arguments | items`` (needs a mapping) while
|
| 206 |
+
others print a JSON string.
|
| 207 |
+
"""
|
| 208 |
if isinstance(value, str):
|
| 209 |
text = value.strip()
|
| 210 |
if not text.startswith(("{", "[")):
|
|
|
|
| 219 |
calls = []
|
| 220 |
for i, item in enumerate(items):
|
| 221 |
arguments = item.get("arguments", item.get("parameters", {}))
|
| 222 |
+
if isinstance(arguments, str):
|
| 223 |
+
try:
|
| 224 |
+
arguments = json.loads(arguments)
|
| 225 |
+
except json.JSONDecodeError:
|
| 226 |
+
arguments = {"value": arguments}
|
| 227 |
+
if not isinstance(arguments, dict):
|
| 228 |
+
arguments = {"value": arguments}
|
| 229 |
calls.append({
|
| 230 |
"id": f"call_{i}",
|
| 231 |
"type": "function",
|
| 232 |
+
"function": {
|
| 233 |
+
"name": item["name"],
|
| 234 |
+
"arguments": arguments,
|
| 235 |
+
"arguments_json": json.dumps(arguments),
|
| 236 |
+
},
|
| 237 |
})
|
| 238 |
return calls
|
| 239 |
|
|
|
|
| 329 |
messages.append({"role": "assistant", "content": str(row["solution"])})
|
| 330 |
return messages, tools
|
| 331 |
|
| 332 |
+
# --- SecOps reasoning: goal/command -> interpretation ------------------
|
| 333 |
+
# dpevzner's rows carry goal + command_sequence + interpretation +
|
| 334 |
+
# classification + safety_and_scope (unified_interpretation is empty in the
|
| 335 |
+
# published revision). Scope metadata goes into the prompt so the model
|
| 336 |
+
# learns the framed, authorised-use context alongside the command knowledge.
|
| 337 |
+
if row.get("goal") and row.get("interpretation"):
|
| 338 |
+
seq = row.get("command_sequence") or {}
|
| 339 |
+
if isinstance(seq, str):
|
| 340 |
+
seq = {"command": seq}
|
| 341 |
+
interp = row.get("interpretation") or {}
|
| 342 |
+
if isinstance(interp, str):
|
| 343 |
+
interp = {"what_it_means": [interp]}
|
| 344 |
+
scope = row.get("safety_and_scope") or {}
|
| 345 |
+
if isinstance(scope, str):
|
| 346 |
+
scope = {}
|
| 347 |
+
|
| 348 |
+
tool = row.get("tool") or {}
|
| 349 |
+
if isinstance(tool, str):
|
| 350 |
+
tool = {"name": tool}
|
| 351 |
+
ask = [f"Environment: {tool.get('name', 'shell')} ({tool.get('platform', 'unknown')})"]
|
| 352 |
+
if seq.get("command"):
|
| 353 |
+
ask.append(f"Command: {seq['command']}")
|
| 354 |
+
ask.append(f"Goal: {row['goal']}")
|
| 355 |
+
if isinstance(scope, dict) and scope:
|
| 356 |
+
ask.append("Scope: " + ", ".join(f"{k}={v}" for k, v in list(scope.items())[:4]))
|
| 357 |
+
if row.get("classification"):
|
| 358 |
+
cls = row["classification"]
|
| 359 |
+
if isinstance(cls, dict):
|
| 360 |
+
ask.append("Context: " + ", ".join(f"{k}={v}" for k, v in list(cls.items())[:3]))
|
| 361 |
+
ask.append("Explain what this command does, what its output means, what it tells you "
|
| 362 |
+
"about the target, and the next step in an authorised assessment.")
|
| 363 |
+
|
| 364 |
+
answer = []
|
| 365 |
+
if seq.get("description"):
|
| 366 |
+
answer.append(f"**What it does.** {seq['description']}")
|
| 367 |
+
if seq.get("expected_output_pattern"):
|
| 368 |
+
answer.append("**Expected output.** " + ", ".join(map(str, seq["expected_output_pattern"])))
|
| 369 |
+
for item in interp.get("what_it_means", []) if isinstance(interp, dict) else []:
|
| 370 |
+
answer.append(f"**What it means.** {item}")
|
| 371 |
+
for item in interp.get("risk_indicators", []) if isinstance(interp, dict) else []:
|
| 372 |
+
answer.append(f"**Risk indicators.** {item}")
|
| 373 |
+
if row.get("ambiguity_analysis"):
|
| 374 |
+
answer.append(f"**Ambiguity.** {row['ambiguity_analysis']}")
|
| 375 |
+
if isinstance(scope, dict) and scope.get("authorization_required"):
|
| 376 |
+
answer.append("**Scope.** Only run this against systems you are authorised to test.")
|
| 377 |
+
if len(answer) < 2:
|
| 378 |
+
return [], tools
|
| 379 |
messages.append({"role": "user", "content": "\n".join(str(a) for a in ask)})
|
| 380 |
+
messages.append({"role": "assistant", "content": "\n\n".join(answer)})
|
| 381 |
return messages, tools
|
| 382 |
|
| 383 |
# --- generic prompt/completion fallback -------------------------------
|
|
|
|
| 394 |
# --------------------------------------------------------------------------
|
| 395 |
def load_source(src: Source, token: str | None, progress: bool = False) -> list[dict]:
|
| 396 |
"""Pull up to ``limit`` rows from one Hub dataset, streaming so we never
|
| 397 |
+
download more than we need.
|
|
|
|
| 398 |
|
| 399 |
+
Datasets move: builder configs get renamed (a config called ``default`` last
|
| 400 |
+
week is ``chatml`` today) and splits get added. Each candidate is tried in
|
| 401 |
+
turn so one rename cannot silently empty a slice of the mix.
|
| 402 |
+
"""
|
| 403 |
+
from datasets import load_dataset
|
| 404 |
|
| 405 |
+
candidates = [
|
| 406 |
+
(src.config, src.split),
|
| 407 |
+
(src.config, "train"),
|
| 408 |
+
(src.config, "test"),
|
| 409 |
+
(None, src.split),
|
| 410 |
+
(None, "train"),
|
| 411 |
+
(None, "test"),
|
| 412 |
+
]
|
| 413 |
+
seen: set = set()
|
| 414 |
+
last_exc: Exception | None = None
|
| 415 |
+
|
| 416 |
+
for config, split in candidates:
|
| 417 |
+
if (config, split) in seen:
|
| 418 |
+
continue
|
| 419 |
+
seen.add((config, split))
|
| 420 |
+
kwargs: dict[str, Any] = {"split": split, "streaming": True}
|
| 421 |
+
if config:
|
| 422 |
+
kwargs["name"] = config
|
| 423 |
+
if token:
|
| 424 |
+
kwargs["token"] = token
|
| 425 |
+
try:
|
| 426 |
+
ds = load_dataset(src.repo, **kwargs)
|
| 427 |
+
rows = []
|
| 428 |
+
for i, row in enumerate(ds):
|
| 429 |
+
if i >= src.limit:
|
| 430 |
+
break
|
| 431 |
+
rows.append(dict(row))
|
| 432 |
+
if progress and i and i % 2500 == 0:
|
| 433 |
+
log.info(" %s: %d rows...", source_name(src), i)
|
| 434 |
+
if not rows:
|
| 435 |
+
last_exc = ValueError(f"config={config} split={split} streamed 0 rows")
|
| 436 |
+
continue
|
| 437 |
+
if (config, split) != (src.config, src.split):
|
| 438 |
+
log.info(" %s: fell back to config=%s split=%s",
|
| 439 |
+
src.repo, config, split)
|
| 440 |
+
return rows
|
| 441 |
+
except Exception as exc: # noqa: BLE001 - try the next candidate
|
| 442 |
+
last_exc = exc
|
| 443 |
+
|
| 444 |
+
raise last_exc if last_exc else RuntimeError(f"could not load {src.repo}")
|
| 445 |
+
|
| 446 |
+
|
| 447 |
+
_RENDER_STYLE: str | None = None
|
| 448 |
+
STYLE_ATTEMPTS = (("flat", "dict"), ("nested", "dict"), ("flat", "string"), ("nested", "string"))
|
| 449 |
+
|
| 450 |
+
|
| 451 |
+
def _apply_arg_style(messages: list[dict], arg_style: str) -> list[dict]:
|
| 452 |
+
"""Copy messages, swapping tool-call arguments between dict and JSON string."""
|
| 453 |
+
if arg_style != "string":
|
| 454 |
+
return messages
|
| 455 |
+
out = []
|
| 456 |
+
for message in messages:
|
| 457 |
+
if message.get("tool_calls"):
|
| 458 |
+
message = dict(message)
|
| 459 |
+
message["tool_calls"] = [
|
| 460 |
+
{
|
| 461 |
+
"id": call["id"],
|
| 462 |
+
"type": "function",
|
| 463 |
+
"function": {
|
| 464 |
+
"name": call["function"]["name"],
|
| 465 |
+
"arguments": call["function"].get(
|
| 466 |
+
"arguments_json", json.dumps(call["function"]["arguments"])
|
| 467 |
+
),
|
| 468 |
+
},
|
| 469 |
+
}
|
| 470 |
+
for call in message["tool_calls"]
|
| 471 |
+
]
|
| 472 |
+
out.append(message)
|
| 473 |
+
return out
|
| 474 |
|
| 475 |
|
| 476 |
def render_record(tokenizer, messages: list[dict], tools: list[dict] | None = None) -> str:
|
| 477 |
+
"""Render with the model's native chat template.
|
| 478 |
+
|
| 479 |
+
Tool templates disagree in two independent ways: whether tool schemas are
|
| 480 |
+
flat (`{"name", "parameters"}`) or OpenAI-nested (`{"type": "function", ...}`),
|
| 481 |
+
and whether tool-call arguments are a mapping or a JSON string. Qwen3-Coder
|
| 482 |
+
renders ``<function=NAME><parameter=...>`` blocks and iterates
|
| 483 |
+
``arguments | items``, so a JSON string there is a hard error. Rather than
|
| 484 |
+
hard-code one convention, detect it once and reuse it for the rest of the
|
| 485 |
+
run - with the model, the mix and the template all free to change.
|
| 486 |
+
"""
|
| 487 |
+
global _RENDER_STYLE
|
| 488 |
+
|
| 489 |
+
order = []
|
| 490 |
+
if _RENDER_STYLE:
|
| 491 |
+
order.append(tuple(_RENDER_STYLE.split("+")))
|
| 492 |
+
order += [style for style in STYLE_ATTEMPTS if style not in order]
|
| 493 |
+
|
| 494 |
+
last_exc: Exception | None = None
|
| 495 |
+
for tool_style, arg_style in order:
|
| 496 |
+
try:
|
| 497 |
+
text = tokenizer.apply_chat_template(
|
| 498 |
+
_apply_arg_style(messages, arg_style),
|
| 499 |
+
tools=_tools_for_style(tools, tool_style),
|
| 500 |
+
tokenize=False,
|
| 501 |
+
add_generation_prompt=False,
|
| 502 |
+
)
|
| 503 |
+
_RENDER_STYLE = f"{tool_style}+{arg_style}"
|
| 504 |
+
return text
|
| 505 |
+
except Exception as exc: # noqa: BLE001 - try the next convention
|
| 506 |
+
last_exc = exc
|
| 507 |
+
raise last_exc if last_exc else RuntimeError("render failed")
|
| 508 |
+
|
| 509 |
|
| 510 |
|
| 511 |
def build_dataset(tokenizer, sources: list[Source], token: str | None, validate: bool):
|