securecoder-scripts / train_securecoder.py
Taimwe's picture
Format-adaptive tool rendering + resilient dataset loading
fffea22 verified
Raw History Blame
35.3 kB
# /// script
# requires-python = ">=3.10"
# dependencies = [
# "unsloth",
# "datasets",
# "trl>=0.22",
# "transformers>=4.57",
# "trackio",
# "huggingface_hub",
# ]
# ///
"""SecureCoder: QLoRA fine-tune for code + tool calling + cybersecurity.
Default base: Qwen/Qwen3-Coder-30B-A3B-Instruct (Apache-2.0, 30.5B MoE, ~3B
active) - a MoE that trains like a small model and runs like a useful one.
Runs anywhere; same file for local validation, a GPU smoke test, and the real
run:
uv run train_securecoder.py --validate-only # no GPU needed
uv run train_securecoder.py --smoke --output-repo you/securecoder-smoke
uv run train_securecoder.py --num-epochs 1 --output-repo you/securecoder-30b-pro
Launch on Hugging Face Jobs (see README.md for why the URL form is used):
hf jobs run -d --flavor l40sx1 --timeout 12h --secrets HF_TOKEN \\
ghcr.io/astral-sh/uv:python3.12-bookworm \\
uv run --no-project https://huggingface.co/USER/securecoder-scripts/resolve/main/train_securecoder.py \\
-- --num-epochs 1 --output-repo USER/securecoder-30b-pro
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import random
import sys
import time
from dataclasses import dataclass
from typing import Any
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("securecoder")
# --------------------------------------------------------------------------
# Data mix
# --------------------------------------------------------------------------
@dataclass
class Source:
"""One dataset feeding the mix.
``kind`` selects the converter ('auto' sniffs columns). ``limit`` is how
many rows are taken - the sources differ wildly in size, so the cap *is*
the recipe. Adjust the numbers, not the code.
"""
repo: str
limit: int
kind: str = "auto"
config: str | None = None
split: str = "train"
note: str = ""
MIX: list[Source] = [
# ---- tool calling ----------------------------------------------------
Source("NousResearch/hermes-function-calling-v1", 9000, "tools", "func_calling",
note="Hermes FC: conversations + JSON tool schemas"),
Source("NousResearch/hermes-function-calling-v1", 3000, "tools", "func_calling_singleturn",
note="single-turn tool selection"),
Source("lockon/xlam-function-calling-60k", 10000, "xlam", "dataset",
note="xLAM: query/answers/tools API-call pairs"),
# ---- coding ----------------------------------------------------------
Source("ise-uiuc/Magicoder-OSS-Instruct-75K", 10000, "magicoder",
note="self-instruct code problems + solutions"),
# ---- cybersecurity ---------------------------------------------------
Source("Trendyol/Trendyol-Cybersecurity-Instruction-Tuning-Dataset", 8000, "sua",
note="security instruction tuning"),
Source("AlicanKiraz0/Cybersecurity-Dataset-Fenrir-v2.1", 5000, "sua",
note="broad security Q&A"),
Source("Humanlearning/CyberSecurity_OWASP-sft-dataset", 3000, "messages",
note="OWASP / secure-coding SFT"),
Source("dpevzner/Cybersecurity_Reasoning_Dataset", 3000, "secops", "chatml", "train",
note="command interpretation reasoning (goal -> unified_interpretation)"),
Source("MrClipperz134/CTF-Instruct", 3000, "io",
note="CTF instruction/output"),
Source("TrueNix/ctf-solver-dataset", 3000, "messages",
note="CTF solving trajectories"),
# ---- capability replay ----------------------------------------------
Source("mlabonne/FineTome-100k", 3000, "messages",
note="general instruct replay so chat ability does not drift"),
]
def source_name(src: Source) -> str:
return src.repo + (f" [{src.config}]" if src.config else "")
# --------------------------------------------------------------------------
# Schema sniffing -> OpenAI-style chat messages
# --------------------------------------------------------------------------
ROLE_ALIASES = {
"human": "user", "user": "user", "gpt": "assistant", "assistant": "assistant",
"system": "system", "tool": "tool", "function": "tool", "function_call": "tool",
"observation": "tool", "chatgpt": "assistant",
}
def _as_list(value: Any) -> list:
"""Accept a JSON string or an already-parsed list."""
if value is None:
return []
if isinstance(value, list):
return value
if isinstance(value, str):
try:
parsed = json.loads(value)
except json.JSONDecodeError:
return []
return parsed if isinstance(parsed, list) else [parsed]
return []
TYPE_ALIASES = {
"str": "string", "string": "string", "text": "string",
"int": "integer", "integer": "integer", "long": "integer",
"float": "number", "double": "number", "number": "number",
"bool": "boolean", "boolean": "boolean",
"list": "array", "array": "array", "dict": "object", "object": "object",
}
def _normalise_parameters(params: Any) -> dict:
"""Coerce a tool's parameter spec into valid JSON Schema.
Two shapes appear in the wild: proper ``{"type": "object", "properties": {}}``
and the flat ``{"arg": {"description": ..., "type": "str"}}`` form used by
xLAM. Qwen's chat template reads ``parameters.properties``, so the flat form
has to be wrapped or rendering raises.
"""
if not isinstance(params, dict) or not params:
return {"type": "object", "properties": {}}
if "properties" in params:
params.setdefault("type", "object")
return params
properties: dict[str, Any] = {}
required: list[str] = []
for name, spec in params.items():
if isinstance(spec, dict):
cleaned = {
k: v for k, v in spec.items()
if k in ("type", "description", "enum", "default", "title", "items")
}
cleaned["type"] = TYPE_ALIASES.get(str(cleaned.get("type", "")).lower(), "string")
properties[name] = cleaned
if "default" not in spec:
required.append(name)
else:
properties[name] = {"type": "string"}
required.append(name)
schema: dict[str, Any] = {"type": "object", "properties": properties}
if required:
schema["required"] = required
return schema
def _normalise_tool_schema(tool: Any) -> dict | None:
"""Canonical **flat** tool schema: {"name", "description", "parameters"}.
Qwen3-Coder's chat template walks ``tool.parameters.properties`` directly, so
the flat form is what we store; ``_tools_for_style`` re-wraps it into the
OpenAI ``{"type": "function", "function": {...}}`` shape for templates that
want that instead.
"""
if not isinstance(tool, dict):
return None
fn = tool.get("function") if isinstance(tool.get("function"), dict) else tool
if not isinstance(fn, dict) or not fn.get("name"):
return None
return {
"name": fn["name"],
"description": fn.get("description", ""),
"parameters": _normalise_parameters(fn.get("parameters")),
}
def _tools_for_style(tools: list[dict], style: str) -> list[dict] | None:
if not tools:
return None
if style == "nested":
return [{"type": "function", "function": t} for t in tools]
return tools
def _parse_calls(value: Any) -> list[dict] | None:
"""Return OpenAI-style tool calls if ``value`` is one or more function calls.
``arguments`` is kept as a **dict** plus a JSON-string copy, because templates
disagree: Qwen3-Coder iterates ``arguments | items`` (needs a mapping) while
others print a JSON string.
"""
if isinstance(value, str):
text = value.strip()
if not text.startswith(("{", "[")):
return None
try:
value = json.loads(text)
except json.JSONDecodeError:
return None
items = value if isinstance(value, list) else [value]
if not items or not all(isinstance(i, dict) and "name" in i for i in items):
return None
calls = []
for i, item in enumerate(items):
arguments = item.get("arguments", item.get("parameters", {}))
if isinstance(arguments, str):
try:
arguments = json.loads(arguments)
except json.JSONDecodeError:
arguments = {"value": arguments}
if not isinstance(arguments, dict):
arguments = {"value": arguments}
calls.append({
"id": f"call_{i}",
"type": "function",
"function": {
"name": item["name"],
"arguments": arguments,
"arguments_json": json.dumps(arguments),
},
})
return calls
def _tools_from_row(row: dict) -> list[dict]:
raw = row.get("tools")
candidates = [raw] if isinstance(raw, dict) else _as_list(raw)
tools = []
for candidate in candidates:
norm = _normalise_tool_schema(candidate)
if norm:
tools.append(norm)
return tools
def _messages_from_any(row: dict, kind: str) -> tuple[list[dict], list[dict]]:
"""Convert one dataset row into (messages, tools).
Returns empty messages when a row cannot be converted confidently; the
caller counts those, so a silent schema change shows up in the logs instead
of quietly training on nothing.
"""
tools = _tools_from_row(row)
messages: list[dict] = []
# --- explicit chat formats (Hermes, OWASP SFT, ctf-solver, FineTome) ---
if isinstance(row.get("conversations"), list) or isinstance(row.get("messages"), list):
for turn in row.get("conversations") or row.get("messages") or []:
if not isinstance(turn, dict):
continue
role = ROLE_ALIASES.get(str(turn.get("role") or turn.get("from") or "").lower())
if role is None:
continue
content = turn.get("content", turn.get("value", ""))
if role == "assistant":
calls = _parse_calls(content)
if calls:
messages.append({"role": "assistant", "content": None, "tool_calls": calls})
continue
if role == "tool":
if not isinstance(content, str):
content = json.dumps(content)
messages.append({"role": "tool", "content": content,
"tool_call_id": turn.get("tool_call_id", "call_0")})
continue
if not isinstance(content, str):
content = json.dumps(content) if content is not None else ""
if content.strip():
messages.append({"role": role, "content": content})
return messages, tools
# --- xLAM: query + answers + tools ------------------------------------
if kind == "xlam" or (row.get("query") and row.get("answers")):
answers = _as_list(row.get("answers"))
calls = _parse_calls(answers)
if calls is None and answers:
# answers can be a list of JSON strings instead of one JSON array
flat: list = []
for item in answers:
flat.extend(_as_list(item))
calls = _parse_calls(flat)
if calls and row.get("query"):
messages.append({"role": "user", "content": str(row["query"])})
messages.append({"role": "assistant", "content": None, "tool_calls": calls})
return messages, tools
# --- system / user / assistant ----------------------------------------
if row.get("user") and row.get("assistant"):
if row.get("system"):
messages.append({"role": "system", "content": str(row["system"])})
messages.append({"role": "user", "content": str(row["user"])})
calls = _parse_calls(row["assistant"])
if calls:
messages.append({"role": "assistant", "content": None, "tool_calls": calls})
else:
messages.append({"role": "assistant", "content": str(row["assistant"])})
return messages, tools
# --- instruction / output (CTF-Instruct) ------------------------------
if row.get("instruction") and (row.get("output") or row.get("response")):
user = str(row["instruction"])
if row.get("input"):
user = f"{user}\n\n{row['input']}"
messages.append({"role": "user", "content": user})
messages.append({"role": "assistant",
"content": str(row.get("output") or row.get("response"))})
return messages, tools
# --- Magicoder: problem / solution ------------------------------------
if row.get("problem") and row.get("solution"):
messages.append({"role": "user", "content":
"You are an exceptionally intelligent coding assistant that consistently "
"delivers reliable and accurate responses.\n\n" + str(row["problem"])})
messages.append({"role": "assistant", "content": str(row["solution"])})
return messages, tools
# --- SecOps reasoning: goal/command -> interpretation ------------------
# dpevzner's rows carry goal + command_sequence + interpretation +
# classification + safety_and_scope (unified_interpretation is empty in the
# published revision). Scope metadata goes into the prompt so the model
# learns the framed, authorised-use context alongside the command knowledge.
if row.get("goal") and row.get("interpretation"):
seq = row.get("command_sequence") or {}
if isinstance(seq, str):
seq = {"command": seq}
interp = row.get("interpretation") or {}
if isinstance(interp, str):
interp = {"what_it_means": [interp]}
scope = row.get("safety_and_scope") or {}
if isinstance(scope, str):
scope = {}
tool = row.get("tool") or {}
if isinstance(tool, str):
tool = {"name": tool}
ask = [f"Environment: {tool.get('name', 'shell')} ({tool.get('platform', 'unknown')})"]
if seq.get("command"):
ask.append(f"Command: {seq['command']}")
ask.append(f"Goal: {row['goal']}")
if isinstance(scope, dict) and scope:
ask.append("Scope: " + ", ".join(f"{k}={v}" for k, v in list(scope.items())[:4]))
if row.get("classification"):
cls = row["classification"]
if isinstance(cls, dict):
ask.append("Context: " + ", ".join(f"{k}={v}" for k, v in list(cls.items())[:3]))
ask.append("Explain what this command does, what its output means, what it tells you "
"about the target, and the next step in an authorised assessment.")
answer = []
if seq.get("description"):
answer.append(f"**What it does.** {seq['description']}")
if seq.get("expected_output_pattern"):
answer.append("**Expected output.** " + ", ".join(map(str, seq["expected_output_pattern"])))
for item in interp.get("what_it_means", []) if isinstance(interp, dict) else []:
answer.append(f"**What it means.** {item}")
for item in interp.get("risk_indicators", []) if isinstance(interp, dict) else []:
answer.append(f"**Risk indicators.** {item}")
if row.get("ambiguity_analysis"):
answer.append(f"**Ambiguity.** {row['ambiguity_analysis']}")
if isinstance(scope, dict) and scope.get("authorization_required"):
answer.append("**Scope.** Only run this against systems you are authorised to test.")
if len(answer) < 2:
return [], tools
messages.append({"role": "user", "content": "\n".join(str(a) for a in ask)})
messages.append({"role": "assistant", "content": "\n\n".join(answer)})
return messages, tools
# --- generic prompt/completion fallback -------------------------------
for pkey, ckey in (("prompt", "completion"), ("question", "answer"), ("input", "output")):
if row.get(pkey) and row.get(ckey):
messages.append({"role": "user", "content": str(row[pkey])})
messages.append({"role": "assistant", "content": str(row[ckey])})
return messages, tools
return [], tools
# --------------------------------------------------------------------------
# Loading, rendering, dataset construction
# --------------------------------------------------------------------------
def load_source(src: Source, token: str | None, progress: bool = False) -> list[dict]:
"""Pull up to ``limit`` rows from one Hub dataset, streaming so we never
download more than we need.
Datasets move: builder configs get renamed (a config called ``default`` last
week is ``chatml`` today) and splits get added. Each candidate is tried in
turn so one rename cannot silently empty a slice of the mix.
"""
from datasets import load_dataset
candidates = [
(src.config, src.split),
(src.config, "train"),
(src.config, "test"),
(None, src.split),
(None, "train"),
(None, "test"),
]
seen: set = set()
last_exc: Exception | None = None
for config, split in candidates:
if (config, split) in seen:
continue
seen.add((config, split))
kwargs: dict[str, Any] = {"split": split, "streaming": True}
if config:
kwargs["name"] = config
if token:
kwargs["token"] = token
try:
ds = load_dataset(src.repo, **kwargs)
rows = []
for i, row in enumerate(ds):
if i >= src.limit:
break
rows.append(dict(row))
if progress and i and i % 2500 == 0:
log.info(" %s: %d rows...", source_name(src), i)
if not rows:
last_exc = ValueError(f"config={config} split={split} streamed 0 rows")
continue
if (config, split) != (src.config, src.split):
log.info(" %s: fell back to config=%s split=%s",
src.repo, config, split)
return rows
except Exception as exc: # noqa: BLE001 - try the next candidate
last_exc = exc
raise last_exc if last_exc else RuntimeError(f"could not load {src.repo}")
_RENDER_STYLE: str | None = None
STYLE_ATTEMPTS = (("flat", "dict"), ("nested", "dict"), ("flat", "string"), ("nested", "string"))
def _apply_arg_style(messages: list[dict], arg_style: str) -> list[dict]:
"""Copy messages, swapping tool-call arguments between dict and JSON string."""
if arg_style != "string":
return messages
out = []
for message in messages:
if message.get("tool_calls"):
message = dict(message)
message["tool_calls"] = [
{
"id": call["id"],
"type": "function",
"function": {
"name": call["function"]["name"],
"arguments": call["function"].get(
"arguments_json", json.dumps(call["function"]["arguments"])
),
},
}
for call in message["tool_calls"]
]
out.append(message)
return out
def render_record(tokenizer, messages: list[dict], tools: list[dict] | None = None) -> str:
"""Render with the model's native chat template.
Tool templates disagree in two independent ways: whether tool schemas are
flat (`{"name", "parameters"}`) or OpenAI-nested (`{"type": "function", ...}`),
and whether tool-call arguments are a mapping or a JSON string. Qwen3-Coder
renders ``<function=NAME><parameter=...>`` blocks and iterates
``arguments | items``, so a JSON string there is a hard error. Rather than
hard-code one convention, detect it once and reuse it for the rest of the
run - with the model, the mix and the template all free to change.
"""
global _RENDER_STYLE
order = []
if _RENDER_STYLE:
order.append(tuple(_RENDER_STYLE.split("+")))
order += [style for style in STYLE_ATTEMPTS if style not in order]
last_exc: Exception | None = None
for tool_style, arg_style in order:
try:
text = tokenizer.apply_chat_template(
_apply_arg_style(messages, arg_style),
tools=_tools_for_style(tools, tool_style),
tokenize=False,
add_generation_prompt=False,
)
_RENDER_STYLE = f"{tool_style}+{arg_style}"
return text
except Exception as exc: # noqa: BLE001 - try the next convention
last_exc = exc
raise last_exc if last_exc else RuntimeError("render failed")
def build_dataset(tokenizer, sources: list[Source], token: str | None, validate: bool):
"""Returns (records, stats). ``records`` are {"text", "source"} dicts."""
records: list[dict] = []
stats: list[dict] = []
for src in sources:
entry = {"source": source_name(src), "note": src.note, "kept": 0, "skipped": 0,
"tool_samples": 0, "chars": 0, "error": None}
try:
rows = load_source(src, token, progress=validate)
for row in rows:
messages, tools = _messages_from_any(row, src.kind)
has_answer = any(
(m.get("content") or m.get("tool_calls")) for m in messages
if m["role"] == "assistant"
)
if not messages or not has_answer:
entry["skipped"] += 1
continue
try:
text = render_record(tokenizer, messages, tools)
except Exception as exc: # noqa: BLE001 - bad template input, skip row
if entry["skipped"] < 3:
log.warning(" render failed (%s): %s", source_name(src), exc)
entry["skipped"] += 1
continue
if len(text) < 40 or len(text) > 120_000:
entry["skipped"] += 1
continue
records.append({"text": text, "source": source_name(src)})
entry["kept"] += 1
entry["chars"] += len(text)
if tools:
entry["tool_samples"] += 1
except Exception as exc: # noqa: BLE001 - one bad dataset must not kill the run
entry["error"] = repr(exc)
log.error(" %s failed: %s", source_name(src), exc)
stats.append(entry)
log.info(" %-58s kept=%-6d skipped=%-5d tools=%-5d",
entry["source"], entry["kept"], entry["skipped"], entry["tool_samples"])
random.shuffle(records)
return records, stats
def print_stats(stats: list[dict], records: list[dict], tokenizer=None) -> None:
total_chars = sum(r["chars"] for r in stats if not r["error"])
print("\n" + "=" * 78)
print("DATA MIX")
print("=" * 78)
print(f"{'source':<58}{'kept':>7}{'skip':>7}{'tools':>7}")
for s in stats:
print(f"{s['source']:<58}{s['kept']:>7}{s['skipped']:>7}{s['tool_samples']:>7}")
if s["error"]:
print(f" !! {s['error'][:120]}")
tool_rows = sum(s["tool_samples"] for s in stats)
print("-" * 78)
print(f"total rows : {len(records):,}")
print(f"tool rows : {tool_rows:,} ({100 * tool_rows / max(len(records), 1):.1f}%)")
print(f"total chars: {total_chars:,} (~{total_chars // 4:,} tokens)")
# --------------------------------------------------------------------------
# Model + training
# --------------------------------------------------------------------------
# Attention + router only by default: on a 128-expert MoE, adapting every expert
# MLP means ~800M trainable parameters, which dominates VRAM and step time.
# Pass --target-modules all-linear when you want the MLP/expert capacity too.
ATTENTION_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj", "gate"]
def load_model_and_tokenizer(args):
from unsloth import FastLanguageModel
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=args.base_model,
max_seq_length=args.max_seq_length,
dtype=None,
load_in_4bit=not args.no_4bit,
)
targets: Any = args.target_modules
if isinstance(targets, str) and targets != "all-linear":
targets = [t.strip() for t in targets.split(",") if t.strip()]
model = FastLanguageModel.get_peft_model(
model,
r=args.lora_r,
target_modules=targets,
lora_alpha=args.lora_alpha,
lora_dropout=0.0,
bias="none",
use_gradient_checkpointing="unsloth",
random_state=args.seed,
use_rslora=False,
)
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
log.info("trainable parameters: %s (%.2f%% of the model)",
f"{trainable:,}", 100 * trainable / max(sum(p.numel() for p in model.parameters()), 1))
return model, tokenizer
def make_sft_config(**kwargs):
"""TRL renamed max_seq_length -> max_length; support both."""
from trl import SFTConfig
try:
return SFTConfig(max_length=kwargs.pop("max_seq_length"), **kwargs)
except TypeError:
kwargs["max_seq_length"] = kwargs.get("max_seq_length")
return SFTConfig(**kwargs)
def build_sft_config(args, has_eval: bool, steps_per_epoch: int | None):
import torch
bf16 = torch.cuda.is_bf16_supported()
cfg: dict[str, Any] = dict(
output_dir=args.output_dir,
per_device_train_batch_size=args.batch_size,
gradient_accumulation_steps=args.grad_accum,
warmup_ratio=0.03,
learning_rate=args.learning_rate,
max_grad_norm=1.0,
weight_decay=0.01,
lr_scheduler_type=args.lr_scheduler,
optim="adamw_8bit",
logging_steps=args.logging_steps,
save_steps=args.save_steps,
save_total_limit=2,
seed=args.seed,
bf16=bf16,
fp16=not bf16,
max_seq_length=args.max_seq_length,
dataset_text_field="text",
packing=args.packing,
report_to=args.report_to,
run_name=args.run_name or "securecoder",
remove_unused_columns=False,
)
if args.max_steps > 0:
cfg["max_steps"] = args.max_steps
else:
cfg["num_train_epochs"] = args.num_epochs
if has_eval:
cfg["eval_strategy"] = "steps"
cfg["eval_steps"] = args.save_steps
cfg["per_device_eval_batch_size"] = 1
cfg["do_eval"] = True
return make_sft_config(**cfg)
def init_trackio(args):
if args.report_to == "none":
return
try:
import trackio
if args.trackio_space:
trackio.init(project=args.trackio_project, space_id=args.trackio_space)
else:
trackio.init(project=args.trackio_project)
log.info("trackio initialised (project=%s space=%s)", args.trackio_project, args.trackio_space)
except Exception as exc: # noqa: BLE001 - monitoring must never kill training
log.warning("trackio init failed (%s); continuing without it", exc)
def train(args, model, tokenizer, records: list[dict]):
from datasets import Dataset
from trl import SFTTrainer
random.Random(args.seed).shuffle(records)
split_at = len(records) - args.eval_samples if args.eval_samples > 0 else len(records)
train_ds = Dataset.from_list(records[:split_at])
eval_ds = Dataset.from_list(records[split_at:]) if args.eval_samples > 0 else None
log.info("train rows=%d eval rows=%d", len(train_ds), len(eval_ds) if eval_ds else 0)
steps_per_epoch = len(train_ds) // max(args.batch_size * args.grad_accum, 1)
cfg = build_sft_config(args, eval_ds is not None, steps_per_epoch)
init_trackio(args)
trainer = SFTTrainer(
model=model,
tokenizer=tokenizer,
train_dataset=train_ds,
eval_dataset=eval_ds,
args=cfg,
)
started = time.time()
stats = trainer.train()
elapsed = time.time() - started
log.info("training finished in %.1f min (final loss %.4f)",
elapsed / 60, stats.metrics.get("train_loss", float("nan")))
if eval_ds is not None:
try:
metrics = trainer.evaluate()
log.info("eval_loss %.4f (train %.4f)",
metrics.get("eval_loss", float("nan")),
stats.metrics.get("train_loss", float("nan")))
except Exception as exc: # noqa: BLE001
log.warning("eval failed: %s", exc)
return trainer, stats, elapsed
def save_and_push(args, model, tokenizer):
from huggingface_hub import HfApi
api = HfApi()
api.create_repo(args.output_repo, repo_type="model", exist_ok=True, private=args.private)
model.save_pretrained(args.output_dir)
tokenizer.save_pretrained(args.output_dir)
log.info("pushing LoRA adapter to %s", args.output_repo)
model.push_to_hub(args.output_repo, tokenizer=tokenizer)
if args.merge_repo:
api.create_repo(args.merge_repo, repo_type="model", exist_ok=True, private=args.private)
log.info("merging to 16-bit and pushing to %s (large upload)", args.merge_repo)
model.push_to_hub_merged(args.merge_repo, tokenizer=tokenizer, save_method="merged_16bit")
# --------------------------------------------------------------------------
# CLI
# --------------------------------------------------------------------------
def parse_args(argv=None):
p = argparse.ArgumentParser(description="SecureCoder QLoRA fine-tune")
p.add_argument("--base-model", default="Qwen/Qwen3-Coder-30B-A3B-Instruct")
p.add_argument("--output-repo", default=None, help="Hub repo for the LoRA adapter")
p.add_argument("--merge-repo", default=None, help="optional Hub repo for a 16-bit merge")
p.add_argument("--output-dir", default="securecoder-out")
p.add_argument("--private", action="store_true", help="create Hub repos as private")
p.add_argument("--max-seq-length", type=int, default=4096)
p.add_argument("--batch-size", type=int, default=2)
p.add_argument("--grad-accum", type=int, default=8)
p.add_argument("--learning-rate", type=float, default=2e-4)
p.add_argument("--lr-scheduler", default="cosine")
p.add_argument("--num-epochs", type=float, default=1.0)
p.add_argument("--max-steps", type=int, default=0, help="overrides --num-epochs when > 0")
p.add_argument("--eval-samples", type=int, default=200, help="0 disables evaluation")
p.add_argument("--logging-steps", type=int, default=10)
p.add_argument("--save-steps", type=int, default=250)
p.add_argument("--packing", action="store_true", default=True)
p.add_argument("--no-packing", dest="packing", action="store_false")
p.add_argument("--seed", type=int, default=3407)
p.add_argument("--lora-r", type=int, default=32)
p.add_argument("--lora-alpha", type=int, default=32)
p.add_argument("--no-4bit", action="store_true")
p.add_argument("--target-modules", default=",".join(ATTENTION_TARGETS),
help="comma-separated suffixes, or 'all-linear' to include expert MLPs")
p.add_argument("--report-to", default="trackio", choices=["trackio", "none"])
p.add_argument("--trackio-project", default="securecoder")
p.add_argument("--trackio-space", default=None, help="e.g. Taimwe/securecoder-trackio")
p.add_argument("--run-name", default=None)
p.add_argument("--validate-only", action="store_true",
help="load a small sample of each source, print the mix, exit (no GPU)")
p.add_argument("--validate-per-source", type=int, default=40)
p.add_argument("--show-samples", type=int, default=3)
p.add_argument("--smoke", action="store_true",
help="tiny end-to-end run: 200 rows/source, 20 steps")
return p.parse_args(argv)
def apply_smoke(args) -> None:
args.max_steps = args.max_steps or 20
args.eval_samples = min(args.eval_samples, 20)
args.save_steps = 20
args.max_seq_length = min(args.max_seq_length, 2048)
global MIX
MIX = [Source(s.repo, 200, s.kind, s.config, s.split, s.note) for s in MIX]
def main(argv=None) -> int:
args = parse_args(argv)
if args.smoke:
apply_smoke(args)
token = os.environ.get("HF_TOKEN")
if args.validate_only:
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(args.base_model)
sources = [Source(s.repo, args.validate_per_source, s.kind, s.config, s.split, s.note)
for s in MIX]
records, stats = build_dataset(tokenizer, sources, token, validate=True)
print_stats(stats, records, tokenizer)
for i, rec in enumerate(records[: args.show_samples], 1):
print("\n" + "-" * 78)
print(f"SAMPLE {i} [{rec['source']}] {len(rec['text'])} chars")
print("-" * 78)
print(rec["text"][:1500])
return 0
import torch
if not torch.cuda.is_available():
log.error("no CUDA device - use --validate-only locally, or run on HF Jobs / Colab")
return 1
log.info("GPU: %s", torch.cuda.get_device_name(0))
if not args.output_repo:
log.error("--output-repo is required (the container/VM is ephemeral)")
return 1
model, tokenizer = load_model_and_tokenizer(args)
records, stats = build_dataset(tokenizer, MIX, token, validate=False)
print_stats(stats, records, tokenizer)
if len(records) < 100:
log.error("only %d usable rows - refusing to train", len(records))
return 1
trainer, stats_train, elapsed = train(args, model, tokenizer, records)
save_and_push(args, model, tokenizer)
print("\n" + "=" * 78)
print(f"DONE rows={len(records):,} time={elapsed / 60:.1f} min "
f"loss={stats_train.metrics.get('train_loss', float('nan')):.4f}")
print(f"adapter: https://huggingface.co/{args.output_repo}")
if args.merge_repo:
print(f"merged : https://huggingface.co/{args.merge_repo}")
print("=" * 78)
return 0
if __name__ == "__main__":
raise SystemExit(main())