Skip to content

Commit 665e8fe

Browse files
committed
ts
1 parent 6416f2e commit 665e8fe

4 files changed

Lines changed: 316 additions & 0 deletions

File tree

res50_backbone.py

Lines changed: 106 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,106 @@
1+
import torch.nn as nn
2+
import torch
3+
4+
5+
class Bottleneck(nn.Module):
6+
expansion = 4
7+
8+
def __init__(self, in_channel, out_channel, stride=1, downsample=None):
9+
super(Bottleneck, self).__init__()
10+
self.conv1 = nn.Conv2d(in_channels=in_channel, out_channels=out_channel,
11+
kernel_size=1, stride=1, bias=False) # squeeze channels
12+
self.bn1 = nn.BatchNorm2d(out_channel)
13+
# -----------------------------------------
14+
self.conv2 = nn.Conv2d(in_channels=out_channel, out_channels=out_channel,
15+
kernel_size=3, stride=stride, bias=False, padding=1)
16+
self.bn2 = nn.BatchNorm2d(out_channel)
17+
# -----------------------------------------
18+
self.conv3 = nn.Conv2d(in_channels=out_channel, out_channels=out_channel*self.expansion,
19+
kernel_size=1, stride=1, bias=False) # unsqueeze channels
20+
self.bn3 = nn.BatchNorm2d(out_channel*self.expansion)
21+
self.relu = nn.ReLU(inplace=True)
22+
self.downsample = downsample
23+
24+
def forward(self, x):
25+
identity = x
26+
if self.downsample is not None:
27+
identity = self.downsample(x)
28+
29+
out = self.conv1(x)
30+
out = self.bn1(out)
31+
out = self.relu(out)
32+
33+
out = self.conv2(out)
34+
out = self.bn2(out)
35+
out = self.relu(out)
36+
37+
out = self.conv3(out)
38+
out = self.bn3(out)
39+
40+
out += identity
41+
out = self.relu(out)
42+
43+
return out
44+
45+
46+
class ResNet(nn.Module):
47+
48+
def __init__(self, block, blocks_num, num_classes=1000, include_top=True):
49+
super(ResNet, self).__init__()
50+
self.include_top = include_top
51+
self.in_channel = 64
52+
53+
self.conv1 = nn.Conv2d(3, self.in_channel, kernel_size=7, stride=2,
54+
padding=3, bias=False)
55+
self.bn1 = nn.BatchNorm2d(self.in_channel)
56+
self.relu = nn.ReLU(inplace=True)
57+
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
58+
self.layer1 = self._make_layer(block, 64, blocks_num[0])
59+
self.layer2 = self._make_layer(block, 128, blocks_num[1], stride=2)
60+
self.layer3 = self._make_layer(block, 256, blocks_num[2], stride=2)
61+
self.layer4 = self._make_layer(block, 512, blocks_num[3], stride=2)
62+
if self.include_top:
63+
self.avgpool = nn.AdaptiveAvgPool2d((1, 1)) # output size = (1, 1)
64+
self.fc = nn.Linear(512 * block.expansion, num_classes)
65+
66+
for m in self.modules():
67+
if isinstance(m, nn.Conv2d):
68+
nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
69+
70+
def _make_layer(self, block, channel, block_num, stride=1):
71+
downsample = None
72+
if stride != 1 or self.in_channel != channel * block.expansion:
73+
downsample = nn.Sequential(
74+
nn.Conv2d(self.in_channel, channel * block.expansion, kernel_size=1, stride=stride, bias=False),
75+
nn.BatchNorm2d(channel * block.expansion))
76+
77+
layers = []
78+
layers.append(block(self.in_channel, channel, downsample=downsample, stride=stride))
79+
self.in_channel = channel * block.expansion
80+
81+
for _ in range(1, block_num):
82+
layers.append(block(self.in_channel, channel))
83+
84+
return nn.Sequential(*layers)
85+
86+
def forward(self, x):
87+
x = self.conv1(x)
88+
x = self.bn1(x)
89+
x = self.relu(x)
90+
x = self.maxpool(x)
91+
92+
x = self.layer1(x)
93+
x = self.layer2(x)
94+
x = self.layer3(x)
95+
x = self.layer4(x)
96+
97+
if self.include_top:
98+
x = self.avgpool(x)
99+
x = torch.flatten(x, 1)
100+
x = self.fc(x)
101+
102+
return x
103+
104+
105+
def resnet50(num_classes=1000, include_top=True):
106+
return ResNet(Bottleneck, [3, 4, 6, 3], num_classes=num_classes, include_top=include_top)

