恒美微站 Logo 恒美微站
  • 首页
  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心
  • 联系我们

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

  • 首页
  • 资讯中心
  • /
  • PyTorch深度学习入门笔记(小土堆)P26-32

相关资讯

软考网络工程师|第 5 章 TCP/UDP 完整备考笔记 2026/8/6 19:51:59
2026年青岛做城市生命线安全工程建设的厂家有哪些? 2026/8/6 19:51:59
2026年合肥做城市生命线安全工程建设的厂家有哪些? 2026/8/6 19:51:59

最新资讯

txtai:一站式AI框架如何用3个核心功能改变语义搜索和LLM应用开发
Stats.js前端性能监控实战与Three.js优化指南
SimplerEnv环境下的机器人任务突破:INTACT-pi0-finetune-bridge性能深度测评
KMinion安全配置实战:TLS加密与SASL认证保护Kafka监控数据
Unity游戏开发中的解释器模式:构建可配置技能系统的核心技术
编导老师智能体:内容创作提效助手

今日推荐

电力系统调度中的源荷不确定性建模与优化实践
VGG-T3技术解析:3D重建速度的革命性突破
深度解析旅游网站建设的意义及其对行业发展的深远影响与核心价值体现

本周热门

ncmdumpGUI:一键解锁网易云音乐ncm文件的终极解决方案
分布式配置中心选型实战:Nacos与Consul在创业场景下的对比
MoneyPrinterPlus实战指南:AI视频批量生成与自动化发布完整解决方案

本月精选

如何用DamaiHelper实现演唱会门票的智能自动化抢购:完整技术解决方案指南
第4篇:59 倍性能差距的索引瓶颈定位——一次教科书级的全表扫描调优
终极歌词批量下载神器:5分钟解决离线音乐库歌词同步难题

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

发布时间:2026/8/6 19:51:59
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

关于恒美微站

恒美微站专注于为个体商户、工作室提供极简自助建站服务,让每个人都能轻松拥有专业网站。

快速链接

  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心

服务项目

  • 可视化建站
  • 拖拽编辑
  • 主题定制
  • SEO 优化
  • 网站托管

联系方式

  • 📍 地址:北京市朝阳区建国路 88 号
  • 📞 电话:400-888-8888
  • ✉️ 邮箱:info@hmyw.cn
  • 🕐 时间:周一至周日 9:00-18:00

© 2024 恒美微站 hmyw.cn 版权所有 | 京 ICP 备 12345678 号