# /// 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 ```` 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())