Spaces:
Runtime error
Runtime error
| """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") | |
| 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) | |
| 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() | |