跳到主要内容
极客日志极客日志面向AI+效率的开发者社区
首页博客我的书AI学习GitHub 精选镜像AI 生图工具UI配色美学关于
搜索内容 / 工具 / 仓库 / 镜像...⌘K搜索
注册
博客列表
PythonAI算法

Python AI 实战:从线性回归到 MNIST 图像分类

Python AI 实战涵盖环境搭建、线性回归、神经网络及 MNIST 图像分类全流程。基于 PyTorch 框架演示核心代码,解析卷积网络结构与训练循环,并对比 Rust 开发体验差异。内容从基础理论延伸至实际项目落地,适合希望快速掌握深度学习基础的开发者参考。

王初壹发布于 2026/3/28更新于 2026/9/448 浏览

Python AI 实战:从线性回归到 MNIST 图像分类

一、环境搭建与 Hello World

开始之前,先准备好开发环境。PyTorch 是目前最流行的深度学习框架之一,安装时建议直接通过 pip 获取稳定版。

pip install torch torchvision numpy matplotlib

1.1 第一个 AI 程序:线性回归

让我们写一个最简单的模型来理解训练流程。这里用线性回归拟合一条直线,虽然简单,但涵盖了数据生成、模型定义、损失计算和参数更新的核心步骤。

import torch
import torch.nn as nn
import numpy as np
import matplotlib.pyplot as plt

# 生成训练数据:y = 2x + 1 + noise
x = torch.linspace(0, 10, 100).unsqueeze(1)
y = 2 * x + 1 + torch.randn(100, 1) * 0.5

# 定义模型
class LinearModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = nn.Linear(1, 1)

    def forward(self, x):
        return self.linear(x)

# 创建实例
model = LinearModel()

# 损失函数和优化器
criterion = nn.MSELoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

# 训练循环
epochs = 100
for epoch in range(epochs):
    outputs = model(x)
    loss = criterion(outputs, y)
    
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

    if (epoch + 1) % 10 == 0:
        print(f'Epoch [{epoch+1}/{epochs}], Loss: {loss.item():.4f}')

# 可视化结果
with torch.no_grad():
    predicted = model(x)

plt.scatter(x.numpy(), y.numpy(), label='Original data')
plt.plot(x.numpy(), predicted.numpy(), 'r-', label='Fitted line')
plt.legend()
plt.show()
print("AI 模型训练完成")

二、从线性回归到神经网络

线性模型只能处理直线关系,现实中的数据往往是非线性的。这时候就需要引入隐藏层,构建神经网络。

2.1 构建神经网络

我们生成一个二次曲线数据,看看网络能否拟合它。

import torch
import torch.nn as nn
import torch.optim as optim
import matplotlib.pyplot as plt

# 生成非线性数据:y = x^2 + noise
x = torch.linspace(-1, 1, 100).unsqueeze(1)
y = x.pow(2) + 0.2 * torch.randn(100, 1)

# 定义神经网络
class NeuralNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.hidden = nn.Linear(1, 10)  # 隐藏层
        self.output = nn.Linear(10, 1)  # 输出层

    def forward(self, x):
        x = torch.relu(self.hidden(x))  # 激活函数
        x = self.output(x)
        return x

model = NeuralNet()
criterion = nn.MSELoss()
optimizer = optim.SGD(model.parameters(), lr=0.01)

# 训练
epochs = 1000
for epoch in range(epochs):
    outputs = model(x)
    loss = criterion(outputs, y)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

    if (epoch + 1) % 100 == 0:
        print(f'Epoch [{epoch+1}/{epochs}], Loss: {loss.item():.4f}')

# 可视化
with torch.no_grad():
    predicted = model(x)

plt.scatter(x.numpy(), y.numpy(), label='Original data')
plt.plot(x.numpy(), predicted.numpy(), 'r-', label='Neural network prediction')
plt.legend()
plt.show()

2.2 理解工作原理

神经网络之所以强大,在于它能通过多层神经元组合提取特征。简单来说:

  • 输入层:接收原始数据。
  • 隐藏层:负责提取特征,层数越多越能捕捉复杂模式。
  • 输出层:给出最终预测。
  • 激活函数:如 ReLU,引入非线性,让网络能拟合曲线而非直线。

