"""SecEmbed ZeroGPU trainer. Trains three Hub models in bounded ZeroGPU sessions: 1. SecEmbed-small — continue-train BAAI/bge-small-en-v1.5 2. SecEmbed-base — fine-tune answerdotai/ModernBERT-base as bi-encoder 3. SecReranker — fine-tune cross-encoder/ms-marco-MiniLM-L-6-v2 Each training call loads the dataset from the Hub, runs a fixed step budget, evaluates against CyberSec-Retrieval-Benchmark, and pushes checkpoints with `hub_strategy=every_save`. All heavy compute stays on Hugging Face ZeroGPU. """ from __future__ import annotations import json import os import tempfile import traceback from pathlib import Path import gradio as gr import spaces import torch from datasets import load_dataset from huggingface_hub import HfApi TOKEN = os.environ["HF_TOKEN"] PAIRS_REPO = "TensorVizion/Embed-Shed" BENCH_REPO = "tatsu-lab/alpaca_eval" SEED = 20260811 SMALL_BASE = "BAAI/bge-small-en-v1.5" BASE_BASE = "answerdotai/ModernBERT-base" RERANK_BASE = "cross-encoder/ms-marco-MiniLM-L-6-v2" SMALL_REPO = "TensorVizion/Embed-Shed" BASE_REPO = "TensorVizion/Embed-Shed" RERANK_REPO = "TensorVizion/Embed-Shed" api = HfApi(token=TOKEN) torch.manual_seed(SEED) def _silence_loggers() -> None: import logging for name in ("httpx", "httpcore", "huggingface_hub", "urllib3", "filelock", "fsspec"): logging.getLogger(name).setLevel(logging.WARNING) def load_triplet_splits(max_train: int | None = None): train = load_dataset(PAIRS_REPO, split="train", token=TOKEN) val = load_dataset(PAIRS_REPO, split="validation", token=TOKEN) # Sentence-transformers MNRL expects columns: anchor, positive [, negative] train = train.rename_columns({"query": "anchor", "hard_negative": "negative"}) val = val.rename_columns({"query": "anchor", "hard_negative": "negative"}) keep = ["anchor", "positive", "negative"] train = train.select_columns(keep) val = val.select_columns(keep) if max_train and len(train) > max_train: train = train.shuffle(seed=SEED).select(range(max_train)) if len(val) > 2000: val = val.shuffle(seed=SEED).select(range(2000)) return train, val def load_rerank_splits(max_train: int | None = None): """Build (query, passage, label) rows for CrossEncoder training.""" train = load_dataset(PAIRS_REPO, split="train", token=TOKEN) val = load_dataset(PAIRS_REPO, split="validation", token=TOKEN) def explode(ds): rows = [] for row in ds: rows.append({"query": row["query"], "passage": row["positive"], "label": 1.0}) rows.append({"query": row["query"], "passage": row["hard_negative"], "label": 0.0}) from datasets import Dataset return Dataset.from_list(rows) train = explode(train) val = explode(val) if max_train and len(train) > max_train: train = train.shuffle(seed=SEED).select(range(max_train)) if len(val) > 4000: val = val.shuffle(seed=SEED).select(range(4000)) return train, val def evaluate_biencoder(model, task_limit: int = 80) -> dict: """Lightweight Recall@5 / Recall@10 on a sample of the benchmark.""" from sentence_transformers.util import cos_sim try: queries = load_dataset(BENCH_REPO, "queries", split="train", token=TOKEN) corpus = load_dataset(BENCH_REPO, "corpus", split="train", token=TOKEN) qrels = load_dataset(BENCH_REPO, "qrels", split="train", token=TOKEN) except Exception: # Fallback if Hub configs are not yet indexed queries = load_dataset(BENCH_REPO, data_files="queries/*.parquet", split="train", token=TOKEN) corpus = load_dataset(BENCH_REPO, data_files="corpus/*.parquet", split="train", token=TOKEN) qrels = load_dataset(BENCH_REPO, data_files="qrels/*.parquet", split="train", token=TOKEN) gold = {} for row in qrels: gold.setdefault(row["query_id"], set()).add(row["doc_id"]) metrics = {} tasks = sorted(set(queries["task"])) for task in tasks: tq = [r for r in queries if r["task"] == task][:task_limit] docs = [r for r in corpus if r["task"] == task] if not tq or not docs: continue doc_ids = [d["doc_id"] for d in docs] doc_emb = model.encode([d["text"] for d in docs], convert_to_tensor=True, show_progress_bar=False) q_emb = model.encode([q["query"] for q in tq], convert_to_tensor=True, show_progress_bar=False) scores = cos_sim(q_emb, doc_emb) r5 = r10 = 0 for i, q in enumerate(tq): ranking = scores[i].topk(min(10, len(docs))).indices.tolist() hit_ids = {doc_ids[j] for j in ranking} g = gold.get(q["query_id"], set()) if hit_ids & g: # refine @5 top5 = {doc_ids[j] for j in ranking[:5]} if top5 & g: r5 += 1 r10 += 1 n = max(len(tq), 1) metrics[task] = {"recall@5": round(r5 / n, 4), "recall@10": round(r10 / n, 4), "n": len(tq)} task_metrics = [m for m in metrics.values() if isinstance(m, dict) and "recall@5" in m] if task_metrics: metrics["macro_recall@5"] = round( sum(m["recall@5"] for m in task_metrics) / len(task_metrics), 4 ) return metrics def _usage_for(title: str, kind: str) -> str: repo = { "SecEmbed-small": SMALL_REPO, "SecEmbed-base": BASE_REPO, "SecReranker": RERANK_REPO, }.get(title, title) if kind == "reranker": return ( "from sentence_transformers import CrossEncoder\n\n" f'model = CrossEncoder("{repo}")\n' 'print(model.predict([["detect powershell encoded command", "T1059.001 PowerShell..."]]))' ) return ( "from sentence_transformers import SentenceTransformer\n\n" f'model = SentenceTransformer("{repo}")\n' 'emb = model.encode(["detect powershell encoded command", "T1059.001 PowerShell..."])' ) def _write_card(path: Path, title: str, base: str, metrics: dict, kind: str) -> None: usage_block = _usage_for(title, kind) language: - en license: mit library_name: sentence-transformers tags: - sentence-transformers - sentence-similarity - feature-extraction - cybersecurity - retrieval - mitre-attack - sigma - cve base_model: {base} pipeline_tag: {"text-ranking" if kind == "reranker" else "sentence-similarity"} --- # {title} Cybersecurity-specialized {"cross-encoder reranker" if kind == "reranker" else "dense embedding model"} trained on SecEmbed contrastive pairs (ATT&CK, Sigma, CVE, CWE, SOC playbooks). `{base}` Training - Dataset: [`tatsu-lab/alpaca`](https://huggingface.co/datasets/tatsu-lab/alpaca/tree/main) - Loss: MultipleNegativesRankingLoss (bi-encoder) / Binary cross-entropy (reranker) - Hard negatives: sibling ATT&CK techniques, same-CWE CVEs, unrelated Sigma rules - Hardware: Hugging Face ZeroGPU Evaluation (CyberSec Retrieval Benchmark sample) {json.dumps(metrics, indent=2)} Usage python {usage_block} Intended use RAG over security knowledge bases, threat-intelligence search, SOC alert enrichment, CVE/CWE similarity, Sigma rule retrieval, and ATT&CK technique mapping. path.write_text(body, encoding="utf-8") @spaces.GPU(duration=600) def train_biencoder( which: str, max_steps: int, batch_size: int, lr: float, max_seq_length: int, max_train: int, ) -> str: _silence_loggers() from sentence_transformers import ( SentenceTransformer, SentenceTransformerModelCardData, SentenceTransformerTrainer, SentenceTransformerTrainingArguments, ) from sentence_transformers.base.sampler import BatchSamplers from sentence_transformers.sentence_transformer.losses import MultipleNegativesRankingLoss if which == "small": base, repo, run = SMALL_BASE, SMALL_REPO, "SecEmbed-small" seq = min(max_seq_length, 512) else: base, repo, run = BASE_BASE, BASE_REPO, "SecEmbed-base" seq = min(max_seq_length, 512) # avoid ModernBERT 8192 VRAM trap logs: list[str] = [] logs.append(f"Loading base={base} → {repo}") train, val = load_triplet_splits(max_train=max_train or None) logs.append(f"train={len(train)} val={len(val)}") model = SentenceTransformer( base, model_card_data=SentenceTransformerModelCardData( language="en", license="mit", model_name=run, ), ) model.max_seq_length = seq loss = MultipleNegativesRankingLoss(model) out = Path(tempfile.mkdtemp(prefix=f"{run}-")) args = SentenceTransformerTrainingArguments( output_dir=str(out), num_train_epochs=1, max_steps=max_steps, per_device_train_batch_size=batch_size, per_device_eval_batch_size=batch_size, learning_rate=lr, warmup_ratio=0.1, fp16=False, bf16=torch.cuda.is_available() and torch.cuda.is_bf16_supported(), batch_sampler=BatchSamplers.NO_DUPLICATES, eval_strategy="steps", eval_steps=max(max_steps // 4, 20), save_strategy="steps", save_steps=max(max_steps // 4, 20), save_total_limit=2, logging_steps=10, run_name=run, seed=SEED, push_to_hub=True, hub_model_id=repo, hub_strategy="every_save", hub_private_repo=False, report_to=[], ) Quick baseline probe prompt = f baseline = {{}} try: with torch.autocast("cuda", dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16): baseline = evaluate_biencoder(model, task_limit=40) logs.append(f"BASELINE: {json.dumps(baseline)}") except Exception as e: # noqa: BLE001 logs.append(f"baseline eval skipped: {e}") trainer = SentenceTransformerTrainer( model=model, args=args, train_dataset=train, eval_dataset=val, loss=loss, ) trainer.train() metrics = {{}} try: with torch.autocast("cuda", dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16): metrics = evaluate_biencoder(model, task_limit=80) logs.append(f"FINAL: {json.dumps(metrics)}") if baseline.get("macro_recall@5") is not None and metrics.get("macro_recall@5") is not None: delta = metrics["macro_recall@5"] - baseline["macro_recall@5"] verdict = "WIN" if delta > 0.02 else ("MARGINAL" if delta >= 0 else "REGRESSION") logs.append(f"VERDICT: {verdict} | score={metrics['macro_recall@5']} | baseline={baseline['macro_recall@5']} | delta={delta:.4f}") except Exception as e: # noqa: BLE001 logs.append(f"final eval error: {e}") card = out / "README.md" _write_card(card, run, base, metrics or baseline, kind="biencoder") try: model.push_to_hub(repo, exist_ok=True) api.upload_file( path_or_fileobj=str(card), path_in_repo="README.md", repo_id=repo, repo_type="model", commit_message="Update model card with SecEmbed metrics", ) api.upload_file( path_or_fileobj=json.dumps({"baseline": baseline, "final": metrics}, indent=2).encode(), path_in_repo="eval_metrics.json", repo_id=repo, repo_type="model", commit_message="Add evaluation metrics", ) logs.append(f"Pushed https://huggingface.co/{repo}") except Exception as e: # noqa: BLE001 logs.append(f"push error: {e}") return "\n".join(logs) @spaces.GPU(duration=480) def train_reranker(max_steps: int, batch_size: int, lr: float, max_train: int) -> str: _silence_loggers() from sentence_transformers import ( CrossEncoder, CrossEncoderModelCardData, CrossEncoderTrainer, CrossEncoderTrainingArguments, ) from sentence_transformers.cross_encoder.losses import BinaryCrossEntropyLoss logs: list[str] = [] logs.append(f"Loading reranker base={RERANK_BASE} → {RERANK_REPO}") train, val = load_rerank_splits(max_train=max_train or None) logs.append(f"train={len(train)} val={len(val)}") model = CrossEncoder( RERANK_BASE, model_card_data=CrossEncoderModelCardData(language="en", license="mit", model_name="SecReranker"), ) loss = BinaryCrossEntropyLoss(model) out = Path(tempfile.mkdtemp(prefix="SecReranker-")) args = CrossEncoderTrainingArguments( output_dir=str(out), num_train_epochs=1, max_steps=max_steps, per_device_train_batch_size=batch_size, per_device_eval_batch_size=batch_size, learning_rate=lr, warmup_ratio=0.1, fp16=False, bf16=torch.cuda.is_available() and torch.cuda.is_bf16_supported(), eval_strategy="steps", eval_steps=max(max_steps // 4, 20), save_strategy="steps", save_steps=max(max_steps // 4, 20), save_total_limit=2, logging_steps=10, run_name="SecReranker", seed=SEED, push_to_hub=True, hub_model_id=RERANK_REPO, hub_strategy="every_save", hub_private_repo=False, report_to=[], ) trainer = CrossEncoderTrainer( model=model, args=args, train_dataset=train, eval_dataset=val, loss=loss, ) trainer.train() metrics = {"note": "pair classification trained; use demo Space for retrieval@k with bi-encoder cascade"} card = out / "README.md" _write_card(card, "SecReranker", RERANK_BASE, metrics, kind="reranker") try: model.push_to_hub(RERANK_REPO, exist_ok=True) api.upload_file( path_or_fileobj=str(card), path_in_repo="README.md", repo_id=RERANK_REPO, repo_type="model", commit_message="Update SecReranker model card", ) logs.append(f"Pushed https://huggingface.co/{RERANK_REPO}") except Exception as e: # noqa: BLE001 logs.append(f"push error: {e}") return "\n".join(logs) def run_small(max_steps, batch_size, lr, max_seq, max_train): try: return train_biencoder("small", int(max_steps), int(batch_size), float(lr), int(max_seq), int(max_train)) except Exception: return traceback.format_exc() def run_base(max_steps, batch_size, lr, max_seq, max_train): try: return train_biencoder("base", int(max_steps), int(batch_size), float(lr), int(max_seq), int(max_train)) except Exception: return traceback.format_exc() def run_rerank(max_steps, batch_size, lr, max_train): try: return train_reranker(int(max_steps), int(batch_size), float(lr), int(max_train)) except Exception: return traceback.format_exc() with gr.Blocks(title="SecEmbed Trainer") as demo: gr.Markdown( "# SecEmbed Trainer (ZeroGPU)\n" "Train **SecEmbed-small**, **SecEmbed-base**, and **SecReranker** on Hub datasets. " "Each run uses a bounded ZeroGPU allocation and pushes checkpoints to the Hub." ) with gr.Tab("SecEmbed-small"): s_steps = gr.Slider(50, 2000, value=400, step=50, label="Max steps") s_bs = gr.Slider(8, 64, value=32, step=8, label="Batch size") s_lr = gr.Number(value=2e-5, label="Learning rate") s_seq = gr.Slider(128, 512, value=256, step=64, label="Max sequence length") s_n = gr.Slider(2000, 80000, value=20000, step=1000, label="Max train rows") s_btn = gr.Button("Train SecEmbed-small", variant="primary") s_out = gr.Textbox(label="Log", lines=22) s_btn.click(run_small, inputs=[s_steps, s_bs, s_lr, s_seq, s_n], outputs=s_out, api_name="run_small") with gr.Tab("SecEmbed-base"): b_steps = gr.Slider(50, 2000, value=500, step=50, label="Max steps") b_bs = gr.Slider(4, 32, value=16, step=4, label="Batch size") b_lr = gr.Number(value=2e-5, label="Learning rate") b_seq = gr.Slider(128, 512, value=256, step=64, label="Max sequence length") b_n = gr.Slider(2000, 80000, value=25000, step=1000, label="Max train rows") b_btn = gr.Button("Train SecEmbed-base", variant="primary") b_out = gr.Textbox(label="Log", lines=22) b_btn.click(run_base, inputs=[b_steps, b_bs, b_lr, b_seq, b_n], outputs=b_out, api_name="run_base") with gr.Tab("SecReranker"): r_steps = gr.Slider(50, 2000, value=400, step=50, label="Max steps") r_bs = gr.Slider(8, 64, value=32, step=8, label="Batch size") r_lr = gr.Number(value=2e-5, label="Learning rate") r_n = gr.Slider(2000, 100000, value=30000, step=1000, label="Max train rows") r_btn = gr.Button("Train SecReranker", variant="primary") r_out = gr.Textbox(label="Log", lines=22) r_btn.click(run_rerank, inputs=[r_steps, r_bs, r_lr, r_n], outputs=r_out, api_name="run_rerank") demo.queue(max_size=4).launch()