Team Ai
Apppublic

pranavKHF/code_vulnerability_detection

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
model.py77 linesDownload Raw Back to src
1"""CodeT5 Vulnerability Detection model2Binary Classication Safe(0) vs Vulnerable(1)"""3 4import torch5import torch.nn as nn6from transformers import T5ForConditionalGeneration, RobertaTokenizer7 8class VulnerabilityCodeT5(nn.Module):9    """CodeT5 model for vulnerability detection"""10    11    def __init__(self, model_name="Salesforce/codet5-base", num_labels=2):12        super().__init__()13 14        self.encoder_decoder = T5ForConditionalGeneration.from_pretrained(model_name)15 16        #Get hidden size from config17        hidden_size = self.encoder_decoder.config.d_model #768 for base18 19        #Classification Head20        self.classifier = nn.Sequential(21            nn.Dropout(0.1),22            nn.Linear(hidden_size, hidden_size),23            nn.ReLU(),24            nn.Dropout(0.1),25            nn.Linear(hidden_size, num_labels)26        )27 28        self.num_labels = num_labels29 30    def forward(self, input_ids, attention_mask, labels=None):31        """32        Forward pass33        Args:34        input_ids : tokenized code [batch_size, seq_len]35        attention_mask : attention mask [batch_size, seq_len]36        labels: ground truth labels [batch_size]37        """38 39        #Get encoder outputs40        encoder_outputs = self.encoder_decoder.encoder(41            input_ids=input_ids,42            attention_mask=attention_mask,43            return_dict=True44        )45 46        #Pool encoder outputs (use first token [CLS])47        hidden_state = encoder_outputs.last_hidden_state # [batch, seq_len, hidden]48        pooled_output = hidden_state[:, 0, :] # [batch, hidden]49 50        #Classification 51        logits = self.classifier(pooled_output) # [batch, num_labels]52 53        #Calculate loss54        loss = None55        if labels is not None:56            loss_fn = nn.CrossEntropyLoss()57            loss = loss_fn(logits, labels)58 59        return {60            'loss': loss,61            'logits': logits,62            'hidden_states': hidden_state63        }64    65    def predict(self, input_ids, attention_mask):66        """Make Predictions"""67        self.eval()68        with torch.no_grad():69            outputs = self.forward(input_ids, attention_mask)70            probs = torch.softmax(outputs["logits"], dim=1)71            predictions = torch.argmax(probs, dim=1)72 73        return predictions, probs74    75def count_parameters(model):76    """Count trainable parameters"""77    return sum(p.numel() for p in model.parameters() if p.requires_grad)