神经网络搭建 的 基本构架介绍(附代码)
参考教程B站up主我是土堆如有侵权或其他问题欢迎留言联系更正或删除。1. dir ( ) 与 help ( ) 函数 —— 学习时常用import torch # 打印 package 下的子分类 print(dir(torch)) print(dir(torch.cuda)) print(dir(torch.cuda.is_available)) # 打印 函数的使用方法 print(help(torch.cuda.is_available))2. 快捷键Tab — 缩进、Ctrl — 查看函数具体信息、Ctrl / — 注释3. GPU训练# 利用GPU训练神经网络提升训练速度远快于CPU import torch # 方法一针对 “神经网络模型” “损失函数” “数据训练集、验证集的值及其对应标签” 使用 GPU 训练 if torch.cuda.is_available(): model_name model_name.cuda() loss_fn loss_fn.cuda() zhi zhi.cuda() label label.cuda() # 方法二设置 “训练设备” Device torch.device(cuda) # 若电脑含多张显卡 # 指定训练设备为第一个cuda Device torch.device(cuda:0) # 指定训练设备为第二个cuda Device torch.device(cuda:1) Device torch.device(cpu) # 语法糖 Device torch.device(cuda if torch.cuda.is_available() else cpu) # 以下两种写法等价 model_name model_name.to(Device) model_name.to(Device) loss_fn loss_fn.to(Device) loss_fn.to(Device) # 注意训练集、验证集的值及其对应标签 仅能以该形式设置训练设备 zhi zhi.to(Device) label label.to(Device)4. pytorch 框架下的 “数据加载”两个实用类dataset ( ) 及 dataloader ( )torch.utils.data.DataLoader功能构建可迭代的数据装载器datasetDataset 类决定数据从哪读取及如何读取batchsize批大小num_works是否多进程读取数据shuffle每个 epoch 是否乱序drop_last当样本数不能被 batchsize 整除时是否舍弃最后一批数据对应代码DataLoader( dataset, batch_size1, shuffleFalse, samplerNone, batch_samplerNone, num_workers0, collate_fnNone, pin_memoryFalse, drop_lastFalse, timeout0, worker_init_fnNone, multiprocessing_contextNone)注Epoch Iteration BatchsizeEpoch:所有训练样本都已输入至模型内一次称为一个 EpochIteration一批 (一个Batch) 样本输入至模型内称为一个 Iteration输入一个 Batch / 经历一次 Iteration更新一次模型参数Batchsize批 (Batch) 大小决定一个 Epoch 有多少个 Iteration也即num (Iteration) num (train_data) / Batchsize5. nn.Module神经网络搭建 的 基本骨架import torch from torch import nn # container —— 神经网络的基础模板可修改其为更复杂的结构 # 继承父类 nn.Module class model_name (nn.Module): def __init__(self): super().__init__() # 前向传播input-神经网络的输入 def forward(self,input): output input 1 return output # 模型实例创建 module_1 model_name() # 创建一个张量作为模型的输入 x torch.tensor(2.) # 输出经神经网络处理后的结果 y module_1(x) print(y)6. 卷积操作理解“卷积” 操作内的 步幅 stride 及 外围填充 padding 的设置空洞卷积可增加 “感受野”设置 dilation 参数Q写出output_1、output_2、output_3的结果import torch import torch.nn.functional as F # 创建被卷积的 tensor外围有n个[] → n维tensor input torch.tensor([[1,2,0,3,1], [0,1,2,3,1], [1,2,1,0,0], [5,2,3,1,1], [2,1,0,1,1]]) # 创建卷积核 kernel torch.tensor([[1,2,1], [0,1,0], [2,1,0]]) # 变幻被卷积的 tensor 和卷积核的 shape input torch.reshape(input, (1,1,5,5)) kernel torch.reshape(kernel,(1,1,3,3)) output_1 F.conv2d(input,kernel,stride1) print(output_1) # 调整步幅 stridestride可以为数字或元组如tuple(1,2) —— 横向步幅1纵向步幅2 output_2 F.conv2d(input,kernel,stride2) print(output_2) # 用 “0” 填充 input 的外围一层padding可以为数字或元组 output_3 F.conv2d(input,kernel,stride2,padding1) print(output_3)7. 池化操作Q写出output 的结果若 ceil_modeFalseoutput 又将为何值import torch from torch import nn from torch.nn import MaxPool2d # 创建被池化的 tensor外围有n个[] → n维tensor input torch.tensor([[1, 2, 0, 3, 1], [0, 1, 2, 3, 1], [1, 2, 1, 0, 0], [5, 2, 3, 1, 1], [2, 1, 0, 1, 1]], dtypetorch.float32) # 定义含 “池化层” 的网络模型 class pool_net(nn.Module): def __init__(self): super(pool_net, self).__init__() # 定义 “池化核” 的大小注意默认情况下stride kernel_size self.pool_layer MaxPool2d(kernel_size3, ceil_modeTrue) def forward(self, input): output self.pool_layer(input) return output # 变幻被池化的 tensor input torch.reshape(input, (-1, 1, 5, 5)) # 将 input 置于定义好的池化网络内 model pool_net() output model(input) print(output)8. 常用的非线性激活函数引入非线性激活函数的目的旨在帮助网络学习数据中的复杂模式对所有隐藏层及输出层添加 “非线性” 的操作使得神经网络的输出更为复杂、表达能力更强注意绝大多数神经网络借助某形式的梯度下降进行参数优化故激活函数需要是可微分的或者至少是几乎完全可微分的9. 线性层全连接层概念各神经元都与上下层各神经元相连一个简单的 “线性层全连接层” 如下所示具体函数# in_features out_features输入出特征数bias偏置项默认为True torch.nn.Linear (in_features, out_features, biasTrue, deviceNone, dtypeNone)tiptorch.flatten ( ) 函数被用于 “拉平” 矩阵10. Drop - out 操作丢弃部分数据避免过拟合11. Sequential 的使用Q练习写出下列图示的网络结构使用 pytorch 框架代码如下1不使用 Sequential 时的解答import torch from torch import nn from torch.nn import Flatten # 搭建图示网络 class M (nn.Module): # 定义后续将使用的网络模块 def __init__ (self): super(M, self).__init__() self.conv_1 nn.Conv2d(3, 32, 5, padding2) self.pool_1 nn.MaxPool2d(2) self.conv_2 nn.Conv2d(32, 32, 5, padding2) self.pool_2 nn.MaxPool2d(2) self.conv_3 nn.Conv2d(32, 64, 5, padding2) self.pool_3 nn.MaxPool2d(2) self.flat Flatten() self.linear_1 nn.Linear(1024, 64) self.linear_2 nn.Linear(64, 10) # 定义前向传播 def forward(self, input): input self.conv_1(input) input self.pool_1(input) input self.conv_2(input) input self.pool_2(input) input self.conv_3(input) input self.pool_3(input) input self.flat(input) input self.linear_1(input) input self.linear_2(input) return input # 初始化上述定义的神经网络 m1 M() # 测试 q torch.ones((64,3,32,32)) t m1(q) print(t.shape) # 预计输出torch.Size([64, 10])2使用 Sequential 时的解答有助于简化网络搭建过程import torch from torch import nn from torch.nn import Flatten, Sequential # 搭建图示网络 class M (nn.Module): # 定义后续将使用的网络模块 def __init__ (self): super(M, self).__init__() # 使用 Sequential有助于简化网络搭建过程如下所示以此类推 self.model1 Sequential ( nn.Conv2d(3, 32, 5, padding2), nn.MaxPool2d(2), nn.Conv2d(32, 32, 5, padding2), nn.MaxPool2d(2), nn.Conv2d(32, 64, 5, padding2), nn.MaxPool2d(2), Flatten(), nn.Linear(1024, 64), nn.Linear(64, 10) ) # 定义前向传播 def forward(self, input): input self.model1(input) return input # 初始化上述定义的神经网络 m1 M() # 测试 q torch.ones((64,3,32,32)) t m1(q) print(t.shape)12. Pytorch 下现有模型 的 引入及使用import torch import torchvision from torch import nn # 以 VGG16 神经网络为例 # 不含预训练参数的网络结构 vgg16_False torchvision.models.vgg16(pretrained False) # 含有预训练参数的网络结构 vgg16_True torchvision.models.vgg16(pretrained True) # 向 pytorch 框架提供的神经网络添加模块进行结构修改 vgg16_False.add_module(linear_1,nn.Linear(1000,10))13. 模型的保存及加载方法一同时保存模型的网络结构 及 模型参数import torch import torchvision # method 1 # 以 torchvision 内的 vgg 模型为例 model_vgg torchvision.models.vgg16(pretrainedFalse) # 模型保存输入 待保存的模型 及 模型保存的路径名称 torch.save(model_vgg, vgg.path) # 模型加载注意若“保存”与“加载”不在同一python文件内则需import model_load torch.load(vgg.path) print(model_load)方法二仅保存模型参数占用内存更小import torch import torchvision # method 2 # 以 torchvision 内的 vgg 模型为例 model_vgg torchvision.models.vgg16(pretrainedFalse) # 模型保存输入 待保存的模型 及 模型保存的路径名称 torch.save(model_vgg.state_dict(), vgg.path) # 模型加载 vgg_16 torchvision.models.vgg16(pretrainedFalse) model_load vgg_16.load_state_dict(torch.load(vgg.path)) print(model_load)14. 损失函数 及 优化器注损失函数 “指导”网络参数的优化更新借助于“优化器”代码示例如下

