Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions configs/default.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,12 @@ model:
temporal_pooling: mean
num_classes: 2

dataloader:
batch_size: 8
num_workers: 4
pin_memory: true
drop_last: false
Comment on lines +36 to +40

evaluation:
metrics:
- accuracy
Expand Down
74 changes: 74 additions & 0 deletions src/training/checkpoint.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
"""
DeepVision AI

Checkpoint Manager
"""

from pathlib import Path
from datetime import datetime
import json
import shutil

import torch
import yaml


class CheckpointManager:
"""
Handles experiment tracking and checkpoint saving.
"""

def __init__(self, root_dir="outputs/runs"):

timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")

self.run_dir = Path(root_dir) / f"run_{timestamp}"

self.run_dir.mkdir(parents=True, exist_ok=True)

@property
def path(self):
return self.run_dir

def save_best_model(self, model):

torch.save(
model.state_dict(),
self.run_dir / "best_model.pth",
)

def save_last_model(self, model):

torch.save(
model.state_dict(),
self.run_dir / "last_model.pth",
)

def save_metrics(self, metrics):

with open(
self.run_dir / "metrics.json",
"w",
) as f:

json.dump(metrics, f, indent=4)

def save_config(self, config):

with open(
self.run_dir / "config.yaml",
"w",
) as f:

yaml.dump(config, f)

def copy_log(self, log_path):

log_path = Path(log_path)

if log_path.exists():

shutil.copy(
log_path,
self.run_dir / log_path.name,
)
47 changes: 47 additions & 0 deletions src/training/early_stopping.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
"""
DeepVision AI

Early Stopping Utility
"""

import numpy as np


class EarlyStopping:
"""
Stops training when validation loss
stops improving.
"""

def __init__(
self,
patience=5,
min_delta=0.0,
):

self.patience = patience
self.min_delta = min_delta

self.best_loss = np.inf

self.counter = 0

self.should_stop = False

def __call__(self, val_loss):

if val_loss < self.best_loss - self.min_delta:

self.best_loss = val_loss

self.counter = 0

else:

self.counter += 1

if self.counter >= self.patience:

self.should_stop = True

return self.should_stop
176 changes: 176 additions & 0 deletions src/training/trainer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,176 @@
"""
DeepVision AI

Generic Trainer
"""

from pathlib import Path

import torch
Comment on lines +7 to +9


class Trainer:

def __init__(
self,
model,
optimizer,
criterion,
train_loader,
val_loader,
device,
logger,
checkpoint_manager,
early_stopping,
scheduler=None,
):

self.model = model.to(device)

self.optimizer = optimizer

self.criterion = criterion

self.train_loader = train_loader

self.val_loader = val_loader

self.device = device

self.logger = logger

self.checkpoint = checkpoint_manager

self.early_stopping = early_stopping

self.scheduler = scheduler

def train_one_epoch(self):

self.model.train()

running_loss = 0.0

correct = 0

total = 0

for batch in self.train_loader:

images = batch["sequence"].to(self.device)

labels = batch["label"].to(self.device)

self.optimizer.zero_grad()

outputs = self.model(images)
Comment on lines +60 to +66

loss = self.criterion(outputs, labels)

loss.backward()

self.optimizer.step()

running_loss += loss.item()

_, predicted = outputs.max(1)

total += labels.size(0)

correct += predicted.eq(labels).sum().item()

epoch_loss = running_loss / len(self.train_loader)

epoch_acc = 100 * correct / total

return epoch_loss, epoch_acc
Comment on lines +82 to +86

def validate(self):

self.model.eval()

running_loss = 0.0

correct = 0

total = 0

with torch.no_grad():

for batch in self.val_loader:

images = batch["sequence"].to(self.device)

labels = batch["label"].to(self.device)

outputs = self.model(images)
Comment on lines +102 to +106

loss = self.criterion(outputs, labels)

running_loss += loss.item()

_, predicted = outputs.max(1)

total += labels.size(0)

correct += predicted.eq(labels).sum().item()

epoch_loss = running_loss / len(self.val_loader)

epoch_acc = 100 * correct / total

return epoch_loss, epoch_acc
Comment on lines +118 to +122

def train(self, epochs):

best_loss = float("inf")

history = []

self.logger.info("Training Started")

for epoch in range(epochs):

train_loss, train_acc = self.train_one_epoch()

val_loss, val_acc = self.validate()

if self.scheduler:

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} | "
f"Val Acc={val_acc:.2f}%"
)

history.append(
{
"epoch": epoch + 1,
"train_loss": train_loss,
"train_acc": train_acc,
"val_loss": val_loss,
"val_acc": val_acc,
}
)

if val_loss < best_loss:

best_loss = val_loss

self.checkpoint.save_best_model(self.model)

self.checkpoint.save_last_model(self.model)

if self.early_stopping(val_loss):

self.logger.info("Early stopping triggered.")

break

self.checkpoint.save_metrics(history)

return history
50 changes: 50 additions & 0 deletions src/utils/logger.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
"""
DeepVision AI

Logger Utility
"""

from pathlib import Path
import logging


def create_logger(
log_dir="outputs/logs",
log_name="train.log",
):
"""
Create a reusable logger.
"""

log_dir = Path(log_dir)

log_dir.mkdir(
parents=True,
exist_ok=True,
)

logger = logging.getLogger("DeepVision")

logger.setLevel(logging.INFO)

logger.handlers.clear()

Comment on lines +28 to +31
formatter = logging.Formatter(
"%(asctime)s | %(levelname)s | %(message)s"
)

file_handler = logging.FileHandler(
log_dir / log_name
)

file_handler.setFormatter(formatter)

console_handler = logging.StreamHandler()

console_handler.setFormatter(formatter)

logger.addHandler(file_handler)

logger.addHandler(console_handler)

return logger
Loading