P26-32)
PyTorch深度学习入门笔记P26-32ZZHow(ZZHow1024)参考课程【PyTorch深度学习快速入门教程【小土堆】】[https://www.bilibili.com/video/BV1hE411t7RN]P26. 完整的模型训练套路一训练部分model.pyimporttorchfromtorchimportnnfromtorch.nnimportSequential# 搭建神经网络classMyModel(torch.nn.Module):def__init__(self):super(MyModel,self).__init__()self.modelSequential(nn.Conv2d(in_channels3,out_channels32,kernel_size5,stride1,padding2),nn.MaxPool2d(kernel_size2),nn.Conv2d(in_channels32,out_channels32,kernel_size5,stride1,padding2),nn.MaxPool2d(kernel_size2),nn.Conv2d(in_channels32,out_channels64,kernel_size5,stride1,padding2),nn.MaxPool2d(kernel_size2),nn.Flatten(),nn.Linear(in_features64*4*4,out_features64),nn.Linear(in_features64,out_features10),)defforward(self,x):xself.model(x)returnx# 测试神经网络模型结构的正确性if__name____main__:modelMyModel()inputtorch.ones([64,3,32,32])outputmodel(input)print(output.shape)train.pyimporttorchimporttorchvision.datasetsfrommodelimportMyModel# 准备数据集train_datatorchvision.datasets.CIFAR10(dataset,trainTrue,transformtorchvision.transforms.ToTensor(),downloadTrue)test_datatorchvision.datasets.CIFAR10(dataset,trainFalse,transformtorchvision.transforms.ToTensor(),downloadTrue)# 获取数据集的长度train_data_sizelen(train_data)test_data_sizelen(test_data)print(f训练数据集的长度为{train_data_size})print(f测试数据集的长度为{test_data_size})# 使用 Dataloader 加载数据集train_dataloadertorch.utils.data.DataLoader(train_data,batch_size64)test_dataloadertorch.utils.data.DataLoader(test_data,batch_size64)# 创建网络模型modelMyModel()# 损失函数loss_fntorch.nn.CrossEntropyLoss()# 优化器learning_rate1e-2optimizertorch.optim.SGD(model.parameters(),lrlearning_rate)# 设置训练网络的参数total_train_step0# 训练次数total_test_step0# 测试次数epoch10# 训练轮次foriinrange(epoch):print(f---第{i1}轮训练开始---)# 训练步骤开始fordataintrain_dataloader:images,targetsdata outputsmodel(images)lossloss_fn(outputs,targets)# 优化器优化模型optimizer.zero_grad()loss.backward()optimizer.step()total_train_step1print(f训练次数{total_train_step}Loss{loss.item()})P27. 完整的模型训练套路二测试验证部分# 测试步骤开始total_test_loss0# 总测试 Losstotal_accuracy0# 总正确率withtorch.no_grad():fordataintest_dataloader:images,targetsdata outputsmodel(images)lossloss_fn(outputs,targets)total_test_lossloss.item()accuracy(outputs.argmax(1)targets).sum()total_accuracyaccuracy writer.add_scalar(test_loss,total_test_loss,total_test_step)writer.add_scalar(test_accuracy,total_accuracy/test_data_size,total_test_step)print(f测试集上的总 Loss{total_test_loss})print(f测试集上的总 正确率{total_accuracy/test_data_size})torch.save(model.state_dict(),os.path.join(model,fmodel_{i}.pth))print(f模型已保存文件名model_{i}.pth)total_test_step1P28. 完整的模型训练套路三训练步骤开始时model.train()测试步骤开始时model.eval()案例演示model.py和train.pyP29. 利用GPU训练一方式一在网络模型、数据输入标注和损失函数后加上.cuda()# 创建网络模型modelMyModel()iftorch.cuda.is_available():modelmodel.cuda()# 损失函数loss_fntorch.nn.CrossEntropyLoss()iftorch.cuda.is_available():loss_fnloss_fn.cuda()# 数据输入标注images,targetsdataiftorch.cuda.is_available():imagesimages.cuda()targetstargets.cuda()案例演示train_gpu_1.pyP30. 利用GPU训练二方式二在网络模型、数据输入标注和损失函数后通过.to(device)转移到对应设备# 训练设备devicecpuiftorch.cuda.is_available():devicecudaeliftorch.mps.is_available():devicempsprint(f训练设备{device})# 创建网络模型modelMyModel()model.to(device)# 损失函数loss_fntorch.nn.CrossEntropyLoss()loss_fnloss_fn.to(device)# 数据输入标注images,targetsdata imagesimages.to(device)targetstargets.to(device)案例演示train_gpu_2.pyP31. 完整的模型验证套路test.pyimportosimporttorchimporttorchvisionfromPILimportImagefrommodelimportMyModel# 测试图片名称image_namedog.png# 测试模型名称model_namemodel_29.pth# 测试设备devicecpuiftorch.cuda.is_available():devicecudaeliftorch.mps.is_available():devicempsprint(f测试设备{device})# 测试图片路径image_pathos.path.join(images,image_name)imageImage.open(image_path)imageimage.convert(RGB)print(image)# 图片预处理transformtorchvision.transforms.Compose([torchvision.transforms.Resize((32,32)),torchvision.transforms.ToTensor()])imagetransform(image)imagetorch.reshape(image,(1,3,32,32))print(image.shape)# 加载模型modelMyModel()model.load_state_dict(torch.load(os.path.join(model,model_name),map_locationtorch.device(device)))# 开始测试model.eval()withtorch.no_grad():outputmodel(image)print(output)print(output.argmax(1))注意若训练模型的设备与当前加载加载模型的设备不一致时需要在torch.load()时指定map_locationtorch.device(device)。案例演示test.py