This commit is contained in:
harisreedhar
2025-02-20 18:09:37 +05:30
committed by henryruhs
parent dcf19634d1
commit b47c6b72ee
10 changed files with 112 additions and 170 deletions
+34 -28
View File
@@ -3,15 +3,15 @@ import os
from typing import Any, Tuple
import lightning
import numpy
import torch
from lightning import Trainer
from lightning.pytorch.callbacks import ModelCheckpoint
from lightning.pytorch.loggers import TensorBoardLogger
from lightning.pytorch.tuner import Tuner
from torch import Tensor, nn
from torch.utils.data import DataLoader, Dataset, TensorDataset, random_split
from torch.utils.data import DataLoader, Dataset, random_split
from .data_loader import DataLoaderRecognition
from .models.embedding_converter import EmbeddingConverter
from .types import Batch, Embedding
@@ -22,23 +22,34 @@ CONFIG.read('config.ini')
class EmbeddingConverterTrainer(lightning.LightningModule):
def __init__(self) -> None:
super(EmbeddingConverterTrainer, self).__init__()
source_path = CONFIG.get('training.model', 'source_path')
target_path = CONFIG.get('training.model', 'target_path')
self.lr = CONFIG.getfloat('training.trainer', 'learning_rate')
self.embedding_converter = EmbeddingConverter()
self.source_embedder = torch.jit.load(source_path) # type:ignore[no-untyped-call]
self.target_embedder = torch.jit.load(target_path) # type:ignore[no-untyped-call]
self.source_embedder.eval()
self.target_embedder.eval()
self.mse_loss = nn.MSELoss()
def forward(self, source_embedding : Embedding) -> Embedding:
return self.embedding_converter(source_embedding)
def training_step(self, batch : Batch, batch_index : int) -> Tensor:
source_tensor, target = batch
output_tensor = self(source_tensor)
loss_training = self.mse_loss(output_tensor, target)
with torch.no_grad():
source_embedding = self.source_embedder(batch)
target_embedding = self.target_embedder(batch)
output_embedding = self(source_embedding)
loss_training = self.mse_loss(output_embedding, target_embedding)
self.log('loss_training', loss_training, prog_bar = True)
return loss_training
def validation_step(self, batch : Batch, batch_index : int) -> Tensor:
source_tensor, target_tensor = batch
output_tensor = self(source_tensor)
validation = self.mse_loss(output_tensor, target_tensor)
with torch.no_grad():
source_embedding = self.source_embedder(batch)
target_embedding = self.target_embedder(batch)
output_embedding = self(source_embedding)
validation = self.mse_loss(output_embedding, target_embedding)
self.log('validation', validation, prog_bar = True)
return validation
@@ -53,36 +64,28 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
'lr_scheduler':
{
'scheduler': scheduler,
'monitor': 'train_loss',
'monitor': 'loss_training',
'interval': 'epoch',
'frequency': 1
}
}
def create_loaders() -> Tuple[DataLoader, DataLoader]:
loader_batch_size = CONFIG.getint('training.loader', 'batch_size')
loader_num_workers = CONFIG.getint('training.loader', 'num_workers')
def create_loaders(dataset : Dataset[Any]) -> Tuple[DataLoader[Any], DataLoader[Any]]:
batch_size = CONFIG.getint('training.loader', 'batch_size')
num_workers = CONFIG.getint('training.loader', 'num_workers')
training_dataset, validate_dataset = split_dataset()
training_loader = DataLoader(training_dataset, batch_size = loader_batch_size, num_workers = loader_num_workers, shuffle = True, pin_memory = True)
validation_loader = DataLoader(validate_dataset, batch_size = loader_batch_size, num_workers = loader_num_workers, shuffle = False, pin_memory = True)
training_dataset, validate_dataset = split_dataset(dataset)
training_loader = DataLoader(training_dataset, batch_size = batch_size, shuffle = True, num_workers = num_workers, drop_last = True, pin_memory = True, persistent_workers = True)
validation_loader = DataLoader(validate_dataset, batch_size = batch_size, shuffle = False, num_workers = num_workers, drop_last = False, pin_memory = True, persistent_workers = True)
return training_loader, validation_loader
def split_dataset() -> Tuple[Dataset[Any], Dataset[Any]]:
input_source_path = CONFIG.get('preparing.input', 'source_path')
input_target_path = CONFIG.get('preparing.input', 'target_path')
def split_dataset(dataset : Dataset[Any]) -> Tuple[Dataset[Any], Dataset[Any]]:
loader_split_ratio = CONFIG.getfloat('training.loader', 'split_ratio')
source_tensor = torch.from_numpy(numpy.load(input_source_path)).float()
target_tensor = torch.from_numpy(numpy.load(input_target_path)).float()
dataset = TensorDataset(source_tensor, target_tensor)
dataset_size = len(dataset)
training_size = int(loader_split_ratio * len(dataset))
validation_size = int(dataset_size - training_size)
training_dataset, validate_dataset = random_split(dataset, [ training_size, validation_size ])
training_size = int(loader_split_ratio * len(dataset)) # type:ignore[operator, arg-type]
validation_size = len(dataset) - training_size # type:ignore[arg-type]
training_dataset, validate_dataset = random_split(dataset, [training_size, validation_size])
return training_dataset, validate_dataset
@@ -112,9 +115,12 @@ def create_trainer() -> Trainer:
def train() -> None:
dataset_path = CONFIG.get('preparing.dataset', 'dataset_path')
dataset_image_pattern = CONFIG.get('preparing.dataset', 'image_pattern')
resume_file_path = CONFIG.get('training.output', 'resume_file_path')
training_loader, validation_loader = create_loaders()
dataset = DataLoaderRecognition(dataset_path, dataset_image_pattern)
training_loader, validation_loader = create_loaders(dataset)
embedding_converter_trainer = EmbeddingConverterTrainer()
trainer = create_trainer()
tuner = Tuner(trainer)