-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
86 lines (71 loc) · 2.76 KB
/
Copy pathtrain.py
File metadata and controls
86 lines (71 loc) · 2.76 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
import torchvision
from torch.utils.data import DataLoader
from torch import nn
import torch
from model import classify_model, vgg_16
# 定义设备
device = torch.device("cuda" if torch.cuda.is_available else "cpu")
print("当前设备 {}".format(device))
# 准备数据集
train_dataset = torchvision.datasets.CIFAR10("dataset",train=True,transform=torchvision.transforms.ToTensor(),download=True)
test_dataset = torchvision.datasets.CIFAR10("dataset",train=False,transform=torchvision.transforms.ToTensor(),download=True)
# 数据集长度
train_data_size = len(train_dataset)
test_data_size = len(test_dataset)
print("训练数据集的长度 {}".format(train_data_size))
print("测试数据集的长度 {}".format(test_data_size))
# 加载数据集
batch_size = 64
train_dataloader = DataLoader(train_dataset,batch_size=batch_size)
test_dataloader = DataLoader(test_dataset,batch_size=batch_size)
model_name = "vgg16" # vgg16 / classify_model
# model = classify_model() # 自定义神经网络
model = vgg_16() # 利用现有的网络
model.to(device)
print(model)
# 损失函数
loss_fn = nn.CrossEntropyLoss()
loss_fn.to(device)
# 优化器
learning_rate = 1e-2
optimizer = torch.optim.SGD(model.parameters(),lr=learning_rate)
# 设置网络参数
total_train_step = 0
total_test_step = 0
# 训练轮数
epoch = 100
for i in range(epoch):
print("-------第{}轮训练开始--------".format(i+1))
# 训练步骤开始
model.train() # 只对特定层起作用,如Dropout
for data in train_dataloader:
imgs, targets = data
imgs = imgs.to(device)
targets = targets.to(device)
output = model(imgs)
loss = loss_fn(output, targets)
# 优化器优化模型
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_train_step = total_train_step + 1
if(total_train_step % 100 == 0):
print("训练次数 {},Loss: {}".format(total_train_step, loss.item()))
# 测试步骤开始
total_test_loss = 0
total_test_acc = 0
model.eval()
with torch.no_grad():
for data in test_dataloader:
imgs, targets = data
imgs = imgs.to(device)
targets = targets.to(device)
output = model(imgs)
loss = loss_fn(output, targets)
total_test_loss = total_test_loss + loss.item()
acc = (output.argmax(1) == targets).sum()
total_test_acc = total_test_acc + acc
print("整体测试集上的Loss: {}".format(total_test_loss))
print("整体测试集上的Acc: {}".format(total_test_acc/test_data_size))
torch.save(model.state_dict(), "./check_point/{}/classify_model_{}.pth".format((model_name), (i+1)))
print("模型已保存")