ComDesignProject/loaders/img_loader.py

75 lines
2.4 KiB
Python

import json
import os
import torch
from torch.utils.data import DataLoader
import cv2
import numpy as np
from torchvision import transforms
import random
device = "cuda" if torch.cuda.is_available() else "cpu"
class Dataset(torch.utils.data.Dataset):
def __init__(self,gallarys_path,query_path,raw_data_transform = None, chosen_query_id:list = None):
self.images_path_list = getFilepath(gallarys_path)
q_paths = getFilepath(query_path)
self.raw_data_transformer = raw_data_transform
# query dict get
self.query_dict = {}
if None == chosen_query_id:
for q_path in q_paths:
q_img = cv2.imread(q_path).astype(np.uint8)
if self.raw_data_transformer is not None:
self.query_dict[q_path] = self.raw_data_transformer(q_img).to(device)
else:
self.query_dict[q_path] = q_img
else:
for q_path in q_paths:
q_index = parseID(q_path)
if q_index in chosen_query_id:
q_img = cv2.imread(q_path).astype(np.uint8)
if self.raw_data_transformer is not None:
self.query_dict[q_path] = self.raw_data_transformer(q_img).to(device)
else:
self.query_dict[q_path] = q_img
def __getitem__(self,index):
image_path: str = self.images_path_list[index]
label_string: str = image_path.split("_")[-1]
label_index = int(label_string.split(".")[0])
image = cv2.imread(image_path).astype(np.uint8)
if self.raw_data_transformer is not None:
image = self.raw_data_transformer(image)
return image, label_index
def __len__(self):
return len(self.images_path_list)
def getFilepath(path):
rlist = []
file_or_dir = os.listdir(path)
file_or_dir.sort(reverse=False)
for file_dir in file_or_dir:
file_or_dir_path = os.path.join(path,file_dir)
if os.path.isdir(file_or_dir_path):
rlist += getFilepath(file_or_dir_path)
else:
if file_dir.endswith(".jpg"):
rlist.append(file_or_dir_path)
return rlist
def parseID(filename:str)->int:
return int(filename.split("_")[-1].split(".")[0])
if __name__ == '__main__':
pass