49 lines
1.3 KiB
Python
49 lines
1.3 KiB
Python
"""Main script to be run by service."""
|
|
|
|
import json
|
|
import logging
|
|
from traceback import print_exc
|
|
|
|
import torch
|
|
from models import VisualCommunicationModel
|
|
from tqdm import tqdm
|
|
from utils import DEVICE, VCDADataset, load_model
|
|
|
|
from shared.data_store import connect_minio
|
|
from shared.mongodb.classes import ModelData
|
|
from shared.utils import setup_logging
|
|
|
|
if __name__ == '__main__':
|
|
# setup logging
|
|
setup_logging()
|
|
# connect to minio
|
|
minio_client = connect_minio()
|
|
# instantiate model
|
|
model: VisualCommunicationModel = load_model(client=minio_client)
|
|
model.eval()
|
|
# setup dataset
|
|
dataset = VCDADataset(
|
|
minio_client=minio_client,
|
|
data_name_list=[
|
|
'02dbaf48d713e4e6d3a6b98fd2dc866e',
|
|
],
|
|
do_augment=False,
|
|
)
|
|
|
|
# make prediction
|
|
with torch.no_grad():
|
|
for i in tqdm(range(len(dataset))):
|
|
try:
|
|
# get image
|
|
image = dataset[i]
|
|
image = torch.unsqueeze(image, 0) # add artificial batch dimension
|
|
image = image.to(DEVICE)
|
|
# make prediction
|
|
pred: ModelData = model(image)
|
|
except Exception:
|
|
print_exc()
|
|
continue
|
|
else:
|
|
print(json.dumps(pred.model_dump(), indent=4))
|
|
logging.debug('finished')
|