50 lines
980 B
Python
50 lines
980 B
Python
from __future__ import annotations
|
|
|
|
from datetime import datetime
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
from dataloader import VCDADataset # noqa: F401
|
|
|
|
from .models import VisualCommunicationModel
|
|
|
|
PRE_WARMUP_LR = 1e-10
|
|
POST_WARMUP_LR = 1e-5
|
|
BATCH_SIZE = 32
|
|
MODEL_NAME = 'visual_communication_model_v1'
|
|
MAX_EPOCHS = 200
|
|
RUN_NAME = datetime.now().strftime('%Y-%m-%d-%H%M') + '_' + MODEL_NAME
|
|
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
|
|
|
|
|
# create loss function
|
|
criterion = nn.MSELoss()
|
|
|
|
|
|
def loss_fn(pred, target):
|
|
"""Create loss function."""
|
|
return criterion(pred, target.unsqueeze(-1))
|
|
|
|
|
|
# create model
|
|
model = VisualCommunicationModel().to(DEVICE)
|
|
|
|
# create optimizer
|
|
optim = torch.optim.Adam(model.parameters(), lr=POST_WARMUP_LR)
|
|
|
|
# create datasets
|
|
# train_dataset
|
|
# validation_dataset
|
|
|
|
# create dataloaders
|
|
|
|
# create trainer and evaluator
|
|
|
|
# setup progressbar
|
|
|
|
# setup lr scheduler
|
|
|
|
# setup checkpoint saving
|
|
|
|
# load in checkpoint if exists
|