binuser007/Toxic_comment_classification_using_Bert
0
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 } 