Download train_securecoder.py from Taimwe/securecoder-scripts: direct link, hf CLI and curl.
- Browser
- Download file 37.8 kB
-
https://huggingface.co/Taimwe/securecoder-scripts/resolve/a2ba1ae024c985ffb58eaed81b24fc7fed01ffad/train_securecoder.py
- Command line
-
hf download hf://Taimwe/securecoder-scripts@a2ba1ae024c985ffb58eaed81b24fc7fed01ffad/train_securecoder.py
-
curl -L -o train_securecoder.py https://huggingface.co/Taimwe/securecoder-scripts/resolve/a2ba1ae024c985ffb58eaed81b24fc7fed01ffad/train_securecoder.py
37.8 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 | |
| # -------------------------------------------------------------------------- | |
| 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("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 projections only by default, for two reasons: | |
| # * Qwen3's MoE router is a custom `Qwen3MoeTopKRouter` module, not nn.Linear, so | |
| # listing it as a LoRA target dies with "Target module ... is not supported". | |
| # * adapting all 128 experts x 32 layers is ~800M trainable parameters, which | |
| # dominates VRAM and step time. | |
| # Use --target-modules all-linear to include the expert MLPs (PEFT then skips the | |
| # modules it cannot adapt instead of failing). | |
| ATTENTION_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj"] | |
| 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()] | |
| peft_kwargs: dict[str, Any] = dict( | |
| r=args.lora_r, | |
| lora_alpha=args.lora_alpha, | |
| lora_dropout=0.0, | |
| bias="none", | |
| use_gradient_checkpointing="unsloth", | |
| random_state=args.seed, | |
| use_rslora=False, | |
| ) | |
| try: | |
| model = FastLanguageModel.get_peft_model(model, target_modules=targets, **peft_kwargs) | |
| except ValueError as exc: | |
| if "is not supported" not in str(exc) or targets == ATTENTION_TARGETS: | |
| raise | |
| log.warning("LoRA targets rejected (%s); retrying with attention projections only", | |
| str(exc).splitlines()[0][:180]) | |
| model = FastLanguageModel.get_peft_model( | |
| model, target_modules=ATTENTION_TARGETS, **peft_kwargs | |
| ) | |
| trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) | |
| total = sum(p.numel() for p in model.parameters()) | |
| log.info("trainable parameters: %s (%.2f%% of the model)", f"{trainable:,}", | |
| 100 * trainable / max(total, 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) | |
| args.report_to = "none" | |
| 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) | |
| init_trackio(args) | |
| cfg = build_sft_config(args, eval_ds is not None, steps_per_epoch) | |
| 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 _push_adapter(api, model, args): | |
| """Unsloth's push_to_hub signature varies between releases (one version | |
| rejects ``tokenizer=``), so fall back to uploading the saved folder - the | |
| tokenizer files sit next to the adapter and get pushed either way.""" | |
| try: | |
| model.push_to_hub(args.output_repo) | |
| return | |
| except TypeError as exc: | |
| log.warning("push_to_hub rejected our arguments (%s); uploading the folder instead", exc) | |
| except Exception as exc: # noqa: BLE001 - never lose a finished run to an upload quirk | |
| log.warning("model.push_to_hub failed (%s); uploading the folder instead", exc) | |
| api.upload_folder(folder_path=args.output_dir, repo_id=args.output_repo, repo_type="model") | |
| def _push_merged(api, model, tokenizer, args): | |
| try: | |
| model.push_to_hub_merged(args.merge_repo, tokenizer=tokenizer, save_method="merged_16bit") | |
| return | |
| except TypeError as exc: | |
| log.warning("push_to_hub_merged rejected tokenizer= (%s); retrying without it", exc) | |
| model.push_to_hub_merged(args.merge_repo, save_method="merged_16bit") | |
| 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) | |
| _push_adapter(api, model, args) | |
| 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) | |
| _push_merged(api, model, tokenizer, args) | |
| # -------------------------------------------------------------------------- | |
| # 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("--mix-scale", type=float, default=1.0, | |
| help="multiply every source limit by this (e.g. 0.5 for a half mix)") | |
| 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: | |
| global MIX | |
| args = parse_args(argv) | |
| if args.smoke: | |
| apply_smoke(args) | |
| elif args.mix_scale != 1.0: | |
| MIX = [Source(s.repo, max(50, int(s.limit * args.mix_scale)), | |
| s.kind, s.config, s.split, s.note) for s in MIX] | |
| log.info("mix scaled by %.2f -> %d rows planned", args.mix_scale, | |
| sum(s.limit for s in MIX)) | |
| 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}") | |
| speed = stats_train.metrics.get("train_samples_per_second") or 0 | |
| if speed: | |
| est_hours = len(records) / speed / 3600 | |
| print(f"throughput: {speed:.1f} rows/s -> a 1-epoch pass over {len(records):,} rows " | |
| f"≈ {est_hours:.1f} h at this rate") | |
| print(f"cost at $1.80/h (l40sx1): ≈ ${est_hours * 1.80:.0f} " | |
| f"at $2.50/h (a100-large): ≈ ${est_hours * 2.50:.0f}") | |
| 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()) | |