CoolFace
Apppublic

nick-localhost/Sign-language-detection

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
train.py119 linesDownload Raw Back to src
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")