三、图像分类入门:MNIST 手写数字

接下来进入真正的计算机视觉领域,使用经典的 MNIST 数据集进行手写数字识别。

3.1 数据准备

PyTorch 提供了便捷的接口加载数据集,关键是要做好预处理,比如归一化。

import torch
import torchvision
import torchvision.transforms as transforms
import matplotlib.pyplot as plt
import numpy as np

# 数据预处理:转 Tensor 并归一化
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,))
])

# 加载数据集
trainset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=64, shuffle=True)
testset = torchvision.datasets.MNIST(root='./data', train=False, download=True, transform=transform)
testloader = torch.utils.data.DataLoader(testset, batch_size=64, shuffle=False)

# 显示一张图片
fig = plt.figure(figsize=(8, 8))
dataiter = iter(trainloader)
images, labels = next(dataiter)

def imshow(img):
    img = img / 2 + 0.5
    npimg = img.numpy()
    fig.add_subplot(111, aspect="equal")
    plt.imshow(np.transpose(npimg, (1, 2, 0)))
    plt.axis('off')

imshow(torchvision.utils.make_grid(images))
print('标签:', ' '.join(f'{labels[j]}' for j in range(4)))

3.2 构建卷积神经网络 (CNN)

图像数据是二维的,卷积层比全连接层更适合提取空间特征。

import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        # 卷积层:1 通道灰度图 -> 32 通道 -> 64 通道
        self.conv1 = nn.Conv2d(1, 32, 3, 1)
        self.conv2 = nn.Conv2d(32, 64, 3, 1)
        # 池化层
        self.pool = nn.MaxPool2d(2, 2)
        # 全连接层
        self.fc1 = nn.Linear(64 * 12 * 12, 128)
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 64 * 12 * 12)
        x = F.relu(self.fc1(x))
        x = self.fc2(x)
        return x

net = Net()
print(net)

3.3 训练与测试

训练循环是重复的过程:前向传播 -> 计算损失 -> 反向传播 -> 更新参数。

import torch.optim as optim

criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)

