-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain_optimized.py
More file actions
172 lines (138 loc) · 5.1 KB
/
Copy pathtrain_optimized.py
File metadata and controls
172 lines (138 loc) · 5.1 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
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
import torchvision
from torch.utils.data import DataLoader
from torch import nn
import torch
from model import classify_model, vgg_16
import matplotlib.pyplot as plt
import os
from tqdm import tqdm
# 定义设备
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"当前设备 {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(f"训练数据集的长度 {train_data_size}")
print(f"测试数据集的长度 {test_data_size}")
# 加载数据集
batch_size = 64
train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_dataloader = DataLoader(test_dataset, batch_size=batch_size)
model_name = "vgg16"
# 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 = 20
# 创建目录保存结果
os.makedirs("./check_point/", exist_ok=True)
os.makedirs(f"./check_point/{model_name}/", exist_ok=True)
# 初始化记录列表
train_losses = []
test_losses = []
test_accuracies = []
# 使用tqdm包装epoch循环
epoch_progress = tqdm(range(epoch), desc="总训练进度", position=0)
for i in epoch_progress:
epoch_progress.set_description(f"第 {i+1}/{epoch} 轮训练")
# 训练步骤
model.train()
epoch_train_loss = 0.0
# 使用tqdm包装训练批次循环
train_batch_progress = tqdm(train_dataloader, desc=f"训练批次", leave=False, position=1)
for data in train_batch_progress:
imgs, targets = data
imgs, targets = imgs.to(device), targets.to(device)
output = model(imgs)
loss = loss_fn(output, targets)
epoch_train_loss += loss.item()
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_train_step += 1
# 更新批次进度条描述
train_batch_progress.set_postfix({"批次损失": f"{loss.item():.4f}"})
# 计算平均训练损失
avg_train_loss = epoch_train_loss / len(train_dataloader)
train_losses.append(avg_train_loss)
# 更新总进度条描述
epoch_progress.set_postfix({"训练损失": f"{avg_train_loss:.4f}"})
# 测试步骤
total_test_loss = 0.0
total_correct = 0
model.eval()
# 使用tqdm包装测试批次循环
test_batch_progress = tqdm(test_dataloader, desc="测试批次", leave=False, position=2)
with torch.no_grad():
for data in test_batch_progress:
imgs, targets = data
imgs, targets = imgs.to(device), targets.to(device)
output = model(imgs)
loss = loss_fn(output, targets)
total_test_loss += loss.item()
_, predicted = torch.max(output.data, 1)
total_correct += (predicted == targets).sum().item()
# 更新测试批次进度条描述
test_batch_progress.set_postfix({"测试损失": f"{loss.item():.4f}"})
# 计算测试指标
avg_test_loss = total_test_loss / len(test_dataloader)
test_losses.append(avg_test_loss)
accuracy = total_correct / test_data_size
test_accuracies.append(accuracy)
# 更新总进度条描述
epoch_progress.set_postfix({
"训练损失": f"{avg_train_loss:.4f}",
"测试损失": f"{avg_test_loss:.4f}",
"测试准确率": f"{accuracy:.4f}"
})
# 保存模型
torch.save(model.state_dict(), f"./check_point/{model_name}/classify_model_{i+1}.pth")
# 每10个epoch保存一次图表
if (i + 1) % 10 == 0 or i == epoch - 1:
# 创建图表
plt.figure(figsize=(12, 5))
# 损失图表
plt.subplot(1, 2, 1)
plt.plot(range(1, i+2), train_losses, label='Training loss')
plt.plot(range(1, i+2), test_losses, label='Test loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.title('Training and testing losses')
plt.legend()
plt.grid(True)
# 准确率图表
plt.subplot(1, 2, 2)
plt.plot(range(1, i+2), test_accuracies, label='Test accuracy', color='green')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.title('Test accuracy')
plt.ylim(0, 1.0)
plt.legend()
plt.grid(True)
plt.tight_layout()
plt.savefig(f'./check_point/{model_name}/training_metrics_epoch_{i+1}.png')
plt.close()
print(f"\n训练指标图表已保存到 ./check_point/{model_name}/training_metrics_epoch_{i+1}.png")
print("\n训练完成,所有指标图表已保存")