binuser007/Toxic_comment_classification_using_Bert
0
1import torch
2import torch.nn as nn
3from transformers import AutoModel
4from typing import Dict, Tuple
5
6class ToxicClassifier(nn.Module):
7 def __init__(self, num_classes: int = 6, dropout: float = 0.3):
8 super(ToxicClassifier, self).__init__()
9
10 # BERT base model - freeze some layers to prevent overfitting
11 self.bert = AutoModel.from_pretrained('bert-base-uncased')
12
13 # Freeze the first 8 layers of BERT
14 for param in list(self.bert.parameters())[:-8]:
15 param.requires_grad = False
16
17 # Simplified architecture focusing on BERT's power
18 self.dropout = nn.Dropout(dropout)
19 self.classifier = nn.Linear(768, num_classes) # 768 is BERT's hidden size
20
21 # Initialize the classifier weights properly
22 torch.nn.init.xavier_uniform_(self.classifier.weight)
23 self.classifier.bias.data.fill_(0.0)
24
25 def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
26 # Get BERT embeddings
27 outputs = self.bert(input_ids, attention_mask=attention_mask)
28 pooled_output = outputs.pooler_output # [batch_size, 768]
29
30 # Apply dropout and classification
31 pooled_output = self.dropout(pooled_output)
32 logits = self.classifier(pooled_output)
33
34 return logits # Return logits directly, BCEWithLogitsLoss will handle the sigmoid