nick-localhost/Sign-language-detection
0
1from data import DETRData
2from model import DETR
3from loss import DETRLoss, HungarianMatcher
4from torch.utils.data import DataLoader
5from torch import optim, load, save
6from colorama import Fore
7from utils.logger import get_logger
8from utils.rich_handlers import TrainingHandler, rich_training_context
9import sys
10import torch
11from utils.boxes import stacker
12
13if __name__ == '__main__':
14 # Initialize logger and handlers
15 logger = get_logger("training")
16 logger.print_banner()
17
18 train_dataset = DETRData('data/train')
19 train_dataloader = DataLoader(train_dataset, batch_size=4, collate_fn=stacker, drop_last=True)
20
21 test_dataset = DETRData('data/test', train=False)
22 test_dataloader = DataLoader(test_dataset, batch_size=4, collate_fn=stacker, drop_last=True)
23
24 num_classes = 3
25 model = DETR(num_classes=num_classes)
26 model.load_pretrained('pretrained/4426_model.pt')
27 model.log_model_info()
28 model.train()
29
30 opt = optim.Adam(model.parameters(), lr=1e-5)
31 scheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(opt, len(train_dataloader)*30, T_mult=2)
32
33 weights= {'class_weighting': 1, 'bbox_weighting': 5, 'giou_weighting': 2}
34 matcher = HungarianMatcher(weights)
35 criterion = DETRLoss(num_classes=num_classes, matcher=matcher, weight_dict=weights, eos_coef=0.1)
36
37 train_batches = len(train_dataloader)
38 test_batches = len(test_dataloader)
39 epochs = 100
40
41 # Log training configuration
42 training_config = {
43 "Total Epochs": epochs,
44 "Batch Size": 4,
45 "Train Batches": train_batches,
46 "Test Batches": test_batches,
47 "Learning Rate": 1e-5,
48 "Optimizer": "Adam",
49 "Scheduler": "CosineAnnealingWarmRestarts"
50 }
51 logger.print_table("๐๏ธ Training Configuration", list(training_config.keys()), [list(training_config.values())])
52
53 # Start training with rich context
54 with rich_training_context() as training_handler:
55 for epoch in range(epochs):
56 with training_handler.create_training_progress() as epoch_progress:
57 epoch_task = epoch_progress.add_task(f"[bold blue] Progress {epoch+1}/{epochs}", train_loss=0.0, test_loss=0.0, total=train_batches)
58 # Training phase
59 model.train()
60 train_epoch_loss = 0.0
61
62 # Create progress bar for current epoch
63 for batch_idx, batch in enumerate(train_dataloader):
64 X, y = batch
65 try:
66 yhat = model(X)
67 yhat_classes = yhat['pred_logits']
68 yhat_bb = yhat['pred_boxes']
69 loss_dict = criterion(yhat, y)
70 weight_dict = criterion.weight_dict
71
72 # Ensure we sum exactly over the expected weighted keys, and keep tensor dtype
73 losses = loss_dict['labels']['loss_ce']*weight_dict['class_weighting'] + loss_dict['boxes']['loss_bbox']*weight_dict['bbox_weighting'] + loss_dict['boxes']['loss_giou']*weight_dict['giou_weighting']
74
75 # Calculate loss
76 train_epoch_loss += losses.item()
77
78 # Zero grads
79 opt.zero_grad()
80
81 # Backward
82 losses.backward()
83 # Apply
84 opt.step()
85
86 # Update progress
87 epoch_progress.update(epoch_task, advance=1, train_loss=round(train_epoch_loss/train_batches,5))
88
89 except Exception as e:
90 logger.error(f"Training error at epoch {epoch}, batch {batch_idx}: {str(e)}")
91 logger.error(f"Batch targets: {str(y)}")
92 sys.exit()
93
94 # Progress lr
95 scheduler.step()
96
97 # Test phase
98 model.eval()
99 test_epoch_loss = 0.0
100 with torch.no_grad():
101 for batch_idx, batch in enumerate(test_dataloader):
102 X, y = batch
103 yhat = model(X)
104 loss_dict = criterion(yhat, y)
105 weight_dict = criterion.weight_dict
106 losses = loss_dict['labels']['loss_ce']*weight_dict['class_weighting'] + loss_dict['boxes']['loss_bbox']*weight_dict['bbox_weighting'] + loss_dict['boxes']['loss_giou']*weight_dict['giou_weighting']
107
108 # Calculate loss
109 test_epoch_loss += losses.item()
110 epoch_progress.update(epoch_task, advance=0, test_loss=round(test_epoch_loss/test_batches,5))
111
112 # Save checkpoints
113 if epoch % 10 == 0 and epoch != 0:
114 checkpoint_path = f"checkpoints/{epoch}_model.pt"
115 save(model.state_dict(), checkpoint_path)
116 training_handler.save_checkpoint_status(checkpoint_path, epoch)
117
118 # Final save
119 save(model.state_dict(), f"checkpoints/{epoch}_model.pt")