相关新闻

Loop for Mac:3步告别混乱窗口,打造高效macOS工作流的终极指南

Loop for Mac:3步告别混乱窗口,打造高效macOS工作流的终极指南

Loop for Mac:3步告别混乱窗口,打造高效macOS工作流的终极指南 【免费下载链接】Loop Window management made elegant. 项目地址: https://gitcode.com/GitHub_Trending/lo/Loop 你是否经常在十几个杂乱的窗口间迷失方向?每天花费宝贵…

2026/7/22 16:32:03 阅读更多 →
AI写小说长篇网文创作工具技术对比:从记忆系统、大纲到伏笔管理

AI写小说长篇网文创作工具技术对比:从记忆系统、大纲到伏笔管理

AI长篇网文创作工具技术对比:从记忆系统、大纲到伏笔管理 AI 写短篇已经不新鲜了。真正难啃的骨头是长篇——几十万字、上百个角色、纵横交错的伏笔。用 ChatGPT 或 Kimi 直接写,写到 30 章左右几乎必崩。本文从技术架构角度,对比 3 款试图解…

2026/7/23 16:06:38 阅读更多 →
AI写小说长篇创作中的上下文局限与外部记忆系统实践

AI写小说长篇创作中的上下文局限与外部记忆系统实践

简介: 本文从注意力衰减、检索盲区和时序冲突三个维度分析长上下文窗口在长篇AI创作中的结构性局限,并结合混合检索(BM25向量)、三层记忆压缩与时序衰减机制,探讨外部记忆系统的工程实现方法。文中以蛙趣拼文创作工作台…

2026/7/20 21:18:48 阅读更多 →

最新新闻

Kimi K3 发布引市场震动,AI 时代模型更迭谁能笑到最后?

Kimi K3 发布引市场震动,AI 时代模型更迭谁能笑到最后?

【Kimi K3 发布引发轰动】像所有 AI 行业人士一样,月之暗面创始人杨植麟喜欢用“AI 一天,人间一年”形容 AI 的迅速迭代,对于刚发布的 Kimi K3 而言,这句话尤为贴切。月之暗面成立三年,过去三年积累或许都不及这三天的…

2026/7/23 23:36:06 阅读更多 →
在半导体及电力电子器件(如 IGBT、SiC/GaN 功率模块、MOSFET 等)的可靠性测试中,无功老化测试机主要用来验证器件在承受高电压、大电流以及高频开关等电应力下的长期稳定性

在半导体及电力电子器件(如 IGBT、SiC/GaN 功率模块、MOSFET 等)的可靠性测试中,无功老化测试机主要用来验证器件在承受高电压、大电流以及高频开关等电应力下的长期稳定性

在半导体及电力电子器件(如 IGBT、SiC/GaN 功率模块、MOSFET 等)的可靠性测试中,无功老化测试机(Burn-in Test System)主要用来验证器件在承受高电压、大电流以及高频开关等电应力下的长期稳定性。 其中,“负载”(Load) 是整个测试回路的核心组成部分,其核心意义在于…

2026/7/23 23:36:06 阅读更多 →
深入了解MIMO

深入了解MIMO

文章目录1. 从SISO到MIMO2. MIMO有哪些类型?3. MU-MIMO和MIMO的区别是什么?4. Wi-Fi中的MIMO是如何工作的?5. 什么是M*N MIMO?6.802.11ac**MIMO(Multiple-Input Multiple-Output)**是指在无线通信领域使用多…

2026/7/23 23:36:06 阅读更多 →
自动感应门雷达人体接近检测解决方案

自动感应门雷达人体接近检测解决方案

一、方案背景自动感应门的人体接近检测是通行体验的核心。人来即开、人走缓关,免去推拉与接触,既方便又卫生,广泛用于商场、写字楼、医院与住宅入户。传统检测手段各有局限:红外对射容易被强光、雨雪与灰尘干扰,且需要…

2026/7/23 23:36:06 阅读更多 →
CentOS7.9‑Kickstart 无人值守安装结构化实战教程

CentOS7.9‑Kickstart 无人值守安装结构化实战教程

教程前言 Kickstart 是红帽系自动化部署工具,通过 ks.cfg 应答文件,预先填写系统安装过程所有选项,批量自动化安装 CentOS 系统;省去手动点击界面,适合服务器批量装机,企业机房批量部署必备技术。 本次环境…

2026/7/23 23:36:06 阅读更多 →
《静默登录》五、HarmonyOS_ArkTS开发避坑与修复指南

《静默登录》五、HarmonyOS_ArkTS开发避坑与修复指南

HarmonyOS ArkTS 开发避坑与修复指南本文基于真实项目(沉浸光感静默登录案例)中遇到的编译错误、运行时异常和架构设计问题,总结出 10 类高频陷阱 及其修复方案。每个陷阱均附有错误代码、正确代码和原理分析,帮助开发者在编码阶段…

2026/7/23 23:35:06 阅读更多 →

日新闻

从单点好评到指数级传播:AI副业主理人必须掌握的4层口碑渗透模型(含ROI测算表)

从单点好评到指数级传播:AI副业主理人必须掌握的4层口碑渗透模型(含ROI测算表)

更多请点击: https://intelliparadigm.com 第一章:从单点好评到指数级传播:AI副业主理人必须掌握的4层口碑渗透模型(含ROI测算表) 当AI副业主理人不再仅满足于单次服务交付,而是主动构建可复用、可裂变、可…

2026/7/23 0:00:25 阅读更多 →
AI写作开头钩子设计:为什么你的AI文案完读率不足18%?——基于2,346篇A/B测试报告的归因分析

AI写作开头钩子设计:为什么你的AI文案完读率不足18%?——基于2,346篇A/B测试报告的归因分析

更多请点击: https://codechina.net 第一章:AI写作开头钩子设计:为什么你的AI文案完读率不足18%?——基于2,346篇A/B测试报告的归因分析 在对2,346篇跨行业AI生成文案的A/B测试数据进行聚类分析后,我们发现&#xff1…

2026/7/23 0:01:26 阅读更多 →
Chitchatter完整指南:免费开源的终极点对点安全聊天工具

Chitchatter完整指南:免费开源的终极点对点安全聊天工具

Chitchatter完整指南:免费开源的终极点对点安全聊天工具 【免费下载链接】chitchatter Secure peer-to-peer chat that is serverless, decentralized, and ephemeral 项目地址: https://gitcode.com/gh_mirrors/ch/chitchatter Chitchatter是一款革命性的安…

2026/7/23 0:01:26 阅读更多 →

周新闻

Go语言静态资源打包方案对比与实践指南

Go语言静态资源打包方案对比与实践指南

1. 项目背景与核心需求在Go语言开发中,我们经常需要处理静态资源文件的打包问题。无论是Web应用的模板文件、前端资源,还是配置文件、证书等,都需要随程序一起分发。传统做法是将这些文件与编译后的二进制文件放在同一目录下,但这…

2026/7/22 8:58:19 阅读更多 →
Go语言实现高性能LDAP认证服务的架构与实践

Go语言实现高性能LDAP认证服务的架构与实践

1. 项目背景与核心价值LDAP(轻量级目录访问协议)作为企业级身份认证的黄金标准,已经服务了超过80%的财富500强公司。我在金融科技领域实施统一认证体系时,发现传统Java方案存在启动慢、内存占用高等痛点。而Go语言凭借其协程并发模…

2026/7/22 19:43:43 阅读更多 →
【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

更多请点击: https://intelliparadigm.com 第一章:AI面试官实战指南的核心价值与适用场景 AI面试官并非替代人类HR的“黑箱工具”,而是以可解释、可审计、可迭代的方式,赋能招聘全链路的关键基础设施。其核心价值在于将主观经验沉…

2026/7/23 17:49:47 阅读更多 →

月新闻