Create README.md
Browse files```
import torch
import torch.nn as nn
import torch.nn.functional as F
import json
from huggingface_hub import hf_hub_download
from transformers import AutoTokenizer, AutoModel, AutoConfig
class SecureBERTFlatClassifier(nn.Module):
def __init__(self, model_name, class_counts):
super().__init__()
config = AutoConfig.from_pretrained(model_name)
if hasattr(config, "reference_compile"): config.reference_compile = False
self.bert = AutoModel.from_pretrained(model_name, config=config)
def make_head(out_features, is_cvss=False):
layers = [
nn.LayerNorm(768), nn.Dropout(0.1),
nn.Linear(768, 768), nn.GELU(), nn.Dropout(0.1),
nn.Linear(768, 768), nn.GELU(), nn.Dropout(0.1),
nn.Linear(768, out_features)
]
if is_cvss: layers.append(nn.Softmax(dim=1))
return nn.Sequential(*layers)
self.cvss_heads = nn.ModuleDict({k: make_head(v, True) for k, v in
{'attack_vector': 4, 'attack_complexity': 2, 'privileges_required': 3,
'user_interaction': 2, 'scope': 2, 'confidentiality': 3, 'integrity': 3, 'availability': 3}.items()})
self.cwe_heads = nn.ModuleDict({k: make_head(v) for k, v in class_counts.items()})
def forward(self, input_ids, attention_mask):
out = self.bert(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
cls_emb = out[:, 0, :]
mask = attention_mask.unsqueeze(-1).expand(out.size()).float()
mean_emb = torch.sum(out * mask, 1) / torch.clamp(mask.sum(1), min=1e-9)
return {**{k: head(cls_emb) for k, head in self.cvss_heads.items()},
**{k: head(mean_emb) for k, head in self.cwe_heads.items()}}
class SecurePredictor:
def __init__(self, repo_id, device=None):
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
# 1. Pobieranie plik贸w
conf_path = hf_hub_download(repo_id=repo_id, filename="config.json")
model_path = hf_hub_download(repo_id=repo_id, filename="pytorch_model.bin")
with open(conf_path, "r", encoding='utf-8') as f:
self.config = json.load(f)
# 2. Inicjalizacja
self.tokenizer = AutoTokenizer.from_pretrained(self.config["base_model"])
counts = {k: len(v) for k, v in self.config["cwe_labels"].items()}
self.model = SecureBERTFlatClassifier(self.config["base_model"], counts)
self.model.load_state_dict(torch.load(model_path, map_location=self.device))
self.model.to(self.device).eval()
def predict(self, text, top_k=3):
inputs = self.tokenizer(text, return_tensors="pt", truncation=True, max_length=512).to(self.device)
with torch.no_grad():
out = self.model(inputs['input_ids'], inputs['attention_mask'])
res = {'cvss': {}, 'cwe': {}}
for k, labels in self.config["cvss_map"].items():
res['cvss'][k] = labels[torch.argmax(out[k]).item()]
for lv in ['pillar', 'class', 'base', 'variant']:
if lv in out:
probs = F.softmax(out[lv], dim=1)
scores, idxs = torch.topk(probs, k=min(top_k, probs.size(1)))
res['cwe'][lv] = [
{"cwe_id": f"CWE-{self.config['cwe_labels'][lv][i]['id']}",
"name": self.config['cwe_labels'][lv][i]['name'],
"conf": f"{s.item():.2%}"} for s, i in zip(scores[0], idxs[0])
]
return res
# --- U呕YCIE ---
predictor = SecurePredictor("bziemba/SecureBERT2.0-final")
result = predictor.predict("User-controlled input is used in a SQL query without sanitization.")
print(json.dumps(result, indent=2, ensure_ascii=False))
```