ssd_model.py

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,46 @@
1+
import torch
2+
from torch import nn, Tensor
3+
4+
5+
6+
class Backbone(nn.Module):
7+
def __init__(self, pretrain_path=None):
8+
super(Backbone, self).__init__() # Q:子类调用父类的方法
9+
net = resnet50()
10+
self.out_channels = [1024, 512, 512, 256, 256, 256] # 6个特征图的channel数目
11+
12+
# 使用预训练模型的权重
13+
if pretrain_path is not None:
14+
net.load_state_dict(torch.load(pretrain_path))
15+
16+
# 从resnet中取前7个模块作为feature extractor
17+
# 参考:https://blog.csdn.net/pengchengliu/article/details/113878358
18+
self.feature_extractor = nn.Sequential(*list(net.children())[:7])
19+
20+
# 对Conv4的第1个Block进行修改
21+
conv4_block1 = self.feature_extractor[-1][0]
22+
conv4_block1.conv1.stride = (1, 1)
23+
conv4_block1.conv2.stride = (1, 1)
24+
conv4_block1.downsample[0].stride = (1, 1)
25+
26+
def forward(self, x):
27+
x = self.feature_extractor(x)
28+
return x
29+
30+
31+
class SSD300(nn.Module):
32+
def __init__(self, backbone=None, num_classes=21):
33+
super(SSD300, self).__init__()
34+
if backbone==None:
35+
raise Exception('backbone is None')
36+
if not hasattr(backbone, 'out_channels'):
37+
raise Exception('backbone not has attribute: out_channels')
38+
self.feature_extractor = backbone
39+
40+
self.num_classes = num_classes
41+
# out_channels = [1024, 512, 512, 256, 256, 256]
42+
self._build_additional_features(self.feature_extractor.out_channels)
43+
44+
45+
def _build_additional_features(self, input_size):
46+

test.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
import torch
2+
from torch import nn, Tensor
3+
from torch.jit.annotations import List
4+
5+
from res50_backbone import resnet50
6+
7+
net = resnet50()
8+
9+
10+
print(*list(net.children())[:7])

train_ssd300.py

Lines changed: 154 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,160 @@
1+
import os
2+
import datetime
3+
4+
import torch
5+
6+
import transforms
7+
from my_dataset import VOCDataSet
8+
9+
10+
def create_model(num_classes=21):
11+
backbone = Backbone() # 特征提取器
12+
model = SSD300(backbone=backbone, num_classes=num_classes)
13+
14+
pre_ssd_path = '***'
15+
if os.path.exists(pre_ssd_path) is False:
16+
raise FileNotFoundError('*** not found in {}'.format(pre_ssd_path))
17+
pre_model_dict = torch.load(pre_ssd_path, map_location='cpu')
18+
pre_weights_dict = pre_model_dict['model']
19+
20+
# 不加载类别预测器的权重,因为是voc和coco不同,但是可以使用回归预测器的权重
21+
del_conf_loc_dict = {}
22+
for k, v in pre_weights_dict.items():
23+
split_key = k.split('.')
24+
if 'conf' in k:
25+
continue
26+
del_conf_loc_dict.update({k: v})
27+
28+
missing_keys, unexpected_keys = model.load_state_dict(del_conf_loc_dict, strict=False)
29+
if len(missing_keys) != 0 or len(unexpected_keys) != 0:
30+
print("missing_keys: ", missing_keys)
31+
print("unexpected_keys: ", unexpected_keys)
32+
33+
return model
34+
135

