TensorVizion's picture
Update app.py
3fbd9e2 verified
Raw
History Blame Contribute Delete
17.5 kB
"""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()