Team Ai
Apppublic

binuser007/Toxic_comment_classification_using_Bert

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
trainer.py86 linesDownload Raw Back to models
1import torch
2from torch.utils.data import DataLoader
3from typing import Dict, List
4from tqdm import tqdm
5from torch.amp import autocast, GradScaler
6
7class ModelTrainer:
8    def __init__(self, model, optimizer, criterion, device, scaler: GradScaler = None, scheduler=None):
9        self.model = model
10        self.optimizer = optimizer
11        self.criterion = criterion
12        self.device = device
13        self.scaler = scaler or GradScaler('cuda')
14        self.use_amp = device.type == 'cuda'
15        self.scheduler = scheduler
16
17    def train_epoch(self, dataloader: DataLoader) -> Dict[str, float]:
18        self.model.train()
19        total_loss = 0
20        
21        for batch in tqdm(dataloader, desc="Training"):
22            input_ids = batch['input_ids'].to(self.device)
23            attention_mask = batch['attention_mask'].to(self.device)
24            labels = batch['labels'].to(self.device)
25
26            self.optimizer.zero_grad()
27            
28            if self.use_amp:
29                with autocast('cuda'):
30                    outputs = self.model(input_ids, attention_mask)
31                    loss = self.criterion(outputs, labels)
32                
33                self.scaler.scale(loss).backward()
34                
35                # Clip gradients
36                self.scaler.unscale_(self.optimizer)
37                torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
38                
39                self.scaler.step(self.optimizer)
40                self.scaler.update()
41            else:
42                outputs = self.model(input_ids, attention_mask)
43                loss = self.criterion(outputs, labels)
44                loss.backward()
45                torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
46                self.optimizer.step()
47            
48            if self.scheduler is not None:
49                self.scheduler.step()
50            
51            total_loss += loss.item()
52
53        return {'loss': total_loss / len(dataloader)}
54
55    def evaluate(self, dataloader: DataLoader) -> Dict[str, float]:
56        self.model.eval()
57        total_loss = 0
58        predictions = []
59        true_labels = []
60
61        with torch.no_grad():
62            for batch in tqdm(dataloader, desc="Evaluating"):
63                input_ids = batch['input_ids'].to(self.device)
64                attention_mask = batch['attention_mask'].to(self.device)
65                labels = batch['labels'].to(self.device)
66
67                if self.use_amp:
68                    with autocast('cuda'):
69                        outputs = self.model(input_ids, attention_mask)
70                        loss = self.criterion(outputs, labels)
71                else:
72                    outputs = self.model(input_ids, attention_mask)
73                    loss = self.criterion(outputs, labels)
74                
75                # Apply sigmoid to get probabilities for predictions
76                probs = torch.sigmoid(outputs)
77                
78                total_loss += loss.item()
79                predictions.extend(probs.cpu().numpy())
80                true_labels.extend(labels.cpu().numpy())
81
82        return {
83            'loss': total_loss / len(dataloader),
84            'predictions': predictions,
85            'true_labels': true_labels
86        }