75 lines
2.4 KiB
Python
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 |