236
def main(parser_data):
37+
# 指定GPU或CPU,Q:如何进行多GPU训练
38+
device = torch.device(parser_data.device if torch.cuda.is_available() else 'cpu')
39+
print('Using device {} training.'.format(device.type))
40+
41+
if not os.path.exists('save_weights'):
42+
os.makedirs('save_weights')
43+
44+
results_file = 'results{}.txt'.format(datetime.datetime.now().strftime('%Y%m%d-%H%M%S'))
45+
46+
data_transform = {
47+
'train': transforms.Compose([
48+
transforms.SSDCropping(), # Q: 裁剪应该旨在增加训练样本,裁剪后GT也要改变?
49+
transforms.Resize(),
50+
transforms.ColorJitter(),
51+
transforms.ToTensor(),
52+
transforms.RandomHorizontalFlips(),
53+
transforms.Normalization(),
54+
# ???:输出default box集合的正样本和负样本(default_box的位置固定,正样本和GT匹配最佳且满足IOU>0.5)
55+
transforms.AssignGTtoDefaultBox()
56+
]),
57+
'val': transforms.Compose([
58+
transforms.Resize(),
59+
transforms.ToTensor(),
60+
transforms.Normalization()
61+
])
62+
}
63+
64+
VOC_root = parser_data.data_path
65+
if os.path.exists(os.path.join(VOC_root, 'VOCdevkit')) is False:
66+
raise FileNotFoundError('VOCdevkit does not in path: {}'.format(VOC_root))
67+
68+
# VOCdevkit/VOC2012/ImageSets/Main/train.txt
69+
train_dataset = VOCDataset(VOC_root, '2012', data_transform['train'], train_set='train.txt')
70+
# 训练时batch size必须大于1.
71+
batch_size = parser_data.batch_size
72+
assert batch_size > 1, 'batch size must be greater than 1' # Q:assert条件不满足,是否就推出程序?
73+
# 防止最后一个batch_size=1,如果是就舍去
74+
drop_last = True if len(train_dataset) % batch_size == 1 else False
75+
nw = min([os.cput_count(), batch_size if batch_size > 1 else 0, 8]) # number of workers
76+
print('Using %g dataloader workers' % nw)
77+
train_data_loader = torch.utils.data.DataLoader(
78+
train_dataset,
79+
batch_size=batch_size,
80+
shuffle=True,
81+
num_workers=nw,
82+
collate_fn=train_dataset.collate_fn, # 核对函数
83+
drop_last=drop_last
84+
)
85+
86+
# VOCdevkit/VOC2012/ImageSets/Main/val.txt
87+
val_dataset = VOCDataset(VOC_root, '2007', data_transform['val'], train_set='val.txt')
88+
val_data_loader = torch.utils.data.Dataloader(
89+
val_dataset,
90+
batch_size=batch_size,
91+
num_workers=nw,
92+
collate_fn=train_dataset.collate_fn # ?:为何不是val_dataset.collate_fn
93+
)
94+
95+
model = create_model(num_classes=parser_data.num_classes+1)
96+
model.to(device)
97+
98+
# define optimizer
99+
params = [p for p in model.parameters() if p.requires_grad] # Q:可以更新的参数部分?
100+
optimizer = torch.optim.SGD(params, lr=0.0005, momentum=0.9, weight_decay=0.0005)
101+
102+
# learning date scheduler,Q:step_size指每个epoch增大5倍?
103+
lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.3)
104+
105+
# 如果指定了上次训练保存的权重文件地址,则续接上次结果训练
106+
if parser_data.resume != '':
107+
check_point = torch.load(parser_data.resume, map_location='cpu') # Q:在cpu上加载保存于GPU的模型?
108+
model.load_state_dict(check_point['model'])
109+
optimizer.load_state_dict(check_point['optimizer'])
110+
lr_scheduler.load_state_dict(check_point['lr_scheduler'])
111+
parser_data.start_epoch = check_point['epoch'] + 1
112+
print('the training process from epoch {}...'.format(parser_data.start_epoch))
113+
114+
train_loss = []
115+
learning_rate = []
116+
val_map = []
117+
118+
# 提前加载验证集数据,以免每次验证都重新加载一次
119+
val_data = get_coco_api_from_dataset(val_data_loader.dataset) # Q:这里是说一次性加载所有的数据?
120+
for epoch in range(parser_data.start_epoch, parser_data.epochs):
121+
mean_loss, lr = utils.train_one_epoch(model=model,optimizer=optimizer, data_loader=train_data_loader, device=device, epoch=epoch, print_freq=50)
122+
train_loss.append(mean_loss.item())
123+
learning_rate.append(lr)
124+
125+
# update learning rate
126+
lr_scheduler.step()
127+
128+
# Q:是否使用data_set参数后就是一直把数据保存在GPU上,从而忽略了data_loader这种以batch加载数据的方式
129+
coco_info = utils.evaluate(model=model, data_loader=val_data_loader, device=device, data_set=val_data)
130+
131+
# write info txt
132+
with open(results_file, 'a') as f:
133+
result_info = [str(round(i, 4)) for i in coco_info + [mean_loss.item()]] + [str(round(lr, 6))]
134+
txt = "epoch:{} {}".format(epoch, ' '.join(result_info))
135+
f.write(txt + "\n")
136+
137+
val_map.append(coco_info[1])
138+
139+
# save weights
140+
save_files = {
141+
'model': model.state_dict(), # Q:state_dict()是什么意思?
142+
'optimizer': optimizer.state_dict(),
143+
'lr_scheduler': lr_scheduler.state_dict(),
144+
'epoch': epoch
145+
}
146+
147+
torch.save(save_files, './save_weights/ssd300-{}.pth'.format(epoch))
148+
149+
# plot loss and lr curve
150+
if len(train_loss) != 0 and len(learning_rate) != 0:
151+
from plot_curve import plot_loss_and_lr
152+
plot_loss_and_lr(train_loss, lr)
3153

154+
# plot mAP curve
155+
if len(val_map) != 0:
156+
from plot_curve import plot_map
157+
plot_map(val_map)
4158

5159

6160
if __name__ == '__main__':

0 commit comments

Comments
 (0)