adapt CMC calc for self-designed extractor

This commit is contained in:
ken4647 2022-11-24 22:17:48 +08:00
parent c405dc4aae
commit 52f6504844
3 changed files with 129 additions and 39 deletions

View File

@ -6,69 +6,72 @@ import numpy as np
import cv2
import os
import param
from deep_sort.deep.feature_extractor import FastReIDExtractor
from fastreid.demo import FeatureExtractionDemo
from fastreid.config import get_cfg
from fastreid.my_extractor import MyExtractor
def extractidfeature(id_path: str, extractor: FeatureExtractionDemo):
def extractidfeature(id_path: str, extractor: MyExtractor):
img_id = cv2.imread(id_path)
tensor_id = extractor.run_on_image(img_id)
tensor_id = extractor([img_id])
return tensor_id.detach().cpu().numpy()
if __name__ == '__main__':
extractor_model_path = "./deep_sort/deep/checkpoint/market_bot_R50-ibn.pth"
cfg_path = "./fastreid/cfgs/Market1501/bagtricks_R50-ibn.yml"
my_option = param.Parameters() # load model in head
extractor = MyExtractor(my_option.weights_reid,use_cuda=True)
cfg = get_cfg()
cfg.MODEL.DEVICE = 'cuda' if torch.cuda. is_available () else 'cpu'
cfg.MODEL.WEIGHTS = extractor_model_path
cfg.merge_from_file(cfg_path)
cfg.freeze()
my_option = param.Parameters() # load model in head
extractor = FeatureExtractionDemo(cfg)
for query_index in range(1,3+1):
query_path = f"./query/cam{query_index}/"
gallery_path = "./gallery/"
for query_index in range(1,3+1): #query是从1开始的
query_path = f"./query/cam{query_index}/" # Query path
gallery_path = "./gallery/" #Gallary path这里直接对gallary文件夹下所有图片进行分析
query_feature_list = []
acck_l_list = []
cmc_k_list = []
query_feature_list = [] # 计算每个query的特征向量保存在列表中
acck_l_list = [] # 二维列表列表的列表其元素也是列表存储单个query的"acck"("cmc")
cmc_k_list = [] # 存储最终的CMC数字所有query取平均
query_list = os.listdir(query_path)
query_list = os.listdir(query_path) # 为每张图片生成路径并保存在列表变量中
query_list.sort()
gallery_list = os.listdir(gallery_path)
gallery_list = os.listdir(gallery_path) # 为每张图片生成路径并保存在列表变量中
with torch.no_grad():
with torch.no_grad():
# 对于每个query这句话加在本循环内每个注释
for path in query_list:
distance_dict = {}
distance_dict = {}
parsed_query_id = path.split("_")[1]
id_feature = extractidfeature(query_path+path, extractor)
query_feature_list.append(id_feature)
query_feature_list.append(id_feature)
counter = 0
# 计算每个gallary图片到目标query的特征距离
for p in gallery_list:
if counter%10 == 9:
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])
temp_acck = [matched[0][0].split("_")[2]==parsed_query_id] # 储存k从1到10的是否
for i in range(1,11):
parsed_gallery_id = matched[i][0].split("_")[2]
temp_acck.append(parsed_gallery_id==parsed_query_id or temp_acck[i-1])
matched = sorted(distance_dict.items(), key=lambda x : x[1]) # 按特征距离排序
# 储存k从前1到前10名的是否存在命中的ID
temp_acck = [matched[0][0].split("_")[2]==parsed_query_id] # ID相同为True否则False
# 获取前k个acck情况
for i in range(1,10+1):
parsed_gallery_id = matched[i][0].split("_")[2] # 解析出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)
for k in range(11):
# 计算不同k下cmc的具体数值
for k in range(10+1):
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)
k_index = [k for k in range(11)]
#plot figure并保存不会显示
plt.clf()
plt.plot(k_index, cmc_k_list)
plt.xlabel("k")

84
fastreid/my_extractor.py Normal file
View File

@ -0,0 +1,84 @@
import torch
from torch import nn
import torchvision.transforms as transforms
import logging
# Define model
# input:128x256x3
# output: 64
class MyNet(nn.Module):
def __init__(self):
super(MyNet, self).__init__()
self.conv1 = nn.Conv2d(3,8,5,1,2)
self.maxpool1 = nn.MaxPool2d(2) # 64x128x8
self.bn1 = nn.BatchNorm2d(8)
self.conv2 = nn.Conv2d(8,16,5,1,2)
self.maxpool2 = nn.MaxPool2d(2) # 32x64x16
self.bn2 = nn.BatchNorm2d(16)
self.conv3 = nn.Conv2d(16,32,5,1,2)
self.conv4 = nn.Conv2d(32,64,5,1,2)
self.maxpool3 = nn.MaxPool2d(2) # 16x32x64
self.bn3 = nn.BatchNorm2d(64)
self.conv5 = nn.Conv2d(64,32,5,1,2) # 16x32x32
self.maxpool4 = nn.MaxPool2d(2) # 8x16x32
self.bn4 = nn.BatchNorm2d(32)
self.conv6 = nn.Conv2d(32,64,8,8) #1x2x64
self.flat = nn.Flatten()
self.linear = nn.Linear(2*64,64) # 64D-feature ID
self.sigmoid = nn.Sigmoid()
self.lrelu = nn.LeakyReLU()
def forward(self, x:torch.Tensor):
x = self.conv1(x)
x = self.lrelu(x)
x = self.maxpool1(x)
x = self.bn1(x)
x = self.conv2(x)
x = self.lrelu(x)
x = self.maxpool2(x)
x = self.bn2(x)
x = self.conv3(x)
x = self.lrelu(x)
x = self.conv4(x)
x = self.lrelu(x)
x = self.maxpool3(x)
x = self.bn3(x)
x = self.conv5(x)
x = self.lrelu(x)
x = self.maxpool4(x)
x = self.bn4(x)
x = self.conv6(x)
x = self.sigmoid(x)
x = self.flat(x)
x = self.linear(x)
return x
class MyExtractor(object):
def __init__(self, model_path, use_cuda=True):
self.device = "cuda" if torch.cuda.is_available() and use_cuda else "cpu"
self.net = MyNet()
self.net.load_state_dict(torch.load(model_path))
self.net = self.net.to(self.device)
# 模型所需要的图片处理,模型改变可能需要改变
self.raw_transformer = transforms.Compose([
transforms.ToPILImage(),
transforms.Resize((256,128)),# hxw
transforms.ToTensor(),
])
def _preprocess(self, im_crops):
im_batch = torch.cat([self.raw_transformer(im).unsqueeze(0) for im in im_crops], dim=0)
return im_batch
def __call__(self, im_crops):
im_batch = self._preprocess(im_crops)
with torch.no_grad():
im_batch = im_batch.to(self.device)
features = self.net(im_batch)
return features

View File

@ -2,13 +2,16 @@ import torch
class Parameters(object):
def __init__(self) -> None:
self.weights = "./best.pt"
self.conf_thres = 0.35
self.iou_thres = 0.30
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
self.nosave = False
self.classes = [0]
self.weights_yolo = "./best.pt"
self.weights_reid = "./model9.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