数据集下载:通过代码自动下载即可
import torch
import torch.nn as nn
import torchvision
import torchvision.transforms as transforms
import numpy as np
import pandas as pd
from matplotlib import pyplot as plt
import torch.nn.functional as F
# 定义归一化参数
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) # 第一个 rgb均值
# 第二个是rgb三通道的方差
train_set = torchvision.datasets.CIFAR10(root='./data',
train=True, download=True, transform=transform)
train_loader = torch.utils.data.DataLoader(train_set,
batch_size=16, shuffle=True, num_workers=0)
# 测试数据
test_set = torchvision.datasets.CIFAR10(root='./data',
train=False, download=True, transform=transform)
test_loader = torch.utils.data.DataLoader(test_set,
batch_size=16, shuffle=False, num_workers=0)
dataiter = iter(train_loader)
images, labels = dataiter.next()
def show_image(img):
# 输入 [c, h, w]
img_np = img.numpy()# 变 numpy
img_np = np.transpose(img_np, (1, 2, 0))# [h, w, c]
plt.imshow(img_np)
show_image(torchvision.utils.make_grid(images))

train_loss_hist = []
test_loss_hist = []
# 建议20次
for epoch in tqdm(range(2)):
cnn.train() #训练部分
running_loss = 0.0 #batch损失累计
for i , data in enumerate(train_loader):
images, labels = data
outputs = cnn(images)
loss = criterion(outputs, labels) # 损失求和
optimizer.zero_grad()
loss.backward()
optimizer.step()
running_loss += loss.item() # 一个batch的损失
if (i%250 == 0):
# 每250个 250 x 16 进行一次记录 和测试
cnn.eval() #验证部分
with torch.no_grad():
for test_data in test_loader:
test_images, test_label = test_data
test_outputs = cnn(test_images)
test_loss = criterion(test_outputs, test_label)
train_loss_hist.append(running_loss/250) # 250个batch的平均损失
test_loss_hist.append(test_loss.item())
running_loss = 0.0
#torch.save(cnn,"model/cnn_image_model_epoch_{}_step{}.pkl".format(epoch,i))
print("epoch:{} step:{} train Loss:{} test Loss:{}".format(epoch+1,i,loss.item(),test_loss.item()))
torch.save(cnn,"model/cnn_image_model_epoch_{}_train Loss_{}_test_Loss_{}.pkl".format(epoch,loss.item(),test_loss.item())) #每个epoch保存一次
print("epoch:{} train Loss:{} test Loss:{}".format(epoch+1,loss.item(),test_loss.item()))
class CNN(nn.Module):
def __init__(self):
# 输入频道的结构 3 x 32 x 32
super(CNN, self).__init__()
# 卷积层
self.conv1 = nn.Conv2d(3, 6, 3)
# 卷积层
self.conv2 = nn.Conv2d(6, 16, 3)
# 全连接层 12544
self.fc1 = nn.Linear(16*28*28, 512)
self.fc2 = nn.Linear(512, 64)
self.fc3 = nn.Linear(64, 10)
def forward(self, x):
x = self.conv1(x)
x = F.relu(x)
x = self.conv2(x)
x = F.relu(x)
# fc reshape
x = x.view(-1, 16*28*28)
x = self.fc1(x)
x = F.relu(x)
x = self.fc2(x)
x = F.relu(x)
x = self.fc3(x)
return x
cnn = CNN()
cnn
CNN(
(conv1): Conv2d(3, 6, kernel_size=(3, 3), stride=(1, 1))
(conv2): Conv2d(6, 16, kernel_size=(3, 3), stride=(1, 1))
(fc1): Linear(in_features=12544, out_features=512, bias=True)
(fc2): Linear(in_features=512, out_features=64, bias=True)
(fc3): Linear(in_features=64, out_features=10, bias=True)
)
import torch.optim as optim
criterion = nn.CrossEntropyLoss()# 损失
optimizer = optim.Adam(cnn.parameters(), lr = 0.001)
train_loss_hist = []
test_loss_hist = []
# 建议20次
for epoch in tqdm(range(2)):
cnn.train() #训练部分
running_loss = 0.0 #batch损失累计
for i , data in enumerate(train_loader):
images, labels = data
outputs = cnn(images)
loss = criterion(outputs, labels) # 损失求和
optimizer.zero_grad()
loss.backward()
optimizer.step()
running_loss += loss.item() # 一个batch的损失
if (i%250 == 0):
# 每250个 250 x 16 进行一次记录 和测试
cnn.eval() #验证部分
with torch.no_grad():
for test_data in test_loader:
test_images, test_label = test_data
test_outputs = cnn(test_images)
test_loss = criterion(test_outputs, test_label)
train_loss_hist.append(running_loss/250) # 250个batch的平均损失
test_loss_hist.append(test_loss.item())
running_loss = 0.0
#torch.save(cnn,"model/cnn_image_model_epoch_{}_step{}.pkl".format(epoch,i))
print("epoch:{} step:{} train Loss:{} test Loss:{}".format(epoch+1,i,loss.item(),test_loss.item()))
torch.save(cnn,"model/cnn_image_model_epoch_{}_train Loss_{}_test_Loss_{}.pkl".format(epoch,loss.item(),test_loss.item())) #每个epoch保存一次
print("epoch:{} train Loss:{} test Loss:{}".format(epoch+1,loss.item(),test_loss.item()))

通过工具类:直接返回numpy格式的数据和标签