Team Ai
Apppublic

binuser007/Toxic_comment_classification_using_Bert

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
toxic_classifier.py34 linesDownload Raw Back to models
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