diff --git a/.idea/.gitignore b/.idea/.gitignore
deleted file mode 100644
index 26d3352..0000000
--- a/.idea/.gitignore
+++ /dev/null
@@ -1,3 +0,0 @@
-# Default ignored files
-/shelf/
-/workspace.xml
diff --git a/.idea/ReIDMixLearning.iml b/.idea/ReIDMixLearning.iml
deleted file mode 100644
index aa55312..0000000
--- a/.idea/ReIDMixLearning.iml
+++ /dev/null
@@ -1,8 +0,0 @@
-
-
-
-
-
-
-
-
\ No newline at end of file
diff --git a/.idea/inspectionProfiles/Project_Default.xml b/.idea/inspectionProfiles/Project_Default.xml
deleted file mode 100644
index aaa7569..0000000
--- a/.idea/inspectionProfiles/Project_Default.xml
+++ /dev/null
@@ -1,33 +0,0 @@
-
-
-
-
-
-
-
-
\ No newline at end of file
diff --git a/.idea/inspectionProfiles/profiles_settings.xml b/.idea/inspectionProfiles/profiles_settings.xml
deleted file mode 100644
index 105ce2d..0000000
--- a/.idea/inspectionProfiles/profiles_settings.xml
+++ /dev/null
@@ -1,6 +0,0 @@
-
-
-
-
-
-
\ No newline at end of file
diff --git a/.idea/misc.xml b/.idea/misc.xml
deleted file mode 100644
index 511c911..0000000
--- a/.idea/misc.xml
+++ /dev/null
@@ -1,7 +0,0 @@
-
-
-
-
-
-
-
\ No newline at end of file
diff --git a/.idea/modules.xml b/.idea/modules.xml
deleted file mode 100644
index b673fd2..0000000
--- a/.idea/modules.xml
+++ /dev/null
@@ -1,8 +0,0 @@
-
-
-
-
-
-
-
-
\ No newline at end of file
diff --git a/evaluate_CMC.py b/evaluate_CMC.py
new file mode 100644
index 0000000..0c5e006
--- /dev/null
+++ b/evaluate_CMC.py
@@ -0,0 +1,121 @@
+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()
+
+if __name__ == '__main__':
+ 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)
+
+ # gallery_index = 3
+ gallery_path = f"./mydataset/gallery/"
+ classes = os.listdir(gallery_path)
+ cmc_sum = []
+ cmc_final = []
+ for query_index in range(1, 6+1): #query是从1开始的
+
+ query_path = f"./mydataset/query/cam{query_index}/" # Query path
+ #Gallary path,这里直接对gallary文件夹下所有图片进行分析
+ print("--------------query cam is ", query_index)
+ for cls in classes:
+ if cls != f"{query_index}":
+ print("——————————gallery cam is", cls)
+ gallery_path = f"./mydataset/gallery/{cls}/"
+ query_feature_list = [] # 计算每个query的特征向量保存在列表中
+ acck_l_list = [] # 二维列表(列表的列表),其元素(也是列表)存储单个query的"acck"("cmc")
+ cmc_k_list = [] # 存储最终的CMC数字(所有query取平均)
+ query_list = os.listdir(query_path) # 为每张图片生成路径并保存在列表变量中
+ query_list.sort()
+ gallery_list = os.listdir(gallery_path) # 为每张图片生成路径并保存在列表变量中
+ with torch.no_grad():
+ # 对于每个query,(这句话加在本循环内每个注释)
+ for path in query_list:
+ distance_dict = {}
+ parsed_query_id = path.split("_")[1]
+ parsed_query_id = parsed_query_id.split(".")[0]
+ parsed_query_id = int(parsed_query_id)
+ if parsed_query_id == 4 and cls == "5":
+ continue
+ print(parsed_query_id)
+ id_feature = extractidfeature(query_path+path, extractor)
+ query_feature_list.append(id_feature)
+ counter = 0
+ # 计算每个gallary图片到目标query的特征距离
+ for p in gallery_list:
+ if counter%10 == 9: # 抽样,相当于设置帧间隔
+ img_path = gallery_path+p
+ gallery_feature = extractidfeature(img_path, extractor)
+ distance_dict[p] = np.sum(np.abs(id_feature-gallery_feature))
+ counter += 1
+ matched = sorted(distance_dict.items(), key=lambda x : x[1]) # 按特征距离排序
+ # 储存k从前1到前10名的是否(存在命中的ID)
+ temp_acck = [int(matched[0][0].split("_")[2].split(".")[0])==parsed_query_id] # ID相同为True,否则False
+ # 获取前k个acck情况
+ for i in range(1,11):
+ parsed_gallery_id = matched[i][0].split("_")[2].split(".")[0]
+ parsed_gallery_id = int(parsed_gallery_id) # 解析出gallary的ID
+ print(parsed_gallery_id) # 解析出gallary的ID
+ temp_acck.append(parsed_gallery_id==parsed_query_id or temp_acck[i-1]) # or用意是表明如果Acc(k-1)已经为True,那么Acck必然为True(递推)
+ print(temp_acck)
+ acck_l_list.append(temp_acck)
+ # 计算不同k下cmc的具体数值
+ for k in range(1, 11):
+ acc_counter = 0.0
+ length = len(acck_l_list)
+ for results in acck_l_list:
+ acc_counter += results[k]
+ cmc_k_list.append(acc_counter/length)
+ #打印以供检查
+ print(cmc_k_list)
+ cmc_sum.append(cmc_k_list)
+ k_index = [k for k in range(1, 11)]
+ #plot figure并保存(不会显示)
+ plt.clf()
+ plt.plot(k_index, cmc_k_list)
+ plt.xlabel("k")
+ plt.ylabel("ACCK")
+ plt.ylim(0.0,1.05)
+ plt.title(f"CMC:query_cam{query_index} and gallery {cls}")
+ for x, y in zip(k_index, cmc_k_list):
+ plt.text(x, y + 0.02, str(round(y, 3)), ha='center', va='bottom', fontsize=10.5)
+ plt.draw()
+ plt.savefig(f"./CMC/query_cam{query_index}_and_gallery_{cls}.jpg")
+ for k in range(10):
+ acc_counter = 0.0
+ length = len(cmc_sum)
+ for results in cmc_sum:
+ acc_counter += results[k]
+ cmc_final.append(acc_counter / length)
+ print(cmc_final)
+ k_index = [k for k in range(1, 11)]
+ # plot figure并保存(不会显示)
+ plt.clf()
+ plt.plot(k_index, cmc_final)
+ plt.xlabel("k")
+ plt.ylabel("ACCK")
+ plt.ylim(0.0, 1.05)
+ plt.title(f"CMC:query and gallery")
+ for x, y in zip(k_index, cmc_final):
+ plt.text(x, y + 0.02, str(round(y,3)), ha='center', va='bottom', fontsize=10.5)
+ plt.draw()
+ plt.savefig(f"./CMC/query_gallery.jpg")
\ No newline at end of file
diff --git a/evaluate_mAP.py b/evaluate_mAP.py
new file mode 100644
index 0000000..875e0d2
--- /dev/null
+++ b/evaluate_mAP.py
@@ -0,0 +1,118 @@
+#####
+# 保存query和前十个配准的gallery图像,并计算mAP(mean average precision)
+######
+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()
+
+
+if __name__ == '__main__':
+ my_option = param.Parameters()
+ dislist = {}
+
+ extractor = my_model.RestNet18() # load model in head
+ extractor.load_state_dict(torch.load(my_option.weights_reid))
+ extractor = extractor.to(my_option.device)
+
+ gallery_index = 4
+ camera_index = 1 # 第几个摄像头
+ query_path = f"./mydataset/query/cam{camera_index}/"
+ galla_path = f"./mydataset/gallery/{gallery_index}/"
+
+ gallary_list = os.listdir(galla_path)
+ count = 0
+ number_query = 12
+
+#####计算mAP使用的变量
+ mAP = 0
+ precision = 0
+ sum_precision = 0
+ AP = 0
+ sum_AP = 0
+ number_true_gallery = 0
+#######
+
+ with torch.no_grad():
+ for i in range(1, 1+number_query):
+ 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"q_{i}_g_{gallery_index}")
+ plt.imshow(img)
+ plt.xticks([])
+ plt.yticks([])
+ ranking = 1
+ ranking_true_gallery = 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[2].split(".")[0])
+ print(i, ID)
+ img = plt.imread(s_path)
+ plt.subplot(1, 11, ranking+1)
+ if ID == i:
+ plt.title(ranking, color='black')
+ precision = ranking_true_gallery/ranking
+ sum_precision += precision
+ ranking_true_gallery += 1
+
+ else:
+ plt.title(ranking, color='red')
+ plt.imshow(img)
+ plt.xticks([])
+ plt.yticks([])
+ ranking += 1
+ # plt.show()
+ plt.draw()
+ plt.savefig(f"query{i}_cam{camera_index}_and_gallery_{gallery_index}.jpg")
+ number_true_gallery = ranking_true_gallery-1
+ if 0 != number_true_gallery:
+ AP = sum_precision / number_true_gallery
+ else:
+ AP = 0
+ sum_precision = 0
+ sum_AP += AP
+ print("query_", i, "AP==", AP)
+ mAP = sum_AP/number_query
+ print("mAP == ", mAP)
+
diff --git a/evaluate_rank10.py b/evaluate_rank10.py
new file mode 100644
index 0000000..4f050c6
--- /dev/null
+++ b/evaluate_rank10.py
@@ -0,0 +1,90 @@
+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()
+
+
\ No newline at end of file
diff --git a/loaders/__pycache__/img_loader.cpython-39.pyc b/loaders/__pycache__/img_loader.cpython-39.pyc
deleted file mode 100644
index 061f279..0000000
Binary files a/loaders/__pycache__/img_loader.cpython-39.pyc and /dev/null differ
diff --git a/main.py b/main.py
deleted file mode 100644
index e69de29..0000000
diff --git a/model/__pycache__/judge_loss.cpython-39.pyc b/model/__pycache__/judge_loss.cpython-39.pyc
deleted file mode 100644
index c34e028..0000000
Binary files a/model/__pycache__/judge_loss.cpython-39.pyc and /dev/null differ
diff --git a/model/__pycache__/my_model.cpython-39.pyc b/model/__pycache__/my_model.cpython-39.pyc
deleted file mode 100644
index e33b881..0000000
Binary files a/model/__pycache__/my_model.cpython-39.pyc and /dev/null differ
diff --git a/param.py b/param.py
new file mode 100644
index 0000000..a07fa6d
--- /dev/null
+++ b/param.py
@@ -0,0 +1,16 @@
+import torch
+
+class Parameters(object):
+ def __init__(self) -> None:
+ self.weights_yolo = "./best.pt"
+ self.weights_reid = "./resnet18.pth"
+ self.conf_thres = 0.35
+ self.iou_thres = 0.70
+ self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
+ self.nosave = False
+ self.classes = [0]
+ self.agnostic_nms = None
+ self.augment = None
+ self.query_index = 1
+ self.gallary_index= 2
+ pass
\ No newline at end of file
diff --git a/train.py b/train.py
index e44142c..064d8dc 100644
--- a/train.py
+++ b/train.py
@@ -46,6 +46,10 @@ if __name__ == '__main__':
raw_transformer = transforms.Compose([
transforms.ToPILImage(),
transforms.Resize((128,64)),# hxw
+ transforms.RandomResizedCrop((100,50)),
+ transforms.RandomAutocontrast(),
+ transforms.RandomHorizontalFlip(),
+ transforms.Resize((128,64)),# hxw
transforms.ToTensor(),
])