90 lines
2.9 KiB
Python
90 lines
2.9 KiB
Python
import shutil
|
|
import matplotlib.pyplot as plt
|
|
import torch
|
|
from torchvision import transforms
|
|
import numpy as np
|
|
import cv2
|
|
import os
|
|
|
|
from model import my_model
|
|
import param
|
|
|
|
|
|
|
|
def extractidfeature(id_path: str, extractor, p_device=torch.device("cuda:0" if torch.cuda.is_available() else "cpu")):
|
|
img_id = cv2.imread(id_path)
|
|
raw_transformer = transforms.Compose([
|
|
transforms.ToPILImage(),
|
|
transforms.Resize((128,64)),# hxw
|
|
transforms.ToTensor(),
|
|
])
|
|
tensor_id = extractor(raw_transformer(img_id).unsqueeze(dim=0).to(p_device))
|
|
|
|
return tensor_id.detach().cpu().numpy()
|
|
|
|
# example
|
|
if __name__ == '__main__':
|
|
dislist = {}
|
|
extractor_model_path = "./deep_sort/deep/checkpoint/market_bot_R50-ibn.pth"
|
|
cfg_path = "./fastreid/cfgs/Market1501/bagtricks_R50-ibn.yml"
|
|
query_path = "D:\\study\\python\\1zhouxue\\UESTC_ReID_Dataset_V3\\query\\cam6"
|
|
galla_path = "D:\\study\\python\\1zhouxue\\UESTC_ReID_Dataset_V3\\gallery"
|
|
camera_index = 6 # 第几个摄像头
|
|
gallary_list = os.listdir(galla_path)
|
|
count = 0
|
|
|
|
my_option = param.Parameters()
|
|
|
|
extractor = my_model.RestNet18() # load model in head
|
|
extractor.load_state_dict(torch.load(my_option.weights_reid))
|
|
extractor = extractor.to(my_option.device)
|
|
|
|
with torch.no_grad():
|
|
for i in range(1, 7):
|
|
count = 0
|
|
dislist.clear()
|
|
id_try = extractidfeature(query_path+f"/{camera_index}_{i}.jpg", extractor)
|
|
for gal in gallary_list:
|
|
count += 1
|
|
if count % 5 == 0:
|
|
id_tar = extractidfeature(galla_path+'/'+gal, extractor)
|
|
dist = np.sum(np.abs(id_try-id_tar))
|
|
dislist[gal] = float(dist)
|
|
if count > 10000:
|
|
break
|
|
matched = sorted(dislist.items(), key=lambda x : x[1])
|
|
print(matched)
|
|
try:
|
|
os.mkdir("match/" + f"{camera_index}_{i}")
|
|
except:
|
|
pass
|
|
|
|
img = plt.imread(query_path+f"/{camera_index}_{i}.jpg")
|
|
plt.subplot(1, 11, 1)
|
|
plt.title(f"query_{i}")
|
|
plt.imshow(img)
|
|
plt.xticks([])
|
|
plt.yticks([])
|
|
x = 1
|
|
for path in matched[0:10]:
|
|
s_path = galla_path+'/'+path[0]
|
|
t_path = "match/" + f"{camera_index}_{i}" + "/" + path[0]
|
|
print(s_path, t_path)
|
|
shutil.copy(s_path, t_path)
|
|
|
|
ID = path[0].split('_')
|
|
# print(ID)
|
|
ID = int(ID[0])
|
|
img = plt.imread(s_path)
|
|
plt.subplot(1, 11, x+1)
|
|
if ID == i:
|
|
plt.title(x)
|
|
else:
|
|
plt.title(x, color='red')
|
|
plt.imshow(img)
|
|
plt.xticks([])
|
|
plt.yticks([])
|
|
x += 1
|
|
plt.show()
|
|
|
|
|