惯性聚合 高效追踪和阅读你感兴趣的博客、新闻、科技资讯
阅读原文 在惯性聚合中打开

推荐订阅源

V
Visual Studio Blog
D
DataBreaches.Net
博客园 - 三生石上(FineUI控件)
博客园_首页
T
Tailwind CSS Blog
美团技术团队
Hugging Face - Blog
Hugging Face - Blog
博客园 - 叶小钗
大猫的无限游戏
大猫的无限游戏
OSCHINA 社区最新新闻
OSCHINA 社区最新新闻
云风的 BLOG
云风的 BLOG
奇客Solidot–传递最新科技情报
奇客Solidot–传递最新科技情报
博客园 - 聂微东
S
SegmentFault 最新的问题
小众软件
小众软件
酷 壳 – CoolShell
酷 壳 – CoolShell
N
Netflix TechBlog - Medium
Jina AI
Jina AI
WordPress大学
WordPress大学
U
Unit 42
J
Java Code Geeks
Blog — PlanetScale
Blog — PlanetScale
钛媒体:引领未来商业与生活新知
钛媒体:引领未来商业与生活新知
The Cloudflare Blog

Data4Fun

要不要把大模型搬回本地 AI 新手村:让大模型学会操作浏览器 Data4Fun Data4Fun Data4Fun Data4Fun Data4Fun 当我使用Claude Code时我在想什么 如何使用Skill_Seeker助力 Skills 开发 使用 pyspark 处理数据的基本流程 AI新手村:Claude Code 飞桨 AI Studio:一步步微调你的大模型 AI 新手村:CLIP AI新手村:LLM 从一张表格谈起 AI新手村:MCP AI 新手村:Embedding AI新手村:Atlas入门 长沙印象 开始炼丹,如何快速训练一个神经网络 初探 YOLOv1
如何快速建立一个神经网络
Shaoyang · 2024-08-28 · via Data4Fun

整个网络的搭建基于pytorch的框架,其中 torch.nn 的命名空间包含了所有构建神经网络需要的基础组件。

基本模块

nn.Flatten 层 把二位图像打平成一维数组,tensor第一位置代表的是通道数,并不参与打平的运算

input_image = torch.rand(3,224,224)
print(input_image.size())
# torch.Size([3, 224, 224])
flatten = nn.Flatten()
flat_image = flatten(input_image)
print(flat_image.size())
# torch.Size([3, 50176])

nn.Linear

通过权重w和偏差b对输入数据进行线性变化,输出结果。这个最基本的神经网络结构

layer1 = nn.Linear(in_features=224*224, out_features=1024)
hidden1 = layer1(flat_image)
print(hidden1.size()
#torch.Size([3, 1024])

nn.ReLU

非线性激活函数,帮助神经网络引入非线性的特性

print(f"Before ReLU: {hidden1}\n\n")
hidden1 = nn.ReLU()(hidden1)
print(f"After ReLU: {hidden1}")
nn.Sequential

一个有序容器,把各个模块顺序连接起来

nn.Softmax

将(-无穷,+无穷)范围的数值,压缩到(0,1),用于表示模型对每个类别的预测概率

softmax = nn.Softmax(dim=1)
pred_probab = softmax(logits)

整体结构

# 引入对应的模块
import os
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
# 选择模型的训练资源
device = (
    "cuda"
    if torch.cuda.is_available()
    else "mps"
    if torch.backends.mps.is_available()
    else "cpu"
)
print(f"Using {device} device")
# 初始化神经网络实例,并把它迁移到对应的设备上
model = NeuralNetwork(224*224).to(device)
print(model)
# 验证输出
X = torch.rand(1, 224, 224, device=device)
logits = model(X)
pred_probab = nn.Softmax(dim=1)(logits)
y_pred = pred_probab.argmax(1)
print(f"Predicted class: {y_pred}")