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

Fix xLAM tool schemas + dpevzner config

Browse files
Files changed (1) hide show
  1. train_securecoder.py +58 -4
train_securecoder.py CHANGED
@@ -83,7 +83,7 @@ MIX: list[Source] = [
83
  note="broad security Q&A"),
84
  Source("Humanlearning/CyberSecurity_OWASP-sft-dataset", 3000, "messages",
85
  note="OWASP / secure-coding SFT"),
86
- Source("dpevzner/Cybersecurity_Reasoning_Dataset", 3000, "secops", "default", "test",
87
  note="command interpretation reasoning (goal -> unified_interpretation)"),
88
  Source("MrClipperz134/CTF-Instruct", 3000, "io",
89
  note="CTF instruction/output"),
@@ -123,6 +123,52 @@ def _as_list(value: Any) -> list:
123
  return []
124
 
125
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
126
  def _normalise_tool_schema(tool: Any) -> dict | None:
127
  """Force a tool schema into the shape Qwen's chat template expects:
128
  {"type": "function", "function": {"name", "description", "parameters"}}."""
@@ -136,11 +182,12 @@ def _normalise_tool_schema(tool: Any) -> dict | None:
136
  "function": {
137
  "name": fn["name"],
138
  "description": fn.get("description", ""),
139
- "parameters": fn.get("parameters") or {"type": "object", "properties": {}},
140
  },
141
  }
142
 
143
 
 
144
  def _parse_calls(value: Any) -> list[dict] | None:
145
  """Return OpenAI tool_calls if ``value`` is one or more JSON function calls."""
146
  if isinstance(value, str):
@@ -215,7 +262,14 @@ def _messages_from_any(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
215
 
216
  # --- xLAM: query + answers + tools ------------------------------------
217
  if kind == "xlam" or (row.get("query") and row.get("answers")):
218
- calls = _parse_calls(_as_list(row.get("answers")))
 
 
 
 
 
 
 
219
  if calls and row.get("query"):
220
  messages.append({"role": "user", "content": str(row["query"])})
221
  messages.append({"role": "assistant", "content": None, "tool_calls": calls})
@@ -252,7 +306,7 @@ def _messages_from_any(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
252
  return messages, tools
253
 
254
  # --- SecOps reasoning: goal/command -> unified_interpretation ---------
255
- if kind == "secops" and row.get("unified_interpretation"):
256
  ask = [f"Tool: {row.get('tool', 'shell')}", f"Goal: {row.get('goal', '')}"]
257
  for key in ("command", "command_sequence", "nmap_context"):
258
  if row.get(key):
 
83
  note="broad security Q&A"),
84
  Source("Humanlearning/CyberSecurity_OWASP-sft-dataset", 3000, "messages",
85
  note="OWASP / secure-coding SFT"),
86
+ Source("dpevzner/Cybersecurity_Reasoning_Dataset", 3000, "secops", "chatml", "train",
87
  note="command interpretation reasoning (goal -> unified_interpretation)"),
88
  Source("MrClipperz134/CTF-Instruct", 3000, "io",
89
  note="CTF instruction/output"),
 
123
  return []
124
 
125
 
126
+ TYPE_ALIASES = {
127
+ "str": "string", "string": "string", "text": "string",
128
+ "int": "integer", "integer": "integer", "long": "integer",
129
+ "float": "number", "double": "number", "number": "number",
130
+ "bool": "boolean", "boolean": "boolean",
131
+ "list": "array", "array": "array", "dict": "object", "object": "object",
132
+ }
133
+
134
+
135
+ def _normalise_parameters(params: Any) -> dict:
136
+ """Coerce a tool's parameter spec into valid JSON Schema.
137
+
138
+ Two shapes appear in the wild: proper ``{"type": "object", "properties": {}}``
139
+ and the flat ``{"arg": {"description": ..., "type": "str"}}`` form used by
140
+ xLAM. Qwen's chat template reads ``parameters.properties``, so the flat form
141
+ has to be wrapped or rendering raises.
142
+ """
143
+ if not isinstance(params, dict) or not params:
144
+ return {"type": "object", "properties": {}}
145
+
146
+ if "properties" in params:
147
+ params.setdefault("type", "object")
148
+ return params
149
+
150
+ properties: dict[str, Any] = {}
151
+ required: list[str] = []
152
+ for name, spec in params.items():
153
+ if isinstance(spec, dict):
154
+ cleaned = {
155
+ k: v for k, v in spec.items()
156
+ if k in ("type", "description", "enum", "default", "title", "items")
157
+ }
158
+ cleaned["type"] = TYPE_ALIASES.get(str(cleaned.get("type", "")).lower(), "string")
159
+ properties[name] = cleaned
160
+ if "default" not in spec:
161
+ required.append(name)
162
+ else:
163
+ properties[name] = {"type": "string"}
164
+ required.append(name)
165
+
166
+ schema: dict[str, Any] = {"type": "object", "properties": properties}
167
+ if required:
168
+ schema["required"] = required
169
+ return schema
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"}}."""
 
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):
 
262
 
263
  # --- xLAM: query + answers + tools ------------------------------------
264
  if kind == "xlam" or (row.get("query") and row.get("answers")):
265
+ answers = _as_list(row.get("answers"))
266
+ calls = _parse_calls(answers)
267
+ if calls is None and answers:
268
+ # answers can be a list of JSON strings instead of one JSON array
269
+ flat: list = []
270
+ for item in answers:
271
+ flat.extend(_as_list(item))
272
+ calls = _parse_calls(flat)
273
  if calls and row.get("query"):
274
  messages.append({"role": "user", "content": str(row["query"])})
275
  messages.append({"role": "assistant", "content": None, "tool_calls": calls})
 
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):