PyTorch深度学习入门笔记(小土堆)P26-32

发布时间:2026/9/26 2:07:44

PyTorch深度学习入门笔记(小土堆)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

相关新闻

软考网络工程师|第 5 章 TCP/UDP 完整备考笔记

软考网络工程师|第 5 章 TCP/UDP 完整备考笔记

2026/8/19 7:35:22

一、TCP 与 UDP 基础对比★★★★1 核心特性总览维度TCP(传输控制协议)UDP(用户数据报协议)连接属性面向连接,传输前建立连接无连接,直接发送报文可靠性可靠传输,重传、确认、有序不可靠尽力交付…

2026年青岛做城市生命线安全工程建设的厂家有哪些?

2026年青岛做城市生命线安全工程建设的厂家有哪些?

2026/9/8 16:42:38

青岛是沿海城市,地下管线规模大,燃气、供水管网建设标准高,汛期防潮防涝压力较大。海风带来的高湿盐雾环境对监测设备提出了特殊要求,既要防得住潮气,也要测得准管网数据,城市生命线安全工程的需求稳定而务…

2026年合肥做城市生命线安全工程建设的厂家有哪些?

2026年合肥做城市生命线安全工程建设的厂家有哪些?

2026/8/21 5:53:36

合肥作为长三角副中心城市,城市规模快速扩张,燃气管网、排水管网延伸迅速。作为「清华方案合肥模式」的诞生地,合肥在城市生命线安全工程建设上起步早、标准高,燃气、桥梁、供水、排水等专项监测的需求持续释放,示范效…

CANN/GE ACL数据集缓冲区添加函数

CANN/GE ACL数据集缓冲区添加函数

2026/9/25 10:06:33

aclmdlAddDatasetBuffer 【免费下载链接】ge GE(Graph Engine)是面向昇腾的图编译器和执行器,提供了计算图优化、多流并行、内存复用和模型下沉等技术手段,加速模型执行效率,减少模型内存占用。 GE 提供对 PyTorch、Te…

用ffmpeg高效批量调整图片尺寸的实战指南

用ffmpeg高效批量调整图片尺寸的实战指南

2026/9/25 9:40:47

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

Transformers 音频特征提取工具库 audio_utils 全解析:从 Mel 刻度换算到对数 Mel 频谱

Transformers 音频特征提取工具库 audio_utils 全解析:从 Mel 刻度换算到对数 Mel 频谱

2026/9/25 10:06:21

Transformers 音频特征提取工具库 audio_utils 全解析:从 Mel 刻度换算到对数 Mel 频谱 【免费下载链接】transformers 🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and mu…

RustFS 多节点集群重启与滚动升级实战:Readiness、Quorum 与 Degraded 模式完全指南

RustFS 多节点集群重启与滚动升级实战:Readiness、Quorum 与 Degraded 模式完全指南

2026/9/25 9:53:52

RustFS 多节点集群重启与滚动升级实战:Readiness、Quorum 与 Degraded 模式完全指南 【免费下载链接】rustfs 🚀2.3x faster than MinIO for 4KB object payloads. RustFS is an open-source, S3-compatible high-performance object storage system sup…

Java Integer缓存揭秘:128陷阱原理、避坑与面试全解

Java Integer缓存揭秘:128陷阱原理、避坑与面试全解

2026/9/25 8:58:17

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

RustFS Scanner 数据用量发布权威性决策:配额准入如何获得可用的权威依据

RustFS Scanner 数据用量发布权威性决策:配额准入如何获得可用的权威依据

2026/9/25 10:00:17

RustFS Scanner 数据用量发布权威性决策:配额准入如何获得可用的权威依据 【免费下载链接】rustfs 🚀2.3x faster than MinIO for 4KB object payloads. RustFS is an open-source, S3-compatible high-performance object storage system supporting mi…

远程协作的工作台整理

远程协作的工作台整理

2026/9/24 16:02:49

远程协作的工作台整理远程协作的核心不是再加一个工具,而是让交接信息足够完整。异步任务要写明目标、输入位置、完成标准和需要决策的人。 工作台的最小配置 将日程、待办、代码和沟通入口收拢到少数固定位置;通知按紧急程度分层。工作台不需要模仿办公…

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

2026/9/25 9:41:47

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能分类:[AI/大模型]细分主题:AI 增强型 CI/CD 流水线自动化与 GitOps 实践:Agent 工作流、工具调用与任务拆解:从原型到生产的验收清单很多团队在尝试用大…

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

2026/9/25 4:22:14

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场分类:[工程技术]细分主题:Kubernetes 生产环境运维与排障实战:可复制的项目复盘模板与决策记录大部分团队的事故复盘报告,最后都变成了躺在 Confluence 或钉…