Use lightning over pytorch_lightning import, Configure tensorboard

This commit is contained in:
henryruhs
2025-03-11 14:43:09 +01:00
parent cd4b10c832
commit d9e10a9f7c
5 changed files with 20 additions and 12 deletions
+9 -5
View File
@@ -2,13 +2,15 @@ import configparser
import os
from typing import Tuple
import pytorch_lightning
import lightning
import torch
import torchvision
from pytorch_lightning import Trainer
from pytorch_lightning.callbacks import ModelCheckpoint
from pytorch_lightning.utilities.types import Optimizer
from lightning import Trainer
from lightning.pytorch.callbacks import ModelCheckpoint
from lightning.pytorch.loggers import TensorBoardLogger
from torch import Tensor
from torch.optim import Optimizer
from torch.utils.data import DataLoader
from .data_loader import DataLoaderVGG
@@ -22,7 +24,7 @@ CONFIG = configparser.ConfigParser()
CONFIG.read('config.ini')
class FaceSwapperTrain(pytorch_lightning.LightningModule, FaceSwapperLoss):
class FaceSwapperTrain(lightning.LightningModule, FaceSwapperLoss):
def __init__(self) -> None:
super().__init__()
FaceSwapperLoss.__init__(self)
@@ -86,9 +88,11 @@ def create_trainer() -> Trainer:
output_directory_path = CONFIG.get('training.output', 'directory_path')
output_file_pattern = CONFIG.get('training.output', 'file_pattern')
trainer_precision = CONFIG.get('training.trainer', 'precision')
logger = TensorBoardLogger('.logs', name = 'face_swapper')
os.makedirs(output_directory_path, exist_ok = True)
return Trainer(
logger = logger,
max_epochs = trainer_max_epochs,
precision = trainer_precision,
callbacks =