fixed types
Code Quality Pipeline / Check Code (pull_request) Successful in 3m15s

This commit is contained in:
brian
2025-01-06 13:42:21 +00:00
parent 16ad2b80ee
commit aff0ae8fc6
16 changed files with 88 additions and 74 deletions
+5 -4
View File
@@ -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 -6
View File
@@ -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)
+4 -6
View File
@@ -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
View File
@@ -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