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 0000000..26fe3cd Binary files /dev/null and b/vit_temporal_pooling.zip differ