# 训练
epochs = 5
for epoch in range(epochs):
    running_loss = 0.0
    for i, data in enumerate(trainloader, 0):
        inputs, labels = data
        
        optimizer.zero_grad()
        outputs = net(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        running_loss += loss.item()
        if i % 100 == 99:
            print(f'[{epoch + 1}, {i + 1}] loss: {running_loss / 100:.3f}')
            running_loss = 0.0
print('训练完成')

# 测试
correct = 0
total = 0
with torch.no_grad():
    for data in testloader:
        images, labels = data
        outputs = net(images)
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

print(f'测试准确率:{100 * correct / total:.2f}%')

# 查看预测结果
dataiter = iter(testloader)
images, labels = next(dataiter)
outputs = net(images)
_, predicted = torch.max(outputs, 1)

print('真实标签:', ' '.join(f'{labels[j]}' for j in range(4)))
print('预测标签:', ' '.join(f'{predicted[j]}' for j in range(4)))

四、从 Rust 开发者角度的思考

如果你习惯 Rust,转战 Python AI 会有些独特的感受。

4.1 体验对比

  • 开发效率:Python 在 AI 领域确实快,库多且封装好,不用操心内存管理。
  • 生态系统:PyTorch 生态成熟,社区活跃,遇到问题容易找到解决方案。
  • 性能:Python 本身慢,但底层算子(如 CUDA)是用 C++/CUDA 写的,实际推理速度并不差。
  • 类型系统:动态类型需要适应,建议配合 mypy 或类型注解保持代码质量。

4.2 学习建议

  • 利用系统思维:Rust 的系统级经验有助于理解显存分配和计算图原理。
  • 注重代码质量:保持 Rust 的代码风格,变量命名清晰,注释到位。
  • 实践项目:理论看多了容易忘,跑通几个 Demo 印象更深。
  • 跨语言结合:未来可以尝试用 Rust 编写高性能算子,Python 做胶水逻辑。

五、总结

从简单的线性回归到复杂的 CNN 图像分类,我们走完了 PyTorch 的基础路径。Python 的动态特性降低了入门门槛,丰富的库让快速原型开发成为可能。对于追求极致性能的底层场景,Rust 依然是王者;但在 AI 应用层,Python 目前的统治地位短期内难以撼动。掌握这两者的互补优势,会让你的技术栈更完整。

目录

  1. Python AI 实战:从线性回归到 MNIST 图像分类
  2. 一、环境搭建与 Hello World
  3. 1.1 第一个 AI 程序:线性回归
  4. 生成训练数据:y = 2x + 1 + noise
  5. 定义模型
  6. 创建实例
  7. 损失函数和优化器
  8. 训练循环
  9. 可视化结果
  10. 二、从线性回归到神经网络
  11. 2.1 构建神经网络
  12. 生成非线性数据:y = x^2 + noise
  13. 定义神经网络
  14. 训练
  15. 可视化
  16. 2.2 理解工作原理
  17. 三、图像分类入门:MNIST 手写数字
  18. 3.1 数据准备
  19. 数据预处理:转 Tensor 并归一化
  20. 加载数据集
  21. 显示一张图片
  22. 3.2 构建卷积神经网络 (CNN)
  23. 3.3 训练与测试
  24. 训练
  25. 测试
  26. 查看预测结果
  27. 四、从 Rust 开发者角度的思考
  28. 4.1 体验对比
  29. 4.2 学习建议
  30. 五、总结

更多推荐文章

查看全部
  • Spring Web MVC 入门:从概念到实践
  • MS-S1 MAX、AI MAX 395 在 Ubuntu 24 利用 Vulkan 版 llama.cpp 运行 gpt-oss 120b
  • AI 鉴伪检测技术:人脸视频与 AIGC 图像伪造识别方案
  • 基于 Cogito-v1-preview-llama-3B 的汽车电子 ECU 诊断逻辑建模实践
  • Spring Boot 零基础入门:快速构建 Java Web 应用
  • AI 时代的软件工程:使用 OpenSpec 驱动自动化开发
  • C++ 哈希表原理、冲突解决及性能分析
  • AI 大模型入门与实战指南:从基础原理到应用部署
  • 2022 年大厂 Android 面试题汇总与解析
  • 基于 SpringBoot+Vue 的校园招聘系统设计与实现
  • 基于 Python 的商品销售数据分析与可视化实战
  • Python AI 实战:从线性回归到 MNIST 图像分类
  • Python 核心语法速查:数据类型与基础操作
  • MaaS 平台与阿里 QWQ 技术:AI 调参实战指南
  • Python 为何成为 AI 开发的首选语言?
  • 浏览器 F5 刷新机制深度解析
  • PHP 使用 Pdo_kdb 驱动连接 Kingbase 数据库实战指南
  • ComfyUI v0.11.1 发布:新增开发者节点、API 强化与 Python 3.14 兼容
  • 从推荐算法转行大模型推理工程:行业趋势与技术路径
  • 网络安全行业主流证书选择指南

相关免费在线工具

  • 加密/解密文本

    使用加密算法(如AES、TripleDES、Rabbit或RC4)加密和解密文本明文。 在线工具,加密/解密文本在线工具,online

  • RSA密钥对生成器

    生成新的随机RSA私钥和公钥pem证书。 在线工具,RSA密钥对生成器在线工具,online

  • Mermaid 预览与可视化编辑

    基于 Mermaid.js 实时预览流程图、时序图等图表,支持源码编辑与即时渲染。 在线工具,Mermaid 预览与可视化编辑在线工具,online

  • 随机西班牙地址生成器

    随机生成西班牙地址(支持马德里、加泰罗尼亚、安达卢西亚、瓦伦西亚筛选),支持数量快捷选择、显示全部与下载。 在线工具,随机西班牙地址生成器在线工具,online

  • Gemini 图片去水印

    基于开源反向 Alpha 混合算法去除 Gemini/Nano Banana 图片水印,支持批量处理与下载。 在线工具,Gemini 图片去水印在线工具,online

  • curl 转代码

    解析常见 curl 参数并生成 fetch、axios、PHP curl 或 Python requests 示例代码。 在线工具,curl 转代码在线工具,online