From 136fa9d694825e49d116329d5228c40d5c81917f Mon Sep 17 00:00:00 2001 From: Arvindhbabu Date: Wed, 8 Jul 2026 16:37:02 +0530 Subject: [PATCH 1/4] feat(training): implement training framework foundation --- src/training/base_trainer.py | 126 +++++++++++++++++++++++++++++++++ src/training/callbacks.py | 65 ++++++++++++++++++ src/training/history.py | 115 +++++++++++++++++++++++++++++++ src/training/metrics.py | 130 +++++++++++++++++++++++++++++++++++ src/training/state.py | 108 +++++++++++++++++++++++++++++ tests/test_base_trainer.py | 26 +++++++ tests/test_callbacks.py | 43 ++++++++++++ tests/test_history.py | 48 +++++++++++++ tests/test_metrics.py | 40 +++++++++++ tests/test_state.py | 39 +++++++++++ 10 files changed, 740 insertions(+) create mode 100644 src/training/base_trainer.py create mode 100644 src/training/history.py create mode 100644 src/training/metrics.py create mode 100644 src/training/state.py create mode 100644 tests/test_base_trainer.py create mode 100644 tests/test_callbacks.py create mode 100644 tests/test_history.py create mode 100644 tests/test_metrics.py create mode 100644 tests/test_state.py diff --git a/src/training/base_trainer.py b/src/training/base_trainer.py new file mode 100644 index 0000000..cfffbbb --- /dev/null +++ b/src/training/base_trainer.py @@ -0,0 +1,126 @@ +""" +DeepVision AI + +Base Trainer + +Defines the common functionality shared by all trainers. +""" + +from abc import ABC, abstractmethod +import time + +import torch + +from src.training.history import History +from src.training.metrics import Metrics +from src.training.callbacks import CallbackManager + + +class BaseTrainer(ABC): + """ + Base class for all DeepVision AI trainers. + """ + + def __init__( + self, + model, + optimizer, + criterion, + train_loader, + val_loader, + device, + logger=None, + scheduler=None, + checkpoint_manager=None, + early_stopping=None, + ): + + self.model = model + self.optimizer = optimizer + self.criterion = criterion + + self.train_loader = train_loader + self.val_loader = val_loader + + self.device = device + + self.logger = logger + self.scheduler = scheduler + + self.checkpoint_manager = checkpoint_manager + self.early_stopping = early_stopping + + # Training state + self.current_epoch = 0 + self.best_val_loss = float("inf") + + # Utilities + self.history = History() + self.metrics = Metrics() + self.callbacks = CallbackManager() + + # AMP (safe on CPU) + self.use_amp = ( + torch.cuda.is_available() + and device.type == "cuda" + ) + + self.scaler = torch.cuda.amp.GradScaler( + enabled=self.use_amp + ) + + # -------------------------------------------------- + # Logging + # -------------------------------------------------- + + def log(self, message): + + if self.logger is not None: + self.logger.info(message) + else: + print(message) + + # -------------------------------------------------- + # Current Learning Rate + # -------------------------------------------------- + + def get_learning_rate(self): + + return self.optimizer.param_groups[0]["lr"] + + # -------------------------------------------------- + # Epoch Timer + # -------------------------------------------------- + + def start_timer(self): + + self._start_time = time.time() + + def stop_timer(self): + + return time.time() - self._start_time + + # -------------------------------------------------- + # Abstract Methods + # -------------------------------------------------- + + @abstractmethod + def train_one_epoch(self): + """ + Train for one epoch. + """ + pass + + @abstractmethod + def validate(self): + """ + Validate one epoch. + """ + pass + + @abstractmethod + def train(self, epochs): + """ + Complete training loop. + """ + pass \ No newline at end of file diff --git a/src/training/callbacks.py b/src/training/callbacks.py index e69de29..450f423 100644 --- a/src/training/callbacks.py +++ b/src/training/callbacks.py @@ -0,0 +1,65 @@ +""" +DeepVision AI + +Training Callback System +""" + + +class Callback: + """ + Base callback class. + """ + + def on_train_begin(self, trainer): + pass + + def on_epoch_begin(self, trainer): + pass + + def on_epoch_end(self, trainer): + pass + + def on_validation_end(self, trainer): + pass + + def on_train_end(self, trainer): + pass + + +class CallbackManager: + """ + Manages all callbacks. + """ + + def __init__(self): + + self.callbacks = [] + + def add(self, callback): + + self.callbacks.append(callback) + + def on_train_begin(self, trainer): + + for cb in self.callbacks: + cb.on_train_begin(trainer) + + def on_epoch_begin(self, trainer): + + for cb in self.callbacks: + cb.on_epoch_begin(trainer) + + def on_epoch_end(self, trainer): + + for cb in self.callbacks: + cb.on_epoch_end(trainer) + + def on_validation_end(self, trainer): + + for cb in self.callbacks: + cb.on_validation_end(trainer) + + def on_train_end(self, trainer): + + for cb in self.callbacks: + cb.on_train_end(trainer) \ No newline at end of file diff --git a/src/training/history.py b/src/training/history.py new file mode 100644 index 0000000..b69d311 --- /dev/null +++ b/src/training/history.py @@ -0,0 +1,115 @@ +""" +DeepVision AI + +Training History Manager +""" + +from pathlib import Path +import json +import csv + + +class History: + """ + Stores and manages training history. + """ + + def __init__(self): + self.records = [] + + # -------------------------------------------------- + # Add epoch + # -------------------------------------------------- + + def add(self, **kwargs): + """ + Add one epoch record. + + Example: + history.add( + epoch=1, + train_loss=0.4, + val_loss=0.3, + train_accuracy=0.95, + val_accuracy=0.94, + ) + """ + self.records.append(kwargs) + + # -------------------------------------------------- + # Length + # -------------------------------------------------- + + def __len__(self): + return len(self.records) + + # -------------------------------------------------- + # Best Epoch + # -------------------------------------------------- + + def best(self, metric="val_loss", mode="min"): + + if len(self.records) == 0: + return None + + if mode == "min": + return min( + self.records, + key=lambda x: x[metric] + ) + + return max( + self.records, + key=lambda x: x[metric] + ) + + # -------------------------------------------------- + # Save JSON + # -------------------------------------------------- + + def save_json(self, path): + + path = Path(path) + + with open(path, "w", encoding="utf-8") as f: + + json.dump( + self.records, + f, + indent=4, + ) + + # -------------------------------------------------- + # Save CSV + # -------------------------------------------------- + + def save_csv(self, path): + + path = Path(path) + + if len(self.records) == 0: + return + + with open( + path, + "w", + newline="", + encoding="utf-8", + ) as f: + + writer = csv.DictWriter( + f, + fieldnames=self.records[0].keys(), + ) + + writer.writeheader() + + writer.writerows(self.records) + + # -------------------------------------------------- + # Return records + # -------------------------------------------------- + + def get(self): + + return self.records \ No newline at end of file diff --git a/src/training/metrics.py b/src/training/metrics.py new file mode 100644 index 0000000..d2557c9 --- /dev/null +++ b/src/training/metrics.py @@ -0,0 +1,130 @@ +""" +DeepVision AI + +Metrics Engine + +Computes classification metrics for training, +validation and evaluation. +""" + +from typing import Dict + +import numpy as np +import torch + +from sklearn.metrics import ( + accuracy_score, + precision_score, + recall_score, + f1_score, +) + + +class Metrics: + """ + Metrics accumulator. + + Example + ------- + metrics = Metrics() + + metrics.update(outputs, labels) + + result = metrics.compute() + + metrics.reset() + """ + + def __init__(self): + + self.reset() + + # -------------------------------------------------------- + # Reset + # -------------------------------------------------------- + + def reset(self): + + self.targets = [] + + self.predictions = [] + + # -------------------------------------------------------- + # Update + # -------------------------------------------------------- + + def update( + self, + outputs: torch.Tensor, + labels: torch.Tensor, + ): + + preds = torch.argmax( + outputs, + dim=1, + ) + + self.predictions.extend( + preds.detach().cpu().numpy().tolist() + ) + + self.targets.extend( + labels.detach().cpu().numpy().tolist() + ) + + # -------------------------------------------------------- + # Compute + # -------------------------------------------------------- + + def compute(self) -> Dict[str, float]: + + y_true = np.array(self.targets) + + y_pred = np.array(self.predictions) + + results = { + + "accuracy": accuracy_score( + y_true, + y_pred, + ), + + "precision": precision_score( + y_true, + y_pred, + zero_division=0, + ), + + "recall": recall_score( + y_true, + y_pred, + zero_division=0, + ), + + "f1": f1_score( + y_true, + y_pred, + zero_division=0, + ), + } + + return results + + # -------------------------------------------------------- + # Pretty Print + # -------------------------------------------------------- + + @staticmethod + def print(results: Dict[str, float]): + + print("=" * 50) + + print("DeepVision AI Metrics") + + print("=" * 50) + + for k, v in results.items(): + + print(f"{k:12}: {v:.4f}") + + print("=" * 50) \ No newline at end of file diff --git a/src/training/state.py b/src/training/state.py new file mode 100644 index 0000000..7bb4540 --- /dev/null +++ b/src/training/state.py @@ -0,0 +1,108 @@ +""" +DeepVision AI + +Training State Manager +""" + +from dataclasses import dataclass, asdict + + +@dataclass +class TrainingState: + """ + Stores the current state of training. + + This object is saved and restored when + resuming training from checkpoints. + """ + + # ----------------------------- + # Progress + # ----------------------------- + + epoch: int = 0 + + global_step: int = 0 + + # ----------------------------- + # Best Metrics + # ----------------------------- + + best_val_loss: float = float("inf") + + best_val_accuracy: float = 0.0 + + # ----------------------------- + # Current Metrics + # ----------------------------- + + train_loss: float = 0.0 + + val_loss: float = 0.0 + + train_accuracy: float = 0.0 + + val_accuracy: float = 0.0 + + # ----------------------------- + # Learning Rate + # ----------------------------- + + learning_rate: float = 0.0 + + # ----------------------------- + # Runtime + # ----------------------------- + + epoch_time: float = 0.0 + + total_training_time: float = 0.0 + + # ----------------------------- + # Flags + # ----------------------------- + + resumed: bool = False + + stopped_early: bool = False + + training_completed: bool = False + + # -------------------------------------------------- + # Convert to Dictionary + # -------------------------------------------------- + + def to_dict(self): + """ + Convert state to dictionary. + """ + + return asdict(self) + + # -------------------------------------------------- + # Update + # -------------------------------------------------- + + def update(self, **kwargs): + """ + Update multiple fields. + """ + + for key, value in kwargs.items(): + + if hasattr(self, key): + + setattr(self, key, value) + + # -------------------------------------------------- + # Reset + # -------------------------------------------------- + + def reset(self): + """ + Reset runtime values while preserving defaults. + """ + + self.__dict__.update( + TrainingState().__dict__ + ) \ No newline at end of file diff --git a/tests/test_base_trainer.py b/tests/test_base_trainer.py new file mode 100644 index 0000000..61570ef --- /dev/null +++ b/tests/test_base_trainer.py @@ -0,0 +1,26 @@ +from src.training.base_trainer import BaseTrainer + + +class DummyTrainer(BaseTrainer): + + def train_one_epoch(self): + pass + + def validate(self): + pass + + def train(self, epochs): + pass + + +def main(): + + print("BaseTrainer imported successfully") + + print(BaseTrainer) + + print(DummyTrainer) + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/tests/test_callbacks.py b/tests/test_callbacks.py new file mode 100644 index 0000000..324d529 --- /dev/null +++ b/tests/test_callbacks.py @@ -0,0 +1,43 @@ +from src.training.callbacks import Callback +from src.training.callbacks import CallbackManager + + +class PrintCallback(Callback): + + def on_train_begin(self, trainer): + print("Training Started") + + def on_epoch_begin(self, trainer): + print("Epoch Started") + + def on_epoch_end(self, trainer): + print("Epoch Finished") + + def on_validation_end(self, trainer): + print("Validation Finished") + + def on_train_end(self, trainer): + print("Training Finished") + + +def main(): + + manager = CallbackManager() + + manager.add(PrintCallback()) + + trainer = object() + + manager.on_train_begin(trainer) + + manager.on_epoch_begin(trainer) + + manager.on_epoch_end(trainer) + + manager.on_validation_end(trainer) + + manager.on_train_end(trainer) + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/tests/test_history.py b/tests/test_history.py new file mode 100644 index 0000000..fe6f884 --- /dev/null +++ b/tests/test_history.py @@ -0,0 +1,48 @@ +from src.training.history import History + + +def main(): + + history = History() + + history.add( + epoch=1, + train_loss=0.45, + val_loss=0.41, + train_accuracy=92.5, + val_accuracy=91.8, + ) + + history.add( + epoch=2, + train_loss=0.32, + val_loss=0.28, + train_accuracy=95.2, + val_accuracy=94.7, + ) + + print("History Length") + + print(len(history)) + + print() + + print("Best Epoch") + + print(history.best()) + + history.save_json( + "outputs/history.json" + ) + + history.save_csv( + "outputs/history.csv" + ) + + print() + + print("History Saved") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/tests/test_metrics.py b/tests/test_metrics.py new file mode 100644 index 0000000..581f211 --- /dev/null +++ b/tests/test_metrics.py @@ -0,0 +1,40 @@ +import torch + +from src.training.metrics import Metrics + + +def main(): + + metrics = Metrics() + + outputs = torch.tensor( + [ + [0.1, 0.9], + [0.8, 0.2], + [0.2, 0.8], + [0.7, 0.3], + ] + ) + + labels = torch.tensor( + [ + 1, + 0, + 1, + 0, + ] + ) + + metrics.update( + outputs, + labels, + ) + + results = metrics.compute() + + metrics.print(results) + + +if __name__ == "__main__": + + main() \ No newline at end of file diff --git a/tests/test_state.py b/tests/test_state.py new file mode 100644 index 0000000..5bed9de --- /dev/null +++ b/tests/test_state.py @@ -0,0 +1,39 @@ +from src.training.state import TrainingState + + +def main(): + + state = TrainingState() + + state.update( + + epoch=5, + + global_step=340, + + best_val_loss=0.248, + + best_val_accuracy=94.75, + + learning_rate=1e-4, + + resumed=True, + + ) + + print() + + print("Training State") + + print("=" * 50) + + for key, value in state.to_dict().items(): + + print(f"{key:25}: {value}") + + print("=" * 50) + + +if __name__ == "__main__": + + main() \ No newline at end of file From 02d976af964237ceb2b5de462d6f4eaaccc91b47 Mon Sep 17 00:00:00 2001 From: Arvindhbabu Date: Mon, 20 Jul 2026 10:13:32 +0530 Subject: [PATCH 2/4] refactor(trainer): migrate to TrainingState --- src/training/trainer.py | 42 +++++++++++++++++++---------------------- 1 file changed, 19 insertions(+), 23 deletions(-) diff --git a/src/training/trainer.py b/src/training/trainer.py index e52e6b6..298eb28 100644 --- a/src/training/trainer.py +++ b/src/training/trainer.py @@ -4,11 +4,14 @@ Generic Trainer """ +from xml.parsers.expat import model + import torch from tqdm import tqdm +from src.training.base_trainer import BaseTrainer -class Trainer: +class Trainer(BaseTrainer): def __init__( self, @@ -24,25 +27,18 @@ def __init__( 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 + super().__init__( + model=model, + optimizer=optimizer, + criterion=criterion, + train_loader=train_loader, + val_loader=val_loader, + device=device, + logger=logger, + scheduler=scheduler, + checkpoint_manager=checkpoint_manager, + early_stopping=early_stopping, + ) def train_one_epoch(self): @@ -205,7 +201,7 @@ def train(self, epochs): f"New Best Model | Val Loss = {val_loss:.4f}" ) - self.checkpoint.save_best_model( + self.checkpoint_manager.save_best_model( model=self.model, optimizer=self.optimizer, scheduler=self.scheduler, @@ -213,7 +209,7 @@ def train(self, epochs): best_loss=best_loss, ) - self.checkpoint.save_last_model( + self.checkpoint_manager.save_last_model( model=self.model, optimizer=self.optimizer, scheduler=self.scheduler, @@ -229,6 +225,6 @@ def train(self, epochs): break - self.checkpoint.save_metrics(history) + self.checkpoint_manager.save_metrics(history) return history \ No newline at end of file From fa3e5c4ea8d2acf981cdbbafc36a2f52be7a0ab2 Mon Sep 17 00:00:00 2001 From: Arvindhbabu Date: Mon, 20 Jul 2026 10:39:07 +0530 Subject: [PATCH 3/4] refactor(trainer): integrate history manager --- src/training/trainer.py | 45 +++++++++++++++++------------------------ tests/test_history.py | 13 ++++++------ 2 files changed, 25 insertions(+), 33 deletions(-) diff --git a/src/training/trainer.py b/src/training/trainer.py index 298eb28..b3de7fb 100644 --- a/src/training/trainer.py +++ b/src/training/trainer.py @@ -4,13 +4,12 @@ Generic Trainer """ -from xml.parsers.expat import model - import torch from tqdm import tqdm from src.training.base_trainer import BaseTrainer + class Trainer(BaseTrainer): def __init__( @@ -45,9 +44,7 @@ def train_one_epoch(self): self.model.train() running_loss = 0.0 - correct = 0 - total = 0 progress = tqdm( @@ -91,7 +88,6 @@ def train_one_epoch(self): ) epoch_loss = running_loss / len(self.train_loader) - epoch_acc = 100.0 * correct / total return epoch_loss, epoch_acc @@ -101,9 +97,7 @@ def validate(self): self.model.eval() running_loss = 0.0 - correct = 0 - total = 0 with torch.no_grad(): @@ -143,20 +137,17 @@ def validate(self): ) epoch_loss = running_loss / len(self.val_loader) - epoch_acc = 100.0 * correct / total return epoch_loss, epoch_acc def train(self, epochs): - best_loss = float("inf") - - history = [] + self.state.best_val_loss = float("inf") self.logger.info("Training Started") - for epoch in range(epochs): + for epoch in range(self.state.epoch, epochs): self.logger.info( f"Epoch {epoch + 1}/{epochs}" @@ -183,19 +174,17 @@ def train(self, epochs): 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, - } + self.history.add( + epoch=epoch + 1, + train_loss=train_loss, + train_acc=train_acc, + val_loss=val_loss, + val_acc=val_acc, ) - if val_loss < best_loss: + if val_loss < self.state.best_val_loss: - best_loss = val_loss + self.state.best_val_loss = val_loss self.logger.info( f"New Best Model | Val Loss = {val_loss:.4f}" @@ -206,7 +195,7 @@ def train(self, epochs): optimizer=self.optimizer, scheduler=self.scheduler, epoch=epoch + 1, - best_loss=best_loss, + best_loss=self.state.best_val_loss, ) self.checkpoint_manager.save_last_model( @@ -214,9 +203,11 @@ def train(self, epochs): optimizer=self.optimizer, scheduler=self.scheduler, epoch=epoch + 1, - best_loss=best_loss, + best_loss=self.state.best_val_loss, ) + self.state.epoch = epoch + 1 + if self.early_stopping(val_loss): self.logger.info( @@ -225,6 +216,8 @@ def train(self, epochs): break - self.checkpoint_manager.save_metrics(history) + self.checkpoint_manager.save_metrics( + self.history.get() + ) - return history \ No newline at end of file + return self.history.get() \ No newline at end of file diff --git a/tests/test_history.py b/tests/test_history.py index fe6f884..346dfaa 100644 --- a/tests/test_history.py +++ b/tests/test_history.py @@ -2,23 +2,22 @@ def main(): - history = History() history.add( epoch=1, - train_loss=0.45, - val_loss=0.41, - train_accuracy=92.5, - val_accuracy=91.8, + train_loss=0.40, + val_loss=0.35, + train_acc=92.5, + val_acc=91.8, ) history.add( epoch=2, train_loss=0.32, val_loss=0.28, - train_accuracy=95.2, - val_accuracy=94.7, + train_acc=95.2, + val_acc=94.7, ) print("History Length") From b655bc27d9a1d94d5c24a39828751910500f858d Mon Sep 17 00:00:00 2001 From: Arvindhbabu Date: Mon, 20 Jul 2026 11:15:39 +0530 Subject: [PATCH 4/4] refactor(trainer): integrate metrics engine --- src/training/base_trainer.py | 5 ++--- src/training/metrics.py | 8 ++++++++ src/training/trainer.py | 36 ++++++++++++++++++++---------------- 3 files changed, 30 insertions(+), 19 deletions(-) diff --git a/src/training/base_trainer.py b/src/training/base_trainer.py index cfffbbb..2a25abf 100644 --- a/src/training/base_trainer.py +++ b/src/training/base_trainer.py @@ -5,7 +5,7 @@ Defines the common functionality shared by all trainers. """ - +from src.training.state import TrainingState from abc import ABC, abstractmethod import time @@ -51,8 +51,7 @@ def __init__( self.early_stopping = early_stopping # Training state - self.current_epoch = 0 - self.best_val_loss = float("inf") + self.state = TrainingState() # Utilities self.history = History() diff --git a/src/training/metrics.py b/src/training/metrics.py index d2557c9..fbebe56 100644 --- a/src/training/metrics.py +++ b/src/training/metrics.py @@ -78,6 +78,14 @@ def update( def compute(self) -> Dict[str, float]: + if len(self.targets) == 0: + return { + "accuracy": 0.0, + "precision": 0.0, + "recall": 0.0, + "f1": 0.0, + } + y_true = np.array(self.targets) y_pred = np.array(self.predictions) diff --git a/src/training/trainer.py b/src/training/trainer.py index b3de7fb..7329099 100644 --- a/src/training/trainer.py +++ b/src/training/trainer.py @@ -44,8 +44,8 @@ def train_one_epoch(self): self.model.train() running_loss = 0.0 - correct = 0 - total = 0 + + self.metrics.reset() progress = tqdm( self.train_loader, @@ -77,18 +77,20 @@ def train_one_epoch(self): running_loss += loss.item() - _, predicted = outputs.max(1) - - total += labels.size(0) - - correct += predicted.eq(labels).sum().item() + self.metrics.update( + outputs, + labels, + ) progress.set_postfix( loss=f"{loss.item():.4f}" ) epoch_loss = running_loss / len(self.train_loader) - epoch_acc = 100.0 * correct / total + + results = self.metrics.compute() + + epoch_acc = results["accuracy"] * 100 return epoch_loss, epoch_acc @@ -97,8 +99,8 @@ def validate(self): self.model.eval() running_loss = 0.0 - correct = 0 - total = 0 + + self.metrics.reset() with torch.no_grad(): @@ -126,18 +128,20 @@ def validate(self): running_loss += loss.item() - _, predicted = outputs.max(1) - - total += labels.size(0) - - correct += predicted.eq(labels).sum().item() + self.metrics.update( + outputs, + labels, + ) progress.set_postfix( loss=f"{loss.item():.4f}" ) epoch_loss = running_loss / len(self.val_loader) - epoch_acc = 100.0 * correct / total + + results = self.metrics.compute() + + epoch_acc = results["accuracy"] * 100 return epoch_loss, epoch_acc