Zimbra-detector-api / app /models /vulnerability_model.py
Arun111123Kumar's picture
Upload folder using huggingface_hub
c5624aa verified
Raw
History Blame Contribute Delete
3.42 kB
import torch
import os
from transformers import AutoTokenizer, AutoModelForSequenceClassification
from peft import PeftModel
import logging
# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class VulnerabilityDetector:
def __init__(self):
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.labels = [
"SQL Injection",
"Cross-Site Scripting (XSS)",
"Insecure Deserialization",
"Command Injection",
"Path Traversal"
]
self.model = None
self.tokenizer = None
self.base_model_path = "checkpoint-800"
self.zimbra_model_path = "checkpoint-297"
def load_model(self):
"""Load both checkpoints in the correct order"""
try:
logger.info("Loading tokenizer...")
self.tokenizer = AutoTokenizer.from_pretrained(self.zimbra_model_path)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
logger.info(f"Loading base model on {self.device}...")
base_model = AutoModelForSequenceClassification.from_pretrained(
"deepseek-ai/deepseek-coder-1.3b-instruct",
num_labels=6,
device_map=self.device,
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32
)
logger.info("Loading checkpoint-800 (Java vulnerability model)...")
model = PeftModel.from_pretrained(base_model, self.base_model_path)
logger.info("Loading checkpoint-297 (Zimbra-specific model)...")
model = PeftModel.from_pretrained(model, self.zimbra_model_path, adapter_name="zimbra")
model.set_adapter("zimbra")
self.model = model
self.model.eval()
logger.info("Model loaded successfully!")
except Exception as e:
logger.error(f"Error loading model: {str(e)}")
raise
def predict(self, code: str):
"""Predict vulnerability for a single code snippet"""
if self.model is None:
self.load_model()
inputs = self.tokenizer(
code,
return_tensors="pt",
truncation=True,
max_length=1024,
padding=True
).to(self.device)
with torch.no_grad():
outputs = self.model(**inputs)
probs = torch.nn.functional.softmax(outputs.logits, dim=-1)
# Get top prediction (only for 5 Zimbra classes)
zimbra_probs = probs[0][:5] # First 5 classes are our Zimbra classes
prediction = zimbra_probs.argmax().item()
confidence = zimbra_probs[prediction].item()
# Get all class probabilities
all_probs = {self.labels[i]: float(zimbra_probs[i]) for i in range(5)}
return {
"prediction": self.labels[prediction],
"confidence": confidence,
"all_probabilities": all_probs,
"is_vulnerable": confidence > 0.6 # Threshold for vulnerability
}
# Singleton instance
model_instance = VulnerabilityDetector()