Taimwe commited on
Commit
fffea22
·
verified ·
1 Parent(s): 39cda51

Format-adaptive tool rendering + resilient dataset loading

Browse files
Files changed (1) hide show
  1. 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
- """Force a tool schema into the shape Qwen's chat template expects:
174
- {"type": "function", "function": {"name", "description", "parameters"}}."""
 
 
 
 
 
175
  if not isinstance(tool, dict):
176
  return None
177
- fn = tool.get("function") if "function" in tool else tool
178
  if not isinstance(fn, dict) or not fn.get("name"):
179
  return None
180
  return {
181
- "type": "function",
182
- "function": {
183
- "name": fn["name"],
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 tool_calls if ``value`` is one or more JSON function calls."""
 
 
 
 
 
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 not isinstance(arguments, str):
208
- arguments = json.dumps(arguments)
 
 
 
 
 
209
  calls.append({
210
  "id": f"call_{i}",
211
  "type": "function",
212
- "function": {"name": item["name"], "arguments": arguments},
 
 
 
 
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 -> unified_interpretation ---------
309
- if row.get("unified_interpretation"):
310
- ask = [f"Tool: {row.get('tool', 'shell')}", f"Goal: {row.get('goal', '')}"]
311
- for key in ("command", "command_sequence", "nmap_context"):
312
- if row.get(key):
313
- ask.append(f"{key}: {row[key]}")
314
- ask.append("Explain what the output means, what it tells you about the target, "
315
- "and what the next step should be.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
316
  messages.append({"role": "user", "content": "\n".join(str(a) for a in ask)})
317
- messages.append({"role": "assistant", "content": str(row["unified_interpretation"])})
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
- kwargs: dict[str, Any] = {"split": src.split, "streaming": True}
338
- if src.config:
339
- kwargs["name"] = src.config
340
- if token:
341
- kwargs["token"] = token
342
 
343
- ds = load_dataset(src.repo, **kwargs)
344
- rows = []
345
- for i, row in enumerate(ds):
346
- if i >= src.limit:
347
- break
348
- rows.append(dict(row))
349
- if progress and i and i % 2500 == 0:
350
- log.info(" %s: %d rows...", source_name(src), i)
351
- return rows
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
352
 
353
 
354
  def render_record(tokenizer, messages: list[dict], tools: list[dict] | None = None) -> str:
355
- """Render to the model's native chat format (Qwen3 emits <tool_call> blocks)."""
356
- return tokenizer.apply_chat_template(
357
- messages,
358
- tools=tools or None,
359
- tokenize=False,
360
- add_generation_prompt=False,
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):