diff --git a/deep_sort/deep/feature_extractor.py b/deep_sort/deep/feature_extractor.py index 5124c7e..82d810d 100644 --- a/deep_sort/deep/feature_extractor.py +++ b/deep_sort/deep/feature_extractor.py @@ -5,7 +5,7 @@ import numpy as np import cv2 import logging -from .model import Net,MyNet +from .model import Net,MyNet,RestNet18 from fastreid.config import get_cfg from fastreid.engine import DefaultTrainer from fastreid.utils.checkpoint import Checkpointer @@ -90,14 +90,14 @@ class FastReIDExtractor(object): 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 = RestNet18() self.net.load_state_dict(torch.load(model_path)) self.net = self.net.to(self.device) logger = logging.getLogger("root.tracker") logger.info("Loading weights from {}... Done!".format(model_path)) self.raw_transformer = transforms.Compose([ transforms.ToPILImage(), - transforms.Resize((256,128)),# hxw + transforms.Resize((128,64)),# hxw transforms.ToTensor(), ]) self.norm = transforms.Compose([ diff --git a/deep_sort/deep/model.py b/deep_sort/deep/model.py index c565966..1245fe2 100644 --- a/deep_sort/deep/model.py +++ b/deep_sort/deep/model.py @@ -150,5 +150,79 @@ if __name__ == '__main__': net = Net() x = torch.randn(4,3,128,64) y = net(x) + +#########ResNet18 +class RestNetBasicBlock(nn.Module): + def __init__(self, in_channels, out_channels, stride): + super(RestNetBasicBlock, self).__init__() + self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1) + self.bn1 = nn.BatchNorm2d(out_channels) + self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=stride, padding=1) + self.bn2 = nn.BatchNorm2d(out_channels) + + def forward(self, x): + output = self.conv1(x) + output = F.relu(self.bn1(output)) + output = self.conv2(output) + output = self.bn2(output) + return F.leaky_relu(x + output) + + +class RestNetDownBlock(nn.Module): + def __init__(self, in_channels, out_channels, stride): + super(RestNetDownBlock, self).__init__() + self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride[0], padding=1) + self.bn1 = nn.BatchNorm2d(out_channels) + self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=stride[1], padding=1) + self.bn2 = nn.BatchNorm2d(out_channels) + self.extra = nn.Sequential( + nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride[0], padding=0), + nn.BatchNorm2d(out_channels) + ) + + def forward(self, x): + extra_x = self.extra(x) + output = self.conv1(x) + out = F.relu(self.bn1(output)) + + out = self.conv2(out) + out = self.bn2(out) + return F.leaky_relu(extra_x + out) + + +class RestNet18(nn.Module): + def __init__(self): + super(RestNet18, self).__init__() + self.conv1 = nn.Conv2d(3, 32, kernel_size=(3, 3), padding=1) + self.bn1 = nn.BatchNorm2d(32) + self.maxpool = nn.MaxPool2d(2) + + self.layer1 = nn.Sequential(RestNetBasicBlock(32, 32, 1), + RestNetBasicBlock(32, 32, 1)) + + self.layer2 = nn.Sequential(RestNetDownBlock(32, 64, [2, 1]), + RestNetBasicBlock(64, 64, 1)) + + self.layer3 = nn.Sequential(RestNetDownBlock(64, 128, [2, 1]), + RestNetBasicBlock(128, 128, 1)) + + self.layer4 = nn.Sequential(RestNetDownBlock(128, 256, [2, 1]), + RestNetBasicBlock(256, 256, 1)) + + self.avgpool = nn.AdaptiveAvgPool2d(output_size=(1, 1)) + + self.fc = nn.Linear(32768, 23) + + def forward(self, x): + x = self.conv1(x) + x = self.layer1(x) + x = self.layer2(x) + x = self.layer3(x) + x = self.layer4(x) + # x = self.avgpool(x) + x = x.reshape(x.shape[0], -1) + # print(x.size()) + x = self.fc(x) + return x diff --git a/param.py b/param.py index 746f303..ff8f530 100644 --- a/param.py +++ b/param.py @@ -3,7 +3,7 @@ import torch class Parameters(object): def __init__(self) -> None: self.weights_yolo = "./best.pt" - self.weights_reid = "./model9.pth" + 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") @@ -12,6 +12,6 @@ class Parameters(object): self.agnostic_nms = None self.augment = None self.query_index = 1 - self.gallary_index= 1 + self.gallary_index= 2 pass \ No newline at end of file