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()