This commit is contained in:
+5
-4
@@ -9,7 +9,7 @@ from models import VisualCommunicationModel
|
||||
from tqdm import tqdm
|
||||
from utils import DEVICE, VCDADataset, load_model
|
||||
|
||||
from shared.datastore import connect_minio
|
||||
from shared.datastore import Datastore
|
||||
from shared.mongodb.classes import ModelData
|
||||
from shared.utils import setup_logging
|
||||
|
||||
@@ -17,13 +17,14 @@ if __name__ == '__main__':
|
||||
# setup logging
|
||||
setup_logging()
|
||||
# connect to minio
|
||||
minio_client = connect_minio()
|
||||
datastore = Datastore()
|
||||
datastore.connect()
|
||||
# instantiate model
|
||||
model: VisualCommunicationModel = load_model(client=minio_client)
|
||||
model: VisualCommunicationModel = load_model(client=datastore._client)
|
||||
model.eval()
|
||||
# setup dataset
|
||||
dataset = VCDADataset(
|
||||
minio_client=minio_client,
|
||||
minio_client=datastore._client,
|
||||
data_name_list=[
|
||||
'02dbaf48d713e4e6d3a6b98fd2dc866e',
|
||||
],
|
||||
|
||||
@@ -4,10 +4,9 @@ import logging
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from minio import Minio
|
||||
|
||||
from model.src.models import VisualCommunicationModel
|
||||
from shared.datastore import get_model
|
||||
from shared.datastore import Datastore
|
||||
|
||||
from .get_model_name import get_model_name
|
||||
|
||||
@@ -15,11 +14,11 @@ DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||
|
||||
|
||||
def load_model(
|
||||
client: Minio,
|
||||
datastore: Datastore,
|
||||
) -> VisualCommunicationModel:
|
||||
"""Instantiate model with weights loaded from latest model saved in
|
||||
MinIO."""
|
||||
assert isinstance(client, Minio)
|
||||
assert isinstance(datastore, Datastore)
|
||||
# instantiate model
|
||||
model = VisualCommunicationModel()
|
||||
# get model object name
|
||||
@@ -27,8 +26,7 @@ def load_model(
|
||||
model_object_name = get_model_name(path=model_name_path)
|
||||
logging.info('using model: %s', model_object_name)
|
||||
# load model from minio
|
||||
model_checkpoint = get_model(
|
||||
client=client,
|
||||
model_checkpoint = datastore.get_model(
|
||||
object_name=model_object_name,
|
||||
)
|
||||
model.load_state_dict(model_checkpoint)
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import random
|
||||
|
||||
from minio import Minio
|
||||
from PIL import Image
|
||||
from torch import Tensor
|
||||
from torch.utils.data import Dataset
|
||||
@@ -16,7 +15,7 @@ from torchvision.transforms.functional import (
|
||||
to_tensor,
|
||||
)
|
||||
|
||||
from shared.datastore import get_image
|
||||
from shared.datastore import Datastore
|
||||
|
||||
# resnet18 original normalization values
|
||||
RESNET_NORMALIZE_MEAN = [0.485, 0.456, 0.406]
|
||||
@@ -28,13 +27,13 @@ class VCDADataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
minio_client: Minio,
|
||||
datastore: Datastore,
|
||||
data_name_list: list[str],
|
||||
do_augment: bool = False,
|
||||
random_annotations: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.minio_client = minio_client
|
||||
self.datastore = datastore
|
||||
self.data_name_list = data_name_list
|
||||
self.do_augment = do_augment
|
||||
self.random_annotations = random_annotations
|
||||
@@ -57,8 +56,7 @@ class VCDADataset(Dataset):
|
||||
def __getitem__(self, idx):
|
||||
# get image from database
|
||||
object_name = self.data_name_list[idx]
|
||||
image = get_image(
|
||||
client=self.minio_client,
|
||||
image = self.datastore.get_image(
|
||||
object_name=object_name,
|
||||
)
|
||||
tensor = self.image_to_tensor(image)
|
||||
|
||||
+5
-4
@@ -18,7 +18,7 @@ from torch.optim.lr_scheduler import ExponentialLR
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from model.src.utils import VCDADataset, get_class
|
||||
from shared.datastore import connect_minio
|
||||
from shared.datastore import Datastore
|
||||
|
||||
|
||||
def parse_arguments():
|
||||
@@ -57,18 +57,19 @@ optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
|
||||
loss_fn = nn.CrossEntropyLoss()
|
||||
|
||||
# create datasets and loaders
|
||||
minio_client = connect_minio()
|
||||
datastore = Datastore()
|
||||
datastore.connect()
|
||||
with open('model/src/dataset/train.csv', encoding='utf-8') as fh:
|
||||
train_data_name_list = fh.read().split('\n')
|
||||
train_dataset = VCDADataset(
|
||||
minio_client=minio_client,
|
||||
datastore=datastore,
|
||||
data_name_list=train_data_name_list,
|
||||
)
|
||||
train_loader = DataLoader(dataset=train_dataset, num_workers=args.loader_workers)
|
||||
|
||||
with open('model/src/dataset/val.csv', encoding='utf-8') as fh:
|
||||
val_data_name_list = fh.read().split('\n')
|
||||
val_dataset = VCDADataset(minio_client=minio_client, data_name_list=val_data_name_list)
|
||||
val_dataset = VCDADataset(datastore=datastore, data_name_list=val_data_name_list)
|
||||
val_loader = DataLoader(dataset=val_dataset, num_workers=args.loader_workers)
|
||||
|
||||
# create trainer and evaluator
|
||||
|
||||
Reference in New Issue
Block a user