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

SecureCoder trainer v1

Browse files
Files changed (1) hide show
  1. train_securecoder.py +642 -0
train_securecoder.py ADDED
@@ -0,0 +1,642 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # requires-python = ">=3.10"
3
+ # dependencies = [
4
+ # "unsloth",
5
+ # "datasets",
6
+ # "trl>=0.22",
7
+ # "transformers>=4.57",
8
+ # "trackio",
9
+ # "huggingface_hub",
10
+ # ]
11
+ # ///
12
+ """SecureCoder: QLoRA fine-tune for code + tool calling + cybersecurity.
13
+
14
+ Default base: Qwen/Qwen3-Coder-30B-A3B-Instruct (Apache-2.0, 30.5B MoE, ~3B
15
+ active) - a MoE that trains like a small model and runs like a useful one.
16
+
17
+ Runs anywhere; same file for local validation, a GPU smoke test, and the real
18
+ run:
19
+
20
+ uv run train_securecoder.py --validate-only # no GPU needed
21
+ uv run train_securecoder.py --smoke --output-repo you/securecoder-smoke
22
+ uv run train_securecoder.py --num-epochs 1 --output-repo you/securecoder-30b-pro
23
+
24
+ Launch on Hugging Face Jobs (see README.md for why the URL form is used):
25
+
26
+ hf jobs run -d --flavor l40sx1 --timeout 12h --secrets HF_TOKEN \\
27
+ ghcr.io/astral-sh/uv:python3.12-bookworm \\
28
+ uv run --no-project https://huggingface.co/USER/securecoder-scripts/resolve/main/train_securecoder.py \\
29
+ -- --num-epochs 1 --output-repo USER/securecoder-30b-pro
30
+ """
31
+
32
+ from __future__ import annotations
33
+
34
+ import argparse
35
+ import json
36
+ import logging
37
+ import os
38
+ import random
39
+ import sys
40
+ import time
41
+ from dataclasses import dataclass
42
+ from typing import Any
43
+
44
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
45
+ log = logging.getLogger("securecoder")
46
+
47
+
48
+ # --------------------------------------------------------------------------
49
+ # Data mix
50
+ # --------------------------------------------------------------------------
51
+ @dataclass
52
+ class Source:
53
+ """One dataset feeding the mix.
54
+
55
+ ``kind`` selects the converter ('auto' sniffs columns). ``limit`` is how
56
+ many rows are taken - the sources differ wildly in size, so the cap *is*
57
+ the recipe. Adjust the numbers, not the code.
58
+ """
59
+
60
+ repo: str
61
+ limit: int
62
+ kind: str = "auto"
63
+ config: str | None = None
64
+ split: str = "train"
65
+ note: str = ""
66
+
67
+
68
+ MIX: list[Source] = [
69
+ # ---- tool calling ----------------------------------------------------
70
+ Source("NousResearch/hermes-function-calling-v1", 9000, "tools", "func_calling",
71
+ note="Hermes FC: conversations + JSON tool schemas"),
72
+ Source("NousResearch/hermes-function-calling-v1", 3000, "tools", "func_calling_singleturn",
73
+ note="single-turn tool selection"),
74
+ Source("lockon/xlam-function-calling-60k", 10000, "xlam", "dataset",
75
+ note="xLAM: query/answers/tools API-call pairs"),
76
+ # ---- coding ----------------------------------------------------------
77
+ Source("ise-uiuc/Magicoder-OSS-Instruct-75K", 10000, "magicoder",
78
+ note="self-instruct code problems + solutions"),
79
+ # ---- cybersecurity ---------------------------------------------------
80
+ Source("Trendyol/Trendyol-Cybersecurity-Instruction-Tuning-Dataset", 8000, "sua",
81
+ note="security instruction tuning"),
82
+ Source("AlicanKiraz0/Cybersecurity-Dataset-Fenrir-v2.1", 5000, "sua",
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"),
90
+ Source("TrueNix/ctf-solver-dataset", 3000, "messages",
91
+ note="CTF solving trajectories"),
92
+ # ---- capability replay ----------------------------------------------
93
+ Source("mlabonne/FineTome-100k", 3000, "messages",
94
+ note="general instruct replay so chat ability does not drift"),
95
+ ]
96
+
97
+
98
+ def source_name(src: Source) -> str:
99
+ return src.repo + (f" [{src.config}]" if src.config else "")
100
+
101
+ # --------------------------------------------------------------------------
102
+ # Schema sniffing -> OpenAI-style chat messages
103
+ # --------------------------------------------------------------------------
104
+ ROLE_ALIASES = {
105
+ "human": "user", "user": "user", "gpt": "assistant", "assistant": "assistant",
106
+ "system": "system", "tool": "tool", "function": "tool", "function_call": "tool",
107
+ "observation": "tool", "chatgpt": "assistant",
108
+ }
109
+
110
+
111
+ def _as_list(value: Any) -> list:
112
+ """Accept a JSON string or an already-parsed list."""
113
+ if value is None:
114
+ return []
115
+ if isinstance(value, list):
116
+ return value
117
+ if isinstance(value, str):
118
+ try:
119
+ parsed = json.loads(value)
120
+ except json.JSONDecodeError:
121
+ return []
122
+ return parsed if isinstance(parsed, list) else [parsed]
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"}}."""
129
+ if not isinstance(tool, dict):
130
+ return None
131
+ fn = tool.get("function") if "function" in tool else tool
132
+ if not isinstance(fn, dict) or not fn.get("name"):
133
+ return None
134
+ return {
135
+ "type": "function",
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):
147
+ text = value.strip()
148
+ if not text.startswith(("{", "[")):
149
+ return None
150
+ try:
151
+ value = json.loads(text)
152
+ except json.JSONDecodeError:
153
+ return None
154
+ items = value if isinstance(value, list) else [value]
155
+ if not items or not all(isinstance(i, dict) and "name" in i for i in items):
156
+ return None
157
+ calls = []
158
+ for i, item in enumerate(items):
159
+ arguments = item.get("arguments", item.get("parameters", {}))
160
+ if not isinstance(arguments, str):
161
+ arguments = json.dumps(arguments)
162
+ calls.append({
163
+ "id": f"call_{i}",
164
+ "type": "function",
165
+ "function": {"name": item["name"], "arguments": arguments},
166
+ })
167
+ return calls
168
+
169
+
170
+ def _tools_from_row(row: dict) -> list[dict]:
171
+ raw = row.get("tools")
172
+ candidates = [raw] if isinstance(raw, dict) else _as_list(raw)
173
+ tools = []
174
+ for candidate in candidates:
175
+ norm = _normalise_tool_schema(candidate)
176
+ if norm:
177
+ tools.append(norm)
178
+ return tools
179
+
180
+ def _messages_from_any(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
181
+ """Convert one dataset row into (messages, tools).
182
+
183
+ Returns empty messages when a row cannot be converted confidently; the
184
+ caller counts those, so a silent schema change shows up in the logs instead
185
+ of quietly training on nothing.
186
+ """
187
+ tools = _tools_from_row(row)
188
+ messages: list[dict] = []
189
+
190
+ # --- explicit chat formats (Hermes, OWASP SFT, ctf-solver, FineTome) ---
191
+ if isinstance(row.get("conversations"), list) or isinstance(row.get("messages"), list):
192
+ for turn in row.get("conversations") or row.get("messages") or []:
193
+ if not isinstance(turn, dict):
194
+ continue
195
+ role = ROLE_ALIASES.get(str(turn.get("role") or turn.get("from") or "").lower())
196
+ if role is None:
197
+ continue
198
+ content = turn.get("content", turn.get("value", ""))
199
+ if role == "assistant":
200
+ calls = _parse_calls(content)
201
+ if calls:
202
+ messages.append({"role": "assistant", "content": None, "tool_calls": calls})
203
+ continue
204
+ if role == "tool":
205
+ if not isinstance(content, str):
206
+ content = json.dumps(content)
207
+ messages.append({"role": "tool", "content": content,
208
+ "tool_call_id": turn.get("tool_call_id", "call_0")})
209
+ continue
210
+ if not isinstance(content, str):
211
+ content = json.dumps(content) if content is not None else ""
212
+ if content.strip():
213
+ messages.append({"role": role, "content": content})
214
+ return messages, tools
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})
222
+ return messages, tools
223
+
224
+ # --- system / user / assistant ----------------------------------------
225
+ if row.get("user") and row.get("assistant"):
226
+ if row.get("system"):
227
+ messages.append({"role": "system", "content": str(row["system"])})
228
+ messages.append({"role": "user", "content": str(row["user"])})
229
+ calls = _parse_calls(row["assistant"])
230
+ if calls:
231
+ messages.append({"role": "assistant", "content": None, "tool_calls": calls})
232
+ else:
233
+ messages.append({"role": "assistant", "content": str(row["assistant"])})
234
+ return messages, tools
235
+
236
+ # --- instruction / output (CTF-Instruct) ------------------------------
237
+ if row.get("instruction") and (row.get("output") or row.get("response")):
238
+ user = str(row["instruction"])
239
+ if row.get("input"):
240
+ user = f"{user}\n\n{row['input']}"
241
+ messages.append({"role": "user", "content": user})
242
+ messages.append({"role": "assistant",
243
+ "content": str(row.get("output") or row.get("response"))})
244
+ return messages, tools
245
+
246
+ # --- Magicoder: problem / solution ------------------------------------
247
+ if row.get("problem") and row.get("solution"):
248
+ messages.append({"role": "user", "content":
249
+ "You are an exceptionally intelligent coding assistant that consistently "
250
+ "delivers reliable and accurate responses.\n\n" + str(row["problem"])})
251
+ messages.append({"role": "assistant", "content": str(row["solution"])})
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):
259
+ ask.append(f"{key}: {row[key]}")
260
+ ask.append("Explain what the output means, what it tells you about the target, "
261
+ "and what the next step should be.")
262
+ messages.append({"role": "user", "content": "\n".join(str(a) for a in ask)})
263
+ messages.append({"role": "assistant", "content": str(row["unified_interpretation"])})
264
+ return messages, tools
265
+
266
+ # --- generic prompt/completion fallback -------------------------------
267
+ for pkey, ckey in (("prompt", "completion"), ("question", "answer"), ("input", "output")):
268
+ if row.get(pkey) and row.get(ckey):
269
+ messages.append({"role": "user", "content": str(row[pkey])})
270
+ messages.append({"role": "assistant", "content": str(row[ckey])})
271
+ return messages, tools
272
+
273
+ return [], tools
274
+
275
+ # --------------------------------------------------------------------------
276
+ # Loading, rendering, dataset construction
277
+ # --------------------------------------------------------------------------
278
+ def load_source(src: Source, token: str | None, progress: bool = False) -> list[dict]:
279
+ """Pull up to ``limit`` rows from one Hub dataset, streaming so we never
280
+ download more than we need."""
281
+ from datasets import load_dataset
282
+
283
+ kwargs: dict[str, Any] = {"split": src.split, "streaming": True}
284
+ if src.config:
285
+ kwargs["name"] = src.config
286
+ if token:
287
+ kwargs["token"] = token
288
+
289
+ ds = load_dataset(src.repo, **kwargs)
290
+ rows = []
291
+ for i, row in enumerate(ds):
292
+ if i >= src.limit:
293
+ break
294
+ rows.append(dict(row))
295
+ if progress and i and i % 2500 == 0:
296
+ log.info(" %s: %d rows...", source_name(src), i)
297
+ return rows
298
+
299
+
300
+ def render_record(tokenizer, messages: list[dict], tools: list[dict] | None = None) -> str:
301
+ """Render to the model's native chat format (Qwen3 emits <tool_call> blocks)."""
302
+ return tokenizer.apply_chat_template(
303
+ messages,
304
+ tools=tools or None,
305
+ tokenize=False,
306
+ add_generation_prompt=False,
307
+ )
308
+
309
+
310
+ def build_dataset(tokenizer, sources: list[Source], token: str | None, validate: bool):
311
+ """Returns (records, stats). ``records`` are {"text", "source"} dicts."""
312
+ records: list[dict] = []
313
+ stats: list[dict] = []
314
+
315
+ for src in sources:
316
+ entry = {"source": source_name(src), "note": src.note, "kept": 0, "skipped": 0,
317
+ "tool_samples": 0, "chars": 0, "error": None}
318
+ try:
319
+ rows = load_source(src, token, progress=validate)
320
+ for row in rows:
321
+ messages, tools = _messages_from_any(row, src.kind)
322
+ has_answer = any(
323
+ (m.get("content") or m.get("tool_calls")) for m in messages
324
+ if m["role"] == "assistant"
325
+ )
326
+ if not messages or not has_answer:
327
+ entry["skipped"] += 1
328
+ continue
329
+ try:
330
+ text = render_record(tokenizer, messages, tools)
331
+ except Exception as exc: # noqa: BLE001 - bad template input, skip row
332
+ if entry["skipped"] < 3:
333
+ log.warning(" render failed (%s): %s", source_name(src), exc)
334
+ entry["skipped"] += 1
335
+ continue
336
+ if len(text) < 40 or len(text) > 120_000:
337
+ entry["skipped"] += 1
338
+ continue
339
+ records.append({"text": text, "source": source_name(src)})
340
+ entry["kept"] += 1
341
+ entry["chars"] += len(text)
342
+ if tools:
343
+ entry["tool_samples"] += 1
344
+ except Exception as exc: # noqa: BLE001 - one bad dataset must not kill the run
345
+ entry["error"] = repr(exc)
346
+ log.error(" %s failed: %s", source_name(src), exc)
347
+
348
+ stats.append(entry)
349
+ log.info(" %-58s kept=%-6d skipped=%-5d tools=%-5d",
350
+ entry["source"], entry["kept"], entry["skipped"], entry["tool_samples"])
351
+
352
+ random.shuffle(records)
353
+ return records, stats
354
+
355
+
356
+ def print_stats(stats: list[dict], records: list[dict], tokenizer=None) -> None:
357
+ total_chars = sum(r["chars"] for r in stats if not r["error"])
358
+ print("\n" + "=" * 78)
359
+ print("DATA MIX")
360
+ print("=" * 78)
361
+ print(f"{'source':<58}{'kept':>7}{'skip':>7}{'tools':>7}")
362
+ for s in stats:
363
+ print(f"{s['source']:<58}{s['kept']:>7}{s['skipped']:>7}{s['tool_samples']:>7}")
364
+ if s["error"]:
365
+ print(f" !! {s['error'][:120]}")
366
+ tool_rows = sum(s["tool_samples"] for s in stats)
367
+ print("-" * 78)
368
+ print(f"total rows : {len(records):,}")
369
+ print(f"tool rows : {tool_rows:,} ({100 * tool_rows / max(len(records), 1):.1f}%)")
370
+ print(f"total chars: {total_chars:,} (~{total_chars // 4:,} tokens)")
371
+
372
+ # --------------------------------------------------------------------------
373
+ # Model + training
374
+ # --------------------------------------------------------------------------
375
+ # Attention + router only by default: on a 128-expert MoE, adapting every expert
376
+ # MLP means ~800M trainable parameters, which dominates VRAM and step time.
377
+ # Pass --target-modules all-linear when you want the MLP/expert capacity too.
378
+ ATTENTION_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj", "gate"]
379
+
380
+
381
+ def load_model_and_tokenizer(args):
382
+ from unsloth import FastLanguageModel
383
+
384
+ model, tokenizer = FastLanguageModel.from_pretrained(
385
+ model_name=args.base_model,
386
+ max_seq_length=args.max_seq_length,
387
+ dtype=None,
388
+ load_in_4bit=not args.no_4bit,
389
+ )
390
+
391
+ targets: Any = args.target_modules
392
+ if isinstance(targets, str) and targets != "all-linear":
393
+ targets = [t.strip() for t in targets.split(",") if t.strip()]
394
+
395
+ model = FastLanguageModel.get_peft_model(
396
+ model,
397
+ r=args.lora_r,
398
+ target_modules=targets,
399
+ lora_alpha=args.lora_alpha,
400
+ lora_dropout=0.0,
401
+ bias="none",
402
+ use_gradient_checkpointing="unsloth",
403
+ random_state=args.seed,
404
+ use_rslora=False,
405
+ )
406
+ trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
407
+ log.info("trainable parameters: %s (%.2f%% of the model)",
408
+ f"{trainable:,}", 100 * trainable / max(sum(p.numel() for p in model.parameters()), 1))
409
+ return model, tokenizer
410
+
411
+
412
+ def make_sft_config(**kwargs):
413
+ """TRL renamed max_seq_length -> max_length; support both."""
414
+ from trl import SFTConfig
415
+
416
+ try:
417
+ return SFTConfig(max_length=kwargs.pop("max_seq_length"), **kwargs)
418
+ except TypeError:
419
+ kwargs["max_seq_length"] = kwargs.get("max_seq_length")
420
+ return SFTConfig(**kwargs)
421
+
422
+
423
+ def build_sft_config(args, has_eval: bool, steps_per_epoch: int | None):
424
+ import torch
425
+
426
+ bf16 = torch.cuda.is_bf16_supported()
427
+ cfg: dict[str, Any] = dict(
428
+ output_dir=args.output_dir,
429
+ per_device_train_batch_size=args.batch_size,
430
+ gradient_accumulation_steps=args.grad_accum,
431
+ warmup_ratio=0.03,
432
+ learning_rate=args.learning_rate,
433
+ max_grad_norm=1.0,
434
+ weight_decay=0.01,
435
+ lr_scheduler_type=args.lr_scheduler,
436
+ optim="adamw_8bit",
437
+ logging_steps=args.logging_steps,
438
+ save_steps=args.save_steps,
439
+ save_total_limit=2,
440
+ seed=args.seed,
441
+ bf16=bf16,
442
+ fp16=not bf16,
443
+ max_seq_length=args.max_seq_length,
444
+ dataset_text_field="text",
445
+ packing=args.packing,
446
+ report_to=args.report_to,
447
+ run_name=args.run_name or "securecoder",
448
+ remove_unused_columns=False,
449
+ )
450
+ if args.max_steps > 0:
451
+ cfg["max_steps"] = args.max_steps
452
+ else:
453
+ cfg["num_train_epochs"] = args.num_epochs
454
+
455
+ if has_eval:
456
+ cfg["eval_strategy"] = "steps"
457
+ cfg["eval_steps"] = args.save_steps
458
+ cfg["per_device_eval_batch_size"] = 1
459
+ cfg["do_eval"] = True
460
+ return make_sft_config(**cfg)
461
+
462
+
463
+ def init_trackio(args):
464
+ if args.report_to == "none":
465
+ return
466
+ try:
467
+ import trackio
468
+
469
+ if args.trackio_space:
470
+ trackio.init(project=args.trackio_project, space_id=args.trackio_space)
471
+ else:
472
+ trackio.init(project=args.trackio_project)
473
+ log.info("trackio initialised (project=%s space=%s)", args.trackio_project, args.trackio_space)
474
+ except Exception as exc: # noqa: BLE001 - monitoring must never kill training
475
+ log.warning("trackio init failed (%s); continuing without it", exc)
476
+
477
+
478
+ def train(args, model, tokenizer, records: list[dict]):
479
+ from datasets import Dataset
480
+ from trl import SFTTrainer
481
+
482
+ random.Random(args.seed).shuffle(records)
483
+ split_at = len(records) - args.eval_samples if args.eval_samples > 0 else len(records)
484
+ train_ds = Dataset.from_list(records[:split_at])
485
+ eval_ds = Dataset.from_list(records[split_at:]) if args.eval_samples > 0 else None
486
+ log.info("train rows=%d eval rows=%d", len(train_ds), len(eval_ds) if eval_ds else 0)
487
+
488
+ steps_per_epoch = len(train_ds) // max(args.batch_size * args.grad_accum, 1)
489
+ cfg = build_sft_config(args, eval_ds is not None, steps_per_epoch)
490
+ init_trackio(args)
491
+
492
+ trainer = SFTTrainer(
493
+ model=model,
494
+ tokenizer=tokenizer,
495
+ train_dataset=train_ds,
496
+ eval_dataset=eval_ds,
497
+ args=cfg,
498
+ )
499
+
500
+ started = time.time()
501
+ stats = trainer.train()
502
+ elapsed = time.time() - started
503
+ log.info("training finished in %.1f min (final loss %.4f)",
504
+ elapsed / 60, stats.metrics.get("train_loss", float("nan")))
505
+
506
+ if eval_ds is not None:
507
+ try:
508
+ metrics = trainer.evaluate()
509
+ log.info("eval_loss %.4f (train %.4f)",
510
+ metrics.get("eval_loss", float("nan")),
511
+ stats.metrics.get("train_loss", float("nan")))
512
+ except Exception as exc: # noqa: BLE001
513
+ log.warning("eval failed: %s", exc)
514
+
515
+ return trainer, stats, elapsed
516
+
517
+
518
+ def save_and_push(args, model, tokenizer):
519
+ from huggingface_hub import HfApi
520
+
521
+ api = HfApi()
522
+ api.create_repo(args.output_repo, repo_type="model", exist_ok=True, private=args.private)
523
+
524
+ model.save_pretrained(args.output_dir)
525
+ tokenizer.save_pretrained(args.output_dir)
526
+ log.info("pushing LoRA adapter to %s", args.output_repo)
527
+ model.push_to_hub(args.output_repo, tokenizer=tokenizer)
528
+
529
+ if args.merge_repo:
530
+ api.create_repo(args.merge_repo, repo_type="model", exist_ok=True, private=args.private)
531
+ log.info("merging to 16-bit and pushing to %s (large upload)", args.merge_repo)
532
+ model.push_to_hub_merged(args.merge_repo, tokenizer=tokenizer, save_method="merged_16bit")
533
+
534
+ # --------------------------------------------------------------------------
535
+ # CLI
536
+ # --------------------------------------------------------------------------
537
+ def parse_args(argv=None):
538
+ p = argparse.ArgumentParser(description="SecureCoder QLoRA fine-tune")
539
+
540
+ p.add_argument("--base-model", default="Qwen/Qwen3-Coder-30B-A3B-Instruct")
541
+ p.add_argument("--output-repo", default=None, help="Hub repo for the LoRA adapter")
542
+ p.add_argument("--merge-repo", default=None, help="optional Hub repo for a 16-bit merge")
543
+ p.add_argument("--output-dir", default="securecoder-out")
544
+ p.add_argument("--private", action="store_true", help="create Hub repos as private")
545
+
546
+ p.add_argument("--max-seq-length", type=int, default=4096)
547
+ p.add_argument("--batch-size", type=int, default=2)
548
+ p.add_argument("--grad-accum", type=int, default=8)
549
+ p.add_argument("--learning-rate", type=float, default=2e-4)
550
+ p.add_argument("--lr-scheduler", default="cosine")
551
+ p.add_argument("--num-epochs", type=float, default=1.0)
552
+ p.add_argument("--max-steps", type=int, default=0, help="overrides --num-epochs when > 0")
553
+ p.add_argument("--eval-samples", type=int, default=200, help="0 disables evaluation")
554
+ p.add_argument("--logging-steps", type=int, default=10)
555
+ p.add_argument("--save-steps", type=int, default=250)
556
+ p.add_argument("--packing", action="store_true", default=True)
557
+ p.add_argument("--no-packing", dest="packing", action="store_false")
558
+ p.add_argument("--seed", type=int, default=3407)
559
+
560
+ p.add_argument("--lora-r", type=int, default=32)
561
+ p.add_argument("--lora-alpha", type=int, default=32)
562
+ p.add_argument("--no-4bit", action="store_true")
563
+ p.add_argument("--target-modules", default=",".join(ATTENTION_TARGETS),
564
+ help="comma-separated suffixes, or 'all-linear' to include expert MLPs")
565
+
566
+ p.add_argument("--report-to", default="trackio", choices=["trackio", "none"])
567
+ p.add_argument("--trackio-project", default="securecoder")
568
+ p.add_argument("--trackio-space", default=None, help="e.g. Taimwe/securecoder-trackio")
569
+ p.add_argument("--run-name", default=None)
570
+
571
+ p.add_argument("--validate-only", action="store_true",
572
+ help="load a small sample of each source, print the mix, exit (no GPU)")
573
+ p.add_argument("--validate-per-source", type=int, default=40)
574
+ p.add_argument("--show-samples", type=int, default=3)
575
+ p.add_argument("--smoke", action="store_true",
576
+ help="tiny end-to-end run: 200 rows/source, 20 steps")
577
+ return p.parse_args(argv)
578
+
579
+
580
+ def apply_smoke(args) -> None:
581
+ args.max_steps = args.max_steps or 20
582
+ args.eval_samples = min(args.eval_samples, 20)
583
+ args.save_steps = 20
584
+ args.max_seq_length = min(args.max_seq_length, 2048)
585
+ global MIX
586
+ MIX = [Source(s.repo, 200, s.kind, s.config, s.split, s.note) for s in MIX]
587
+
588
+
589
+ def main(argv=None) -> int:
590
+ args = parse_args(argv)
591
+ if args.smoke:
592
+ apply_smoke(args)
593
+ token = os.environ.get("HF_TOKEN")
594
+
595
+ if args.validate_only:
596
+ from transformers import AutoTokenizer
597
+
598
+ tokenizer = AutoTokenizer.from_pretrained(args.base_model)
599
+ sources = [Source(s.repo, args.validate_per_source, s.kind, s.config, s.split, s.note)
600
+ for s in MIX]
601
+ records, stats = build_dataset(tokenizer, sources, token, validate=True)
602
+ print_stats(stats, records, tokenizer)
603
+ for i, rec in enumerate(records[: args.show_samples], 1):
604
+ print("\n" + "-" * 78)
605
+ print(f"SAMPLE {i} [{rec['source']}] {len(rec['text'])} chars")
606
+ print("-" * 78)
607
+ print(rec["text"][:1500])
608
+ return 0
609
+
610
+ import torch
611
+
612
+ if not torch.cuda.is_available():
613
+ log.error("no CUDA device - use --validate-only locally, or run on HF Jobs / Colab")
614
+ return 1
615
+ log.info("GPU: %s", torch.cuda.get_device_name(0))
616
+
617
+ if not args.output_repo:
618
+ log.error("--output-repo is required (the container/VM is ephemeral)")
619
+ return 1
620
+
621
+ model, tokenizer = load_model_and_tokenizer(args)
622
+ records, stats = build_dataset(tokenizer, MIX, token, validate=False)
623
+ print_stats(stats, records, tokenizer)
624
+ if len(records) < 100:
625
+ log.error("only %d usable rows - refusing to train", len(records))
626
+ return 1
627
+
628
+ trainer, stats_train, elapsed = train(args, model, tokenizer, records)
629
+ save_and_push(args, model, tokenizer)
630
+
631
+ print("\n" + "=" * 78)
632
+ print(f"DONE rows={len(records):,} time={elapsed / 60:.1f} min "
633
+ f"loss={stats_train.metrics.get('train_loss', float('nan')):.4f}")
634
+ print(f"adapter: https://huggingface.co/{args.output_repo}")
635
+ if args.merge_repo:
636
+ print(f"merged : https://huggingface.co/{args.merge_repo}")
637
+ print("=" * 78)
638
+ return 0
639
+
640
+
641
+ if __name__ == "__main__":
642
+ raise SystemExit(main())