# Function to evaluate model performance def evaluate_model(model, dataloader, device): model.eval() # Set model to evaluation mode all_preds = [] all_labels = [] # Disable gradient calculations with torch.no_grad(): for batch in dataloader: input_ids = batch['input_ids'].to(device) attention_mask = batch['attention_mask'].to(device) labels = batch['labels'].to(device) # Forward pass to get logits outputs = model(input_ids, attention_mask=attention_mask) logits = outputs.logits # Get predictions preds = torch.argmax(logits, dim=1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.cpu().numpy()) # Calculate evaluation metrics accuracy = accuracy_score(all_labels, all_preds) precision, recall, f1, _ = precision_recall_fscore_support(all_labels, all_preds, average='binary') return accuracy, precision, recall, f1