pranavKHF/code_vulnerability_detection
0
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) 