From 02007c51bea219ce506eaa67696a82cd98760852 Mon Sep 17 00:00:00 2001 From: Arvindhbabu Date: Wed, 8 Jul 2026 16:21:42 +0530 Subject: [PATCH] feat: complete Sprint 5 training infrastructure --- configs/default.yaml | 34 ++- src/models/__init__.py | 1 + src/models/base_model.py | 35 +++ src/models/efficientnet_bilstm.py | 18 ++ src/models/model_factory.py | 37 +++ src/models/vit_temporal_pooling.py | 95 ++++++++ src/training/checkpoint.py | 145 ++++++++++- src/training/optimizer_factory.py | 56 +++++ src/training/scheduler_factory.py | 45 ++++ src/training/trainer.py | 94 ++++++-- tests/test_model_factory.py | 19 ++ tests/test_optimizer_scheduler.py | 34 +++ tests/test_training_pipeline.py | 88 +++++++ train.py | 374 +++++++++++++++++++++++++++++ vit_temporal_pooling.zip | Bin 0 -> 9450 bytes 15 files changed, 1042 insertions(+), 33 deletions(-) create mode 100644 src/models/base_model.py create mode 100644 src/models/efficientnet_bilstm.py create mode 100644 src/models/model_factory.py create mode 100644 src/models/vit_temporal_pooling.py create mode 100644 src/training/optimizer_factory.py create mode 100644 src/training/scheduler_factory.py create mode 100644 tests/test_model_factory.py create mode 100644 tests/test_optimizer_scheduler.py create mode 100644 tests/test_training_pipeline.py create mode 100644 train.py create mode 100644 vit_temporal_pooling.zip diff --git a/configs/default.yaml b/configs/default.yaml index 67e68ad..904e4ee 100644 --- a/configs/default.yaml +++ b/configs/default.yaml @@ -6,6 +6,8 @@ paths: raw_data: data/raw processed_data: data/processed sequences: data/intermediate/sequences + + outputs: outputs checkpoints: outputs/checkpoints logs: outputs/logs metrics: outputs/metrics @@ -13,31 +15,47 @@ paths: dataset: name: FFPP_CelebDF + image_size: 224 + sequence_length: 30 + fps: 5 + +dataloader: batch_size: 8 - num_workers: 4 + + num_workers: 0 + + pin_memory: true + + drop_last: false training: epochs: 30 + learning_rate: 0.0001 + weight_decay: 0.0001 + optimizer: AdamW + scheduler: CosineAnnealingLR model: - backbone: vit_b_16 + name: vit_temporal_pooling + pretrained: true + freeze_backbone: true - temporal_pooling: mean + num_classes: 2 -dataloader: - batch_size: 8 - num_workers: 4 - pin_memory: true - drop_last: false + pooling: mean + + hidden_dim: 256 + + dropout: 0.5 evaluation: metrics: diff --git a/src/models/__init__.py b/src/models/__init__.py index e69de29..6663b2a 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -0,0 +1 @@ +from .model_factory import ModelFactory \ No newline at end of file diff --git a/src/models/base_model.py b/src/models/base_model.py new file mode 100644 index 0000000..085a526 --- /dev/null +++ b/src/models/base_model.py @@ -0,0 +1,35 @@ +""" +DeepVision AI + +Base Model Interface +""" + +from abc import ABC, abstractmethod + +import torch.nn as nn + + +class BaseModel(nn.Module, ABC): + """ + Base class for all DeepVision AI models. + """ + + def __init__(self): + super().__init__() + + @abstractmethod + def forward(self, x): + """ + Forward pass. + + Args: + x (torch.Tensor) + + Returns: + torch.Tensor + """ + pass + + @property + def name(self): + return self.__class__.__name__ \ No newline at end of file diff --git a/src/models/efficientnet_bilstm.py b/src/models/efficientnet_bilstm.py new file mode 100644 index 0000000..52c1ed1 --- /dev/null +++ b/src/models/efficientnet_bilstm.py @@ -0,0 +1,18 @@ +""" +Placeholder. + +Will be integrated in Sprint 4.2 +""" + +import torch.nn as nn + + +class EfficientNetBiLSTM(nn.Module): + + def __init__(self): + super().__init__() + + def forward(self, x): + raise NotImplementedError( + "EfficientNet-BiLSTM integration begins in Sprint 4.2" + ) \ No newline at end of file diff --git a/src/models/model_factory.py b/src/models/model_factory.py new file mode 100644 index 0000000..c4237df --- /dev/null +++ b/src/models/model_factory.py @@ -0,0 +1,37 @@ +""" +DeepVision AI + +Model Factory +""" + +from src.models.vit_temporal_pooling import ViTTemporalPooling +from src.models.efficientnet_bilstm import EfficientNetBiLSTM + + +class ModelFactory: + + @staticmethod + def create(config): + + model_name = config["model"]["name"].lower() + + + if model_name == "vit_temporal_pooling": + + return ViTTemporalPooling( + image_size=config["dataset"]["image_size"], + num_classes=config["model"]["num_classes"], + pretrained=config["model"]["pretrained"], + freeze_backbone=config["model"]["freeze_backbone"], + pooling=config["model"]["pooling"], + hidden_dim=config["model"]["hidden_dim"], + dropout=config["model"]["dropout"], + ) + + elif model_name == "efficientnet_bilstm": + + return EfficientNetBiLSTM() + + raise ValueError( + f"Unsupported model: {model_name}" + ) \ No newline at end of file diff --git a/src/models/vit_temporal_pooling.py b/src/models/vit_temporal_pooling.py new file mode 100644 index 0000000..cf29905 --- /dev/null +++ b/src/models/vit_temporal_pooling.py @@ -0,0 +1,95 @@ +""" +DeepVision AI + +Vision Transformer with Temporal Mean Pooling +""" + +import torch +import torch.nn as nn +from torchvision.models import vit_b_16, ViT_B_16_Weights + +from src.models.base_model import BaseModel + + +class ViTTemporalPooling(BaseModel): + """ + Vision Transformer + Temporal Pooling + """ + + def __init__( + self, + image_size=224, + num_classes=2, + pretrained=True, + freeze_backbone=True, + pooling="mean", + hidden_dim=256, + dropout=0.5, + ): + super().__init__() + + self.image_size = image_size + self.pooling = pooling + + # Load ViT backbone + if pretrained: + weights = ViT_B_16_Weights.IMAGENET1K_V1 + else: + weights = None + + self.backbone = vit_b_16(weights=weights) + + # Remove classification head + self.backbone.heads = nn.Identity() + + # Feature dimension of ViT-B/16 + self.feature_dim = 768 + + # Freeze backbone if requested + if freeze_backbone: + for param in self.backbone.parameters(): + param.requires_grad = False + + self.classifier = nn.Sequential( + nn.Linear(self.feature_dim, hidden_dim), + nn.ReLU(inplace=True), + nn.Dropout(dropout), + nn.Linear(hidden_dim, num_classes), + ) + + def temporal_pool(self, features): + """ + features: + (B,T,768) + """ + + if self.pooling == "mean": + return features.mean(dim=1) + + elif self.pooling == "max": + return features.max(dim=1)[0] + + else: + raise ValueError( + f"Unsupported pooling: {self.pooling}" + ) + + def forward(self, x): + """ + x: + (B,T,C,H,W) + """ + + B, T, C, H, W = x.shape + + x = x.reshape(B * T, C, H, W) + + features = self.backbone(x) + + features = features.reshape(B, T, self.feature_dim) + + video_features = self.temporal_pool(features) + + logits = self.classifier(video_features) + + return logits \ No newline at end of file diff --git a/src/training/checkpoint.py b/src/training/checkpoint.py index 17f00e3..3499383 100644 --- a/src/training/checkpoint.py +++ b/src/training/checkpoint.py @@ -15,7 +15,8 @@ class CheckpointManager: """ - Handles experiment tracking and checkpoint saving. + Handles experiment tracking, checkpoint saving, + experiment reproducibility and resume training. """ def __init__(self, root_dir="outputs/runs"): @@ -30,37 +31,152 @@ def __init__(self, root_dir="outputs/runs"): def path(self): return self.run_dir - def save_best_model(self, model): + # -------------------------------------------------- + # Save Best Model + # -------------------------------------------------- + + def save_best_model( + self, + model, + optimizer=None, + scheduler=None, + epoch=None, + best_loss=None, + ): + + checkpoint = { + "epoch": epoch, + "best_loss": best_loss, + "model_state_dict": model.state_dict(), + } + + if optimizer is not None: + checkpoint["optimizer_state_dict"] = optimizer.state_dict() + + if scheduler is not None: + checkpoint["scheduler_state_dict"] = scheduler.state_dict() torch.save( - model.state_dict(), + checkpoint, self.run_dir / "best_model.pth", ) - def save_last_model(self, model): + # -------------------------------------------------- + # Save Last Model + # -------------------------------------------------- + + def save_last_model( + self, + model, + optimizer=None, + scheduler=None, + epoch=None, + best_loss=None, + ): + + checkpoint = { + "epoch": epoch, + "best_loss": best_loss, + "model_state_dict": model.state_dict(), + } + + if optimizer is not None: + checkpoint["optimizer_state_dict"] = optimizer.state_dict() + + if scheduler is not None: + checkpoint["scheduler_state_dict"] = scheduler.state_dict() torch.save( - model.state_dict(), + checkpoint, self.run_dir / "last_model.pth", ) + # -------------------------------------------------- + # Resume Training + # -------------------------------------------------- + + def load_checkpoint( + self, + checkpoint_path, + model, + optimizer=None, + scheduler=None, + map_location="cpu", + ): + + checkpoint = torch.load( + checkpoint_path, + map_location=map_location, + ) + + model.load_state_dict( + checkpoint["model_state_dict"] + ) + + if ( + optimizer is not None + and "optimizer_state_dict" in checkpoint + ): + optimizer.load_state_dict( + checkpoint["optimizer_state_dict"] + ) + + if ( + scheduler is not None + and "scheduler_state_dict" in checkpoint + ): + scheduler.load_state_dict( + checkpoint["scheduler_state_dict"] + ) + + epoch = checkpoint.get("epoch", 0) + + best_loss = checkpoint.get( + "best_loss", + float("inf"), + ) + + return epoch, best_loss + + # -------------------------------------------------- + # Save Metrics + # -------------------------------------------------- + def save_metrics(self, metrics): with open( self.run_dir / "metrics.json", "w", + encoding="utf-8", ) as f: - json.dump(metrics, f, indent=4) + json.dump( + metrics, + f, + indent=4, + ) + + # -------------------------------------------------- + # Save Config + # -------------------------------------------------- def save_config(self, config): with open( self.run_dir / "config.yaml", "w", + encoding="utf-8", ) as f: - yaml.dump(config, f) + yaml.safe_dump( + config, + f, + sort_keys=False, + ) + + # -------------------------------------------------- + # Copy Log + # -------------------------------------------------- def copy_log(self, log_path): @@ -71,4 +187,19 @@ def copy_log(self, log_path): shutil.copy( log_path, self.run_dir / log_path.name, + ) + + # -------------------------------------------------- + # Save Arbitrary File + # -------------------------------------------------- + + def save_file(self, file_path): + + file_path = Path(file_path) + + if file_path.exists(): + + shutil.copy( + file_path, + self.run_dir / file_path.name, ) \ No newline at end of file diff --git a/src/training/optimizer_factory.py b/src/training/optimizer_factory.py new file mode 100644 index 0000000..c53e22b --- /dev/null +++ b/src/training/optimizer_factory.py @@ -0,0 +1,56 @@ +""" +DeepVision AI + +Optimizer Factory +""" + +import torch + + +class OptimizerFactory: + """ + Creates optimizer from configuration. + """ + + @staticmethod + def create(model, config): + + optimizer_name = config["training"]["optimizer"].lower() + + lr = config["training"]["learning_rate"] + + weight_decay = config["training"]["weight_decay"] + + trainable_params = filter( + lambda p: p.requires_grad, + model.parameters() + ) + + if optimizer_name == "adam": + + return torch.optim.Adam( + trainable_params, + lr=lr, + weight_decay=weight_decay, + ) + + if optimizer_name == "adamw": + + return torch.optim.AdamW( + trainable_params, + lr=lr, + weight_decay=weight_decay, + ) + + if optimizer_name == "sgd": + + return torch.optim.SGD( + trainable_params, + lr=lr, + momentum=0.9, + weight_decay=weight_decay, + ) + + raise ValueError( + f"Unsupported optimizer: {optimizer_name}" + ) \ No newline at end of file diff --git a/src/training/scheduler_factory.py b/src/training/scheduler_factory.py new file mode 100644 index 0000000..56d774b --- /dev/null +++ b/src/training/scheduler_factory.py @@ -0,0 +1,45 @@ +""" +DeepVision AI + +Scheduler Factory +""" + +import torch + + +class SchedulerFactory: + """ + Creates LR scheduler from configuration. + """ + + @staticmethod + def create(optimizer, config): + + scheduler_name = config["training"]["scheduler"].lower() + + epochs = config["training"]["epochs"] + + if scheduler_name == "cosineannealinglr": + + return torch.optim.lr_scheduler.CosineAnnealingLR( + optimizer, + T_max=epochs, + ) + + if scheduler_name == "steplr": + + return torch.optim.lr_scheduler.StepLR( + optimizer, + step_size=10, + gamma=0.1, + ) + + if scheduler_name == "reducelronplateau": + + return torch.optim.lr_scheduler.ReduceLROnPlateau( + optimizer, + mode="min", + patience=3, + ) + + return None \ No newline at end of file diff --git a/src/training/trainer.py b/src/training/trainer.py index 0141f7e..e52e6b6 100644 --- a/src/training/trainer.py +++ b/src/training/trainer.py @@ -4,9 +4,8 @@ Generic Trainer """ -from pathlib import Path - import torch +from tqdm import tqdm class Trainer: @@ -55,13 +54,25 @@ def train_one_epoch(self): total = 0 - for batch in self.train_loader: + progress = tqdm( + self.train_loader, + desc="Training", + leave=False, + ) - images = batch["sequence"].to(self.device) + for batch in progress: - labels = batch["label"].to(self.device) + images = batch["sequence"].to( + self.device, + non_blocking=True, + ) - self.optimizer.zero_grad() + labels = batch["label"].to( + self.device, + non_blocking=True, + ) + + self.optimizer.zero_grad(set_to_none=True) outputs = self.model(images) @@ -79,9 +90,13 @@ def train_one_epoch(self): correct += predicted.eq(labels).sum().item() + progress.set_postfix( + loss=f"{loss.item():.4f}" + ) + epoch_loss = running_loss / len(self.train_loader) - epoch_acc = 100 * correct / total + epoch_acc = 100.0 * correct / total return epoch_loss, epoch_acc @@ -97,11 +112,23 @@ def validate(self): with torch.no_grad(): - for batch in self.val_loader: + progress = tqdm( + self.val_loader, + desc="Validation", + leave=False, + ) + + for batch in progress: - images = batch["sequence"].to(self.device) + images = batch["sequence"].to( + self.device, + non_blocking=True, + ) - labels = batch["label"].to(self.device) + labels = batch["label"].to( + self.device, + non_blocking=True, + ) outputs = self.model(images) @@ -115,9 +142,13 @@ def validate(self): correct += predicted.eq(labels).sum().item() + progress.set_postfix( + loss=f"{loss.item():.4f}" + ) + epoch_loss = running_loss / len(self.val_loader) - epoch_acc = 100 * correct / total + epoch_acc = 100.0 * correct / total return epoch_loss, epoch_acc @@ -131,16 +162,25 @@ def train(self, epochs): for epoch in range(epochs): + self.logger.info( + f"Epoch {epoch + 1}/{epochs}" + ) + train_loss, train_acc = self.train_one_epoch() val_loss, val_acc = self.validate() - if self.scheduler: + if self.scheduler is not None: + + if self.scheduler.__class__.__name__ == "ReduceLROnPlateau": - self.scheduler.step() + self.scheduler.step(val_loss) + + else: + + self.scheduler.step() self.logger.info( - f"Epoch {epoch+1}/{epochs} | " f"Train Loss={train_loss:.4f} | " f"Train Acc={train_acc:.2f}% | " f"Val Loss={val_loss:.4f} | " @@ -161,13 +201,31 @@ def train(self, epochs): best_loss = val_loss - self.checkpoint.save_best_model(self.model) - - self.checkpoint.save_last_model(self.model) + self.logger.info( + f"New Best Model | Val Loss = {val_loss:.4f}" + ) + + self.checkpoint.save_best_model( + model=self.model, + optimizer=self.optimizer, + scheduler=self.scheduler, + epoch=epoch + 1, + best_loss=best_loss, + ) + + self.checkpoint.save_last_model( + model=self.model, + optimizer=self.optimizer, + scheduler=self.scheduler, + epoch=epoch + 1, + best_loss=best_loss, + ) if self.early_stopping(val_loss): - self.logger.info("Early stopping triggered.") + self.logger.info( + "Early stopping triggered." + ) break diff --git a/tests/test_model_factory.py b/tests/test_model_factory.py new file mode 100644 index 0000000..1f07d63 --- /dev/null +++ b/tests/test_model_factory.py @@ -0,0 +1,19 @@ +from src.utils.config import load_config +from src.models.model_factory import ModelFactory + + +def main(): + + config = load_config() + + model = ModelFactory.create(config) + + print("\nModel Loaded Successfully\n") + + print(type(model)) + + print(model.name) + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/tests/test_optimizer_scheduler.py b/tests/test_optimizer_scheduler.py new file mode 100644 index 0000000..198f22d --- /dev/null +++ b/tests/test_optimizer_scheduler.py @@ -0,0 +1,34 @@ +import torch.nn as nn + +from src.utils.config import load_config +from src.training.optimizer_factory import OptimizerFactory +from src.training.scheduler_factory import SchedulerFactory + + +def main(): + + config = load_config() + + model = nn.Linear(10, 2) + + optimizer = OptimizerFactory.create( + model, + config, + ) + + scheduler = SchedulerFactory.create( + optimizer, + config, + ) + + print("Optimizer") + print(type(optimizer)) + + print() + + print("Scheduler") + print(type(scheduler)) + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/tests/test_training_pipeline.py b/tests/test_training_pipeline.py new file mode 100644 index 0000000..67e175f --- /dev/null +++ b/tests/test_training_pipeline.py @@ -0,0 +1,88 @@ +""" +DeepVision AI + +End-to-End Training Smoke Test +""" + +import torch +import torch.nn as nn + +from src.utils.config import load_config +from src.datasets.dataloader import create_dataloaders +from src.models.model_factory import ModelFactory + + +def main(): + + print("=" * 60) + print(" DeepVision AI - Integration Test ") + print("=" * 60) + + config = load_config() + + print("\nLoading dataloaders...") + + train_loader, _, _ = create_dataloaders(config) + + batch = next(iter(train_loader)) + + sequences = batch["sequence"] + + labels = batch["label"] + + print(f"Input Shape : {sequences.shape}") + + print("\nLoading model...") + + model = ModelFactory.create(config) + + model.train() + + criterion = nn.CrossEntropyLoss() + + optimizer = torch.optim.AdamW( + model.parameters(), + lr=config["training"]["learning_rate"], + ) + + print("\nForward Pass...") + + outputs = model(sequences) + + print(f"Output Shape : {outputs.shape}") + + print("\nComputing Loss...") + + loss = criterion(outputs, labels) + + print(f"Loss : {loss.item():.4f}") + + print("\nBackward Pass...") + + optimizer.zero_grad() + + loss.backward() + + optimizer.step() + + print("\nBackward Pass Successful") + + print("\nChecking Gradients...") + + gradients = 0 + + for name, param in model.named_parameters(): + + if param.grad is not None: + + gradients += 1 + + print(f"Layers with gradients : {gradients}") + + print("\nIntegration Test PASSED") + + print("=" * 60) + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/train.py b/train.py new file mode 100644 index 0000000..9532ed1 --- /dev/null +++ b/train.py @@ -0,0 +1,374 @@ +""" +============================================================= +DeepVision AI + +Main Training Script + +Author : Arvindh Babu +============================================================= +""" + +import argparse +from pathlib import Path + +import torch +import torch.nn as nn + +from src.utils.config import load_config +from src.utils.logger import create_logger + +from src.datasets.dataloader import create_dataloaders + +from src.models.model_factory import ModelFactory + +from src.training.optimizer_factory import OptimizerFactory +from src.training.scheduler_factory import SchedulerFactory + +from src.training.checkpoint import CheckpointManager +from src.training.early_stopping import EarlyStopping +from src.training.trainer import Trainer + + +# ============================================================ +# Argument Parser +# ============================================================ + +def parse_arguments(): + """ + Parse command-line arguments. + """ + + parser = argparse.ArgumentParser( + description="DeepVision AI Training" + ) + + parser.add_argument( + "--config", + type=str, + default="configs/default.yaml", + help="Path to configuration YAML", + ) + + parser.add_argument( + "--resume", + type=str, + default=None, + help="Checkpoint to resume training", + ) + + return parser.parse_args() + + +# ============================================================ +# Device +# ============================================================ + +def get_device(): + + if torch.cuda.is_available(): + + device = torch.device("cuda") + + print("=" * 60) + print("CUDA AVAILABLE") + print("=" * 60) + + print( + "GPU :", + torch.cuda.get_device_name(0), + ) + + print( + "CUDA Version :", + torch.version.cuda, + ) + + print("=" * 60) + + else: + + device = torch.device("cpu") + + print("=" * 60) + print("Running on CPU") + print("=" * 60) + + return device + + +# ============================================================ +# Main +# ============================================================ + +def main(): + + args = parse_arguments() + + config = load_config(args.config) + + logger = create_logger() + + logger.info("=" * 60) + logger.info("DeepVision AI") + logger.info("=" * 60) + + device = get_device() + + logger.info(f"Device : {device}") + + # -------------------------------------------------------- + # Dataloaders + # -------------------------------------------------------- + + logger.info("Creating dataloaders...") + + train_loader, val_loader, test_loader = create_dataloaders( + config + ) + + logger.info("DataLoaders Ready") + + logger.info( + f"Training batches : {len(train_loader)}" + ) + + logger.info( + f"Validation batches : {len(val_loader)}" + ) + + logger.info( + f"Test batches : {len(test_loader)}" + ) + + # -------------------------------------------------------- + # Model + # -------------------------------------------------------- + + logger.info("Creating model...") + + model = ModelFactory.create(config) + + model = model.to(device) + + logger.info(f"Model : {model.__class__.__name__}") + + total_params = sum( + p.numel() for p in model.parameters() + ) + + trainable_params = sum( + p.numel() + for p in model.parameters() + if p.requires_grad + ) + + logger.info( + f"Total Parameters : {total_params:,}" + ) + + logger.info( + f"Trainable Parameters : {trainable_params:,}" + ) + + # -------------------------------------------------------- + # Loss + # -------------------------------------------------------- + + criterion = nn.CrossEntropyLoss() + + logger.info("Loss : CrossEntropyLoss") + + # -------------------------------------------------------- + # Optimizer + # -------------------------------------------------------- + + optimizer = OptimizerFactory.create( + model, + config, + ) + + logger.info( + f"Optimizer : {optimizer.__class__.__name__}" + ) + + # -------------------------------------------------------- + # Scheduler + # -------------------------------------------------------- + + scheduler = SchedulerFactory.create( + optimizer, + config, + ) + + if scheduler is not None: + + logger.info( + f"Scheduler : {scheduler.__class__.__name__}" + ) + + else: + + logger.info("Scheduler : None") + + # -------------------------------------------------------- + # Checkpoint Manager + # -------------------------------------------------------- + + checkpoint_manager = CheckpointManager() + + checkpoint_manager.save_config(config) + + logger.info( + f"Experiment Folder : {checkpoint_manager.path}" + ) + + # -------------------------------------------------------- + # Early Stopping + # -------------------------------------------------------- + + early_stopping = EarlyStopping( + patience=5 + ) + + # -------------------------------------------------------- + # Resume Training + # -------------------------------------------------------- + + if args.resume is not None: + + logger.info( + f"Loading checkpoint : {args.resume}" + ) + + epoch, best_loss = checkpoint_manager.load_checkpoint( + checkpoint_path=args.resume, + model=model, + optimizer=optimizer, + scheduler=scheduler, + map_location=device, + ) + + logger.info( + f"Checkpoint Loaded (Epoch {epoch})" + ) + + # -------------------------------------------------------- + # Trainer + # -------------------------------------------------------- + + trainer = Trainer( + + model=model, + + optimizer=optimizer, + + criterion=criterion, + + train_loader=train_loader, + + val_loader=val_loader, + + device=device, + + logger=logger, + + checkpoint_manager=checkpoint_manager, + + early_stopping=early_stopping, + + scheduler=scheduler, + ) + + # -------------------------------------------------------- + # Start Training + # -------------------------------------------------------- + + logger.info("=" * 60) + logger.info("Starting Training") + logger.info("=" * 60) + + history = trainer.train( + epochs=config["training"]["epochs"] + ) + + # -------------------------------------------------------- + # Save Training Log + # -------------------------------------------------------- + + try: + + checkpoint_manager.copy_log( + "outputs/logs/train.log" + ) + + except Exception as e: + + logger.warning( + f"Unable to copy log file: {e}" + ) + + # -------------------------------------------------------- + # Training Summary + # -------------------------------------------------------- + + logger.info("=" * 60) + logger.info("Training Completed Successfully") + logger.info("=" * 60) + + logger.info( + f"Experiment Directory : {checkpoint_manager.path}" + ) + + logger.info( + f"Epochs Completed : {len(history)}" + ) + + if len(history) > 0: + + best_epoch = min( + history, + key=lambda x: x["val_loss"] + ) + + logger.info( + f"Best Validation Loss : {best_epoch['val_loss']:.4f}" + ) + + logger.info( + f"Best Validation Accuracy : {best_epoch['val_acc']:.2f}%" + ) + + logger.info( + "Artifacts Saved:" + ) + + logger.info( + f" • Best Model : {checkpoint_manager.path / 'best_model.pth'}" + ) + + logger.info( + f" • Last Model : {checkpoint_manager.path / 'last_model.pth'}" + ) + + logger.info( + f" • Metrics : {checkpoint_manager.path / 'metrics.json'}" + ) + + logger.info( + f" • Config : {checkpoint_manager.path / 'config.yaml'}" + ) + + logger.info( + f" • Training Log : {checkpoint_manager.path / 'train.log'}" + ) + + logger.info("=" * 60) + + +# ============================================================ +# Entry Point +# ============================================================ + +if __name__ == "__main__": + + main() \ No newline at end of file diff --git a/vit_temporal_pooling.zip b/vit_temporal_pooling.zip new file mode 100644 index 0000000000000000000000000000000000000000..26fe3cd319aae4c753a4733a4d49c2c35a945f80 GIT binary patch literal 9450 zcmbtZ1yq$;*S>U1ONY{35*L&bq~X#?cXxLQNOwphEg=X}lF}W5AR!2NJ`yMM9w(-LXBJC~TF|A&Z2294%^vG802=uy5%0AFxV^o@nzkZ^SAU;u&`5=`Nvj?o|gGIV#|X)IDnSD7~w*wM`j>$GNI_v{+WC zV(}&D(H4`#2hH*UHrU{kdbYYZ4Eav%92LE6uZV8=is-|qG06C)0sAZQ9hJA>0b1_Z z_l+^nZWqL_&!}{@kUys%sh|Zyi`+VVeFUEo9t~38Q=yGu#~DM;H$-|W>~8tk%2_7& z`}5eY$;U$nbDvii%TA1D`gXztpRxJgSFG+GQ*gbFHcq?b)=F7Ol1If(N@ZSllIgvS zE;EX@BFyN+6y6#4(1=Mv-4V7YJWxSR-uF#Gcm(`C^{#Af!RXF3wG%D6u*fpTk;K+V z-6Yk0=JPk0vfn#xJ<1kFzipzWTX7HbV{3&=+yUFnjBNI|{zZaJ?+TP7x$*e?olPry zkJqodsJaKMMQFMx()F5y<08uKo3*F8U~)#*_TJ4xrFODLeK0`P(RYnFU3_PEo3J`W zBMhsaBJHWxEtw&f3G$b)WufD!H8N_Ywr%;gZwDr|UIr94w zh#@3M2O8}(99g_R*s&tv?e{Avrl)BQt(ByJ6l?z zUf|Msy!{?i$H2$sEQb`WI4k7QU%ZFuNoW! z6WN>J2QiMY8)Fit;*i)O@BWe!L-$VpmyXXZtkrknl(Ujd{VW3H zg^+oTpqpP)EvV?Rq!A{Vn%mV-Za7q`JDR_IoS}h>o0~hsR|i)xuInhH#p$P9_8{xH zwaCuEgzCA(0Qmx^UzV)LT-OxPLg&=La6?Vl+HoxLpo_#OtJIS}PO?5tjLtEFVHeei z{)Rg`*_rLx>}li3xZKHPV*aqQ%p0Y8^g!vbLLT%I1Aktm-BeX7x}u*V`#|E?*P#4^W#J%QYwIJ)$?Or(C+zf9eT^4wZkkN%+!8l5G69}^A}Jvu4nzF6lE8J=o`CS( z&Kh>B!NCN-9yV4l=-xK3sl4$07hwe3#TxIU(O8S^UyM1q%BnEJk~RB5$ecd`ilD93JvDBF#G6pag5H zZ-dNDTR{wNFYcYVn;7s+;eWP9C9_0DKRMIv^ASw+<%MPkVNx~0WsyDRcs6N9ITU}G zo0&OTj&O>x`4V|Q3J2~4twr~YskT* z9Q7J*mcC<6v_p+7cIR(mlzf3pStrlDvG_Wo6TW$po}H6_+?vBG|Cu_+Uh}M_Tn=&J z%gKtB4F}U>J4Hq8g>+9<8QW%y;?voe{4S}vT6bHx*P9{Xt7jkwQ62^Wz=FFNzF@ey zFD2Fg?&0j$c{uZ+{f}%LsH`DJB!-hQ8P>?dMH!v{Zhl+_GeXfez+IkMMv)xpmW*ji zH2K5l4>Z2uCr)W|)Cc))QlpDkC5Z48(TA%ek!r-~#5Un1bB)|l&e;#)rjMoHkm9S{ z=S@W25YGT|tnwuRJ({+o_%jj8*?c;O$s_?MNQ6O}8ZOT(%m_>cA4y=8K4+y8f$KK9 zMVFsj7$2L%4hxMBl~Rj%2MJ|Vd-%?H#|EL(P+NIJN5_kD!GwqJg)?*f4%SbJj(7<5 zu~J%&mg0>%i3UVzaHE4sWvIlynEZs5RDR!& z<@Cz{ceBOf>6f1sNaod~JyNO)NvUX}XWz4pl6ZLBL*X~z9;iH85>cvYtd}$I-vCv5 zJ*uvJSR9h_lKIXp7?0Hv-;LRh%+0Y+J`0k-wRWwfB8Fne)SOHNy$LmZg#*p1sKJjX zL^dy25UGoU3rt#%pII7nJORPox;NmCfjNs9aGcXL()WmBwR126-80-86%_N87Wa)= zQMfse)%y6|go5&-B_z1A&}Cs&`;$zg=!u`X;I@!?nrDKiOE+mvDu z)tm{n(HKm~B#Kxj`WYLTM7VNw=%>D5Xn)daM{SC4g0+EHt79CIOuJverutHyv`KCZ@Ev<%v&Y>)fvGfL1yn)A8`#zXZ{ z?m=AOZ@wUuH(G%t@UK+sRbDlr-OQ!RE1a?V-Urw4nDqEjaXLeK)QK0-X_L-LOzcVR z!Kc`h-Ba!wQ4F5W{#X9R-x1EUo7ZIc+uX5lq&drj+LHp;ZNzvd;*l^B$56VGhdC=ar{W zm)R?!e7KvpMICq_CsL6|`U4y?O<`?mKXPDWf0r*Z863iZ$DLq9!MV%`*Lbi?% z&1WJO4Y8iGQtlhMCJR)c2x<2?)5^HvUu_8@t*9py}PvDwEz_*C6ey11V?UDc`4dATUV5H7oBLL$RxuMxl z)f<5(qb{<=V0mBpEc6Vhn{9RWN2OReI^{V9t`fz-n(Sw#`1+ShQSe_BQ$MQ4->9T2 z(0-PXeXGmM`T+mcyNAi-M{F>Jy0H2PAQ%i4osEl4RipzMTUIV_;;>^iXTzA>=hRa9^Y zHCR~W*!6w0?GBev!(-{|p9g&AELtMihwoR5-~O$;zq?zl9n}!(52Dh#gUM3CBsE?l zYVsG|PI}#KQ4kdpkr8FKGQ6s)*8dIF4(UqhCR|oAB-f!i+UuEFU!z^Ld~ANS2-211 zZKk;}yt(JkY(YA-O3|Zwa_}efPMzlD1O8(c2(VB>iZ5l9Zo|deCpsN|p>T?CeFtBy z;;R>W^7(W9m(R45FQ<5Hr?%Q0Z^zXr(COZcPm#@!OCLca&dZU7MJzTec4B(lWa&s{ zt`SsO3pdhdQa35G$CLppXESg8wW|QPH_okxo)7)y#7n)rpU(%Hf7+EDS48 za0e4LR5&_BMVTz7rp~(JZVgt)aaxb3D7U_AmsriE>8hIQiBT1lpFs}DkCWVNA)vGfBFOLVz1pjT< zedyb9gQ3XPa^+v$T=ihHLW858h)TvACbegN&^v5T;f-*i@dSwDKvkT&9v{CqcB|~c z99xyzuJz412(cjBmzML_iv^IuXnC$Ddlc`QW54B$pXgWD!xLD}8-)!dCb6SPXH{UZ zDN{O}HMLOfC3w%TbxcLz_sXdv#^+!UbnIQAsZw$%jKL90cl^YjCXi0UxM7o-G6s*O zc8Za;;J5^@wQ_*;bLEsFI?{%G_xIN$lD7Mfn3>QO2zQ8qVJf?iz6P*7T+k~Q+L9%@ z19;ZLT7?ID>%Lf%As*jXqiM5=q@f^P)#bPI4qG1a3B%QZx$LGcVc(bELe93x-|@AI zD!gpkT1tZEf@@%zzYnFeGf~afa+{eM$J4dbJ@K?)eYQL_jb}97f)B6Zs~Pr^10iM{ z>2a_~@vTV|Tc7l;vZ($v1Y+Pr8g=_lV9ensXsyG-0Sj`z&yOW3zY&ZH$0&WJbbfNL zv#tiidyQywld{BIuk+Pvuhb^9=1Yc2ZN(6JbhrL%;*U&n zb(@duLh|@G{fArE1q>aX|E2L@`%66H`km$fj7`6hA-%wX>b!SqrK}M~EQ<5IsOeyD zaL}5I+CpSSv=uR#ik9}4xIPwa40)}HB{x@R$Vi48TFOkyvzy@Lk@>2eP9=pNm==Apj!-t z03usk6KF$}7}E_I6gMb|*}@-Wqp;AVFk?3ev|wZkRQ(n_mdlkpAZ^^(5|M+^&z3js zYN4c2J}W>(yUS==Jua?i{Dop=Q^%ZH=}VyebZS#`l)2cAgvwU*kmpv_6YHs*fVLK+YU*91@;+2D5+?k;imJH<+s4ZaPvpS58-vsy#0mp()(Y z@OpN)uAm!vPVg&8JMe9gFhkC4;+Jy4I?IoQ?jiW?u=W~glWTcB&STZqK2ECmIVK15 z`Ok%NCnkbFJ?9IDN@s}D;K<-0B(wwSMB;^nFTrOQZnE*N={?aNX7tMAF!65_;vdT9 z$y4qq9!~J52KHM@&`-?+<%|p1_3!56Vc#ZJyG4#9Z?RhuELS{9l_<4(EH=gN4abw+ zy!G1to1IU8rp)W2!^@8x7#7Yfrm8TN$atKC0DU6Rn@8J=LLzC3 znZc{q?4`}=+rqqnO%k(#St>+vQ`ZhoNAoEj64LB? zj)_mvM#W|K0=?suuN-Drh8W|$3tv5okvKG#uCTLSzh`OjclV)<{jdmbQbaH5gB?Fp zLE5pb$h0F&R^PqaTNZnL1%j+HbLg+lp{hzdK}R3TeZ3ov8FCr33MsaepAkQJvh5F; zL|djmPP;Qs$muL52eW3%eS>2H-+iU>lZ~Gyb$kMpw(cH0=o42_lj{wUv1Pu!OfUU= z@QJ4Id#zsxjLRuMqGH!@Izq9~oFE?MiEo&{?Vg_S>;1cw4`^ zjTA&-;|T*Eq}&8F)jm)EWU+0XM36UMvJw5M!$VyY&Q5i1>G((;3Z6`G&5vN=9((JO z+*uVL{`U_zUe%Qni6##Ynquui)sr9U}*%``zWE zO#h-r<}m2)))qu^srS>LpQ8g0&>>l2-+H324J%wC7-d9vWb$jhP;`Uji!lQK!ApWEv${2TOjWbl_7aS zW!$Kbu&@Ic&T*HD=U-LE|7!QYkrz7+e`>6oTgh)W}I*UC`c9sN&UI>#p&kx9FiRHbLnb7#p{;-yICg~1v2Zr zh;lyDyc(tLI+VYgaDq`F6V7uK2nqjxlKIn?Ga(MBYiPP|u6YGpkcsBSsS^Dqwj%xk z;o^Yhiee6o2bpi4TLd|s^I6q7Z#0+iD8S+4%7g#Qq91YLYB0#;>KyDU`e-kJoe8D< zBiP@RjaN_!nN(fuaNoaxR_zA>_q_k!w4Z1fY5k`N2;K-}Zgp-pgkwnPx!C`?ypaWP z(f;2ya&<@oZv@g_zqk=dyXj&NM(`3NGBxQH|T&gpDykKLdV4h$K_opp#P^` z{5BWo&6M+tPkM<94~%PeaeZqA%r~U9a?w&q3+19VBe{gJejUcO4HYm1q@i*V;#}mN zn{he17J;Mtm3QmM`CaX}5XpCra+TxWx?sqe5FzfhC_jDZ|FxpP#{8NpAQJCl)xvcF zrrOW#cVK=`@PV-*g74f!h?7F Y4gMdz0S+E=xLDvXFadxuV#t602ZcxsaR2}S literal 0 HcmV?d00001