From 49ace925245ef47115200ef6c6c6f384530eb663 Mon Sep 17 00:00:00 2001 From: ken4647 Date: Fri, 2 Dec 2022 23:00:40 +0800 Subject: [PATCH] manually merge the branch of evalution --- .idea/.gitignore | 3 - .idea/ReIDMixLearning.iml | 8 -- .idea/inspectionProfiles/Project_Default.xml | 33 ----- .../inspectionProfiles/profiles_settings.xml | 6 - .idea/misc.xml | 7 - .idea/modules.xml | 8 -- evaluate_CMC.py | 121 ++++++++++++++++++ evaluate_mAP.py | 118 +++++++++++++++++ evaluate_rank10.py | 90 +++++++++++++ loaders/__pycache__/img_loader.cpython-39.pyc | Bin 2156 -> 0 bytes main.py | 0 model/__pycache__/judge_loss.cpython-39.pyc | Bin 1250 -> 0 bytes model/__pycache__/my_model.cpython-39.pyc | Bin 4810 -> 0 bytes param.py | 16 +++ train.py | 4 + 15 files changed, 349 insertions(+), 65 deletions(-) delete mode 100644 .idea/.gitignore delete mode 100644 .idea/ReIDMixLearning.iml delete mode 100644 .idea/inspectionProfiles/Project_Default.xml delete mode 100644 .idea/inspectionProfiles/profiles_settings.xml delete mode 100644 .idea/misc.xml delete mode 100644 .idea/modules.xml create mode 100644 evaluate_CMC.py create mode 100644 evaluate_mAP.py create mode 100644 evaluate_rank10.py delete mode 100644 loaders/__pycache__/img_loader.cpython-39.pyc delete mode 100644 main.py delete mode 100644 model/__pycache__/judge_loss.cpython-39.pyc delete mode 100644 model/__pycache__/my_model.cpython-39.pyc create mode 100644 param.py 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 061f2795286136fd64300674dafc6a3a43610432..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 2156 zcmaJ?OK%)S5bo}I?#tdd4<$Sz!2z*=Rtg*vp$J992?8N2B2t7#My<(qdo#=I?7FAd z!PfX>Uy(R)%)vgP{0Dvlzo4(2@)x+kSG|5E2wIxzs_xpVuBxwU^IosbuztIBIEuji zNy5byLb!*o{uM+q#Y>hgcb;)Ci%eh_WiSo1kT36|Eb@18mS7LcblS>VCwJ|vt;M8s z)IDbAi3%&BqOA87Q?W`8nM!cCGFhGU4}z?(Qq?+S*(JTEFRS(op*pI27-m=WRn_}Z zm|qls%myp3@G`bB;I6%%+x**EuCyV^n!KvFW@cKuaI{dl3r6$Bh?Fm`vj9KAsb&-u z0&3|%i5FZ2Dm;v`P(>=nLn4)^6npF)w_QRGYi-H->gd_5)|K2|XtOH|HHt2HX@Mt4 z_wm(#fLPY>1Ku#n!HEODA2k6aQ83Uzu_i=(^CntnKYY0#+wdrA;$wak+hmd+^Zle@ z6J&kC&C`Y*i6+5a%eIeUaS}A~UAFyELpmpX^}VWo=DSjkC^DTB=1}%_n_H-XZz0;;Gia_NT*ogYhq2b!R=i>Pu-dBh+clWh^K7&2BcWfc5H$ce9TM=aUfh!>dJ>X z;QsM%z9W$m8d81+_1ag4@2#Hu^>+iYu`x)f{8c{HQo6R3)0tY7#CIizLtdWTgwz)1 z9T+ndmHChe;q3yU_ze}Fl3f#%^8WR-=g(e?kA@9>K0XCz%=R%_jxZ2VTOb7E0{Hlu zbBr4dm3zPm#shG{n&4OfBP0O<$J~6=a2q$_0ly6!+c$xxgwSMeHXK-i24^^>(Hm|& z#NK|0kI|B1R=b44Q-wi)&Wv@TZ^dvjD=L7Z#)YAo=u%y&`gwto5$E$c zgh6C(;UbF+8OSllNV^ZQk%=(fw$>GZ3i!uzNz~z~6wx3H}7%5U1QYbK-estAqnj@;A@+__NG2ANMJtB8zJ{s_rf zCMzTph3A7ynhqs?f+qY5SkdtWds0<^Luazep4O(cp5=|qJk1wP9)7xMA*o>Y9)LnA zmMA5!&?%N!k`+&>V)LjQv7|pz$z-%*Y4Qp&lyV3g`{9>Gm1&`CYjfOj^$>iz9rgoo z9J)knyg?W_vZfofW-i*`7KvCQV8X|f#7DMSDC4;hMO8Q<41t3_dHm$zSHSOyt9>lh zd66qmORM0X%qo5Q2K#{!xz4N=;tl%eqlceOkG>HnI;*Bfb@@QqVpg4)daiOeEo-Ut zbl%9B>axDOcxgC<8NjezJ@AlR+_}Ei-M86HsE_N$EgI(|XBv0r2YFq|!WDIuzpW&& zbW|V#bQCE60B~?~iD!6&*Tm5cUNXm8dWn~j#LL)4ZFGqrpl@FE% z@LLjaLzcYdl00K(+A{Fy9{Ln5+4Eo9STgY9mO+%H{Eq;c~&n=U1 zYQJYY=vn%e|1hyOdOy4qKQ`fMXRm(<_WR-09?Y(~rSJECdGzPS>TtVvbvT19u)hx{ z94c0u{LIH+tIF2KgkCW_0G?`9`3TfHQ)bjHs27XLz>LEmy8-S52w~EptKC*?K8ABM zcb1NGpo51MdEHdb<09owl?ws-Fzsp)82Y$Sj zYOb^fiC^a!(lQ}Y6T}lV+3AtvQPizWc6!3DDR#Pgo);FxXF{X-sA|f^OX&ZF31al4 q!iwxXD|Gf$tG?Fp<1X2b?MC413i12>bv_RdYQrNz^iUqMOa28gu`e6| diff --git a/model/__pycache__/my_model.cpython-39.pyc b/model/__pycache__/my_model.cpython-39.pyc deleted file mode 100644 index e33b8810fe450b0a417730c740c033f0e57ac6e2..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 4810 zcmb_gOK;rP73L*5BGxf7I5U>&z(B7shv%Mqct{?8=evgr zr=}VP#@CC_d+*E}#=kfjA2ue}@Js&!Aq^>F!{^bAjZNVTQP|8)VH(2Nw0v6_k+D%b za1IT%BCW)fw(ovq$eMJX8`6oyhWE_!>#(@egT*T?4Or^30ZXH_G+}AV7A&pO(t>45 zwqa?PmMJM77~Sb(+?%o56}g!tx!v2AVQ%%dwtF;>538J{KY$V$zL17*er2o)DWv(_ z^DSvf`?=xUUm3RH*ATjvPu-8wzT6Ju8}Xp`q-T!Bqwe6xu<1jH%*e#P@yt9h50Q<1 z=z*2l8#QS964H8M!{$^rdkz_uW@&dP?8sPY_oRVaaaWJbV_r(R5XS=I$jqBhB9%mO zkoNbY+}#SL>?iBTyt5-ydGywT-rag$OSiWorD~YvUJ&$?eij6|*Lxf$2$CwEmpeBH z$DSi^DVxf+OSK@BRA9!<(N4_u?>lcyq9M z3yE7#?y12>)XN@j4rCNR+}sVcvaq$QrV*lvU#bx(+QJcY;=I_KpUC*a@HQt>&phxc zZu0ykh|D-Zx)02KafrV!P3RW&FKM;+P1>}N*f-A^hk(tISjFoVxAR6Eg-><^6~)`C zgHURY;tWLBQfG0<#nW8us`E5!ae+p6ZJ@pgm84_JBpYHzH1M}~cJcvAzaxh!Pr}Uz zDQ^VK!*)!)6$GCHU&V^fih7$Dm@0GIi{dy4#x9RKH5Bhbzer{;Wtn=S zais4;oQC@dAeGi%g|ubux$Qd||1Lo1JE z%-1sOz>#_qiE-_e*VzFl`lZMDBE)S5FiWQ z!0b(Je+$}owDz{H*O}F#P+2L8$to(b4sCr|kH1|7Wr?lhOI&QxjvvN9uO45+;rNdf zAQA|u$LBEZUQ%rw)DJ0GK)PhrEHz>`Z{7&A-s9DQ+C&9gkCH{-y`=_QgY7K$WT7l- zZ0Gg6;nRDAK@2}{Gkm&5m1is!^x&#TB)rAyw2}f=D(_`_FXyG>Qe|JEeWf6b_kH(v z9A;SrUS5NeJ9qj?6sp{b_0e}9^w&2BeVMx;3o3s74$JX=bm!B&7PBaigswuRVPDms ze)@}eGSR+?97^-3F^+o){KN^4cxRHd~mZMxECDs8sX zI;92%R7_v@OhJ0rp{t)z{FH*Mrbw)6nPP?FXB6*K{G8%4#V;tXQ2dhO1Bz=DzoNJb zfqjT1syi*1;VLj>Nh+~=XtVxKY_j4 z{pV4p^J4GfiA+_cI#HucRn-9yc2$j!pw-G=Kp-_$FHls1TU5du z2$tKr9hiJl^(I!zoykzU8wuhK<36lh0PdS6zbpSeOwg`6?bT z7BPY?-c_{xUqA^i)S$dI_#!zSliaK{CjSKGRhYbBaIbvvYFt`lUK7{^PO;+y!Nj;l zjKdRO4lgQ<>%F$AFM2l_u&3N%xzk((RQa<3Snof3!`E8n@pQxztrcSV5=V%IRxrMmnZPYOR!d_S zp9?^&4HQ#z3JpFkirp^y*aEwrT1CVWfYom~Jq^FL5gq%ke(2z-`W?J;3uqq0U!$W_F@vx6XqO)N>>cy)FdKTH zrPVP!dqL-RZ-C?RO7v7=xF~t-_ zSNP3>$MjdEdY23Dap7!HQ}<|ibxke8#dpg#Qr+bVL^#^Ka5qwm`aPGErSiC39#{O? z>oVNR`a99}o%Nz4_1$o1ogJzAkWtLFp1Q%9rxfNM>>4P#1w>J{(E2sLajeb>**5n& z6--XAQFeYL+MA4e9|BePx}Bdu#fl9_3X<5Nsr$oWx5s~1SgVTMb^Hwm+7^<-q>k1`e)eF8Hl1rb>C!m*zpcqkHyR1UZ8x_yNpzB zHuH*>Zd={LK{KjTm3W$cfx)gK8mB5Y 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(), ])