构建神经网络(代码实现)

发布时间:2026/10/2 8:22:36
构建神经网络(代码实现) 神经网络优缺点:优点:1. 精度高.2. 可以近似任意的非线性函数.3. 有大量的框架和库可以调用.缺点:1. 黑箱.2. 训练时间长, 需要大量的算力.3. 网络结构复杂, 需要调整超参.4. 小数据集上表现不佳, 容易发生过拟合.# 导包import torchimport torch.nn as nn # 线性模型, 初始化方法都在这里.from torchsummary import summary # 模型可视化(统计模型参数的), 你要额外安装, 即: pip install torchsummary# todo 1. 创建1个类, 继承 nn.Moduleclass ModelDemo(nn.Module):# todo 1.1 定义 __init__() 方法, 构建神经网络.def __init__(self):# 1. 初始化父类成员.super().__init__()# 2. 创建隐藏层 和 输出层.# 2.1 搭建: 隐藏层1, in_features: 输入的特征数, 即: 上一层神经元个数. out_features: 输出的特征数, 即: 当前层神经元个数.self.linear1 nn.Linear(in_features3, out_features3)# 2.2 搭建: 隐藏层2. 隐藏层一般用 linear 或者 fc(fully connected, 全连接层)表示.self.linear2 nn.Linear(in_features3, out_features2)# 2.3 搭建: 输出层,output代替图中linear3self.output nn.Linear(in_features2, out_features2)# 3. 对隐藏层进行参数初始化# 3.1 隐藏层1: 标准的xavier初始化, 激活函数用 Sigmoidnn.init.xavier_normal_(tensorself.linear1.weight)nn.init.zeros_(tensorself.linear1.bias) # 不考虑偏置项.# 3.2 隐藏层2: 标准的He初始化, 激活函数用 ReLUnn.init.kaiming_normal_(tensorself.linear2.weight, nonlinearityrelu)nn.init.zeros_(tensorself.linear2.bias)# todo 1.2 定义前向传播方法 forward(), 得到预测值. - 注意: 方法名固定, 不能改.def forward(self, x): # x表示输入样本的 特征值.# 1. 第1层(隐藏层1), 加权求和 和 激活函数计算.# 激活函数 线性加权求和x torch.sigmoid(inputself.linear1(x))# 2. 第2层(隐藏层2), 加权求和 和 激活函数计算.x torch.relu(inputself.linear2(x))# 3. 输出层计算, 假设 多分类问题.# dim-1解释: 按行计算, 一个样本一个样本计算.x torch.softmax(inputself.output(x), dim-1)# 4. 返回 预测值.return x# todo 2. 创建模型训练函数.def train():# 1. 创建 神经网络 模型对象.my_model ModelDemo()print(fmy_model - {my_model}) # 打印模型结构.# 2. 构建数据集样本, 随机生成.data torch.randn(size(5, 3))print(fdata - {data})print(fdata.shape - {data.shape}) # torch.Size([5, 3])print(fdata.requires_grad: {data.requires_grad}) # False# 3. 调用 神经网络模型对象 进行模型训练.output my_model(data)print(foutput - {output})print(foutput.shape - {output.shape}) # torch.Size([5, 2])print(foutput.requires_grad: {output.requires_grad}) # Trueprint(♥️ * 15)# 4. 计算模型参数.# model: (自定义的 神经网络)模型对象.# input_size: 输入数据的特征数(即: 样本的特征数)# batch_size: 批次大小(即: 批次训练的样本数)# summary(modelmy_model, input_size(5, 3))summary(modelmy_model, input_size(3,), batch_size5)print(♥️ * 15)# 5. 查看模型参数.for name, param in my_model.named_parameters():print(fname: {name})print(fparam: {param}\n)# todo 3. 测试代码.if __name__ __main__:train()

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询