Repository navigation
Expand file tree
/
Copy pathtrain.py
More file actions
122 lines (103 loc) · 4.58 KB
/
Copy pathtrain.py
File metadata and controls
122 lines (103 loc) · 4.58 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
import os
import time
import torch
import torch.nn as nn
from tqdm import tqdm
from torchvision import transforms, datasets
from torch.utils.data import DataLoader
from argparse import ArgumentParser
from torchinfo import summary
from archs import UNet
from metrics import iou_score
import params
def train(config, train_loader, model, criterion, optimizer):
model.train()
train_loss = 0.0
train_iou = 0.0
with tqdm(train_loader, desc='Training', unit='batch') as pbar:
for images, masks in train_loader:
# 将图像和标签送入
images, masks = images.to(config.device), masks.to(config.device)
outputs = model(images)
loss = criterion(outputs, masks)
optimizer.zero_grad() # 清空梯度
loss.backward() # 反向传播,计算梯度
optimizer.step() # 更新参数
iou = iou_score(outputs, masks)
train_iou += iou * images.size(0)
train_loss += loss.item() * images.size(0)
# time.sleep(10) # Simulate training tim
pbar.set_postfix(loss=loss.item(), iou=iou)
pbar.update(1)
train_loss = train_loss / len(train_loader.dataset)
train_iou = train_iou / len(train_loader.dataset)
return train_loss, train_iou
def validate(config, val_loader, model, criterion):
model.eval()
val_loss = 0.0
val_iou = 0.0
with torch.no_grad():
with tqdm(val_loader, desc='Validation', unit='batch') as pbar:
for images, masks in pbar:
images, masks = images.to(
config.device), masks.to(config.device)
outputs = model(images)
loss = criterion(outputs, masks)
val_loss += loss.item() * images.size(0)
iou = iou_score(outputs, masks)
val_iou += iou * images.size(0)
pbar.set_postfix(loss=loss.item(), iou=iou)
val_loss = val_loss / len(val_loader.dataset)
val_iou = val_iou / len(val_loader.dataset)
return val_loss, val_iou
def main():
config = params.parse_args()
# config.device = torch.device(
# 'cuda' if torch.cuda.is_available() else 'cpu')
if os.path.exists(config.save_path) == False:
os.makedirs(config.save_path)
def image_label_transform(image, label):
image_transform = transforms.Compose([
transforms.Resize((config.input_height, config.input_width)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[
0.229, 0.224, 0.225])
])
label_transform = transforms.Compose([
transforms.Resize((config.input_height, config.input_width)),
transforms.ToTensor(),
transforms.Lambda(lambda x: ((x > 0) & (x < 255)).float())
])
return image_transform(image), label_transform(label)
train_dataset = datasets.VOCSegmentation(
root='./data', image_set='train', download=False, transforms=image_label_transform)
val_dataset = datasets.VOCSegmentation(
root='./data', image_set='val', download=False, transforms=image_label_transform)
train_loader = DataLoader(
train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=4)
val_loader = DataLoader(
val_dataset, batch_size=config.batch_size, shuffle=False, num_workers=4)
model = UNet(num_classes=1).to(config.device)
criterion = nn.BCEWithLogitsLoss().cuda()
optimizer = torch.optim.Adam(model.parameters(), lr=config.learning_rate)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer, mode='min', factor=0.1, patience=5, verbose=True)
summary(model, input_size=(config.batch_size, 3,
config.input_height, config.input_width))
best_iou = 0.0
for epoch in range(config.epochs):
print('*' * 20)
train_loss, train_iou = train(
config, train_loader, model, criterion, optimizer)
val_loss, val_iou = validate(config, val_loader, model, criterion)
print(f'Epoch [{epoch+1}/{config.epochs}], Train Loss: {train_loss:.4f},'
f' Train IOU: {train_iou:.4f}, Val Loss: {val_loss:.4f}, Val IOU: {val_iou:.4f}')
scheduler.step(val_loss)
if val_iou > best_iou:
best_iou = val_iou
torch.save(model.state_dict(),
f'{config.save_path}/model.pth')
print(f'Model saved at epoch {epoch}')
torch.cuda.empty_cache()
if __name__ == '__main__':
main()