当前位置: 首页 > news >正文

Python AI入门:从Hello World到图像分类

Python AI入门:从Hello World到图像分类

一、Python AI的Hello World

1.1 环境搭建

首先,我们需要搭建Python AI的开发环境:

# 安装PyTorchpipinstalltorch torchvision# 安装其他依赖pipinstallnumpy matplotlib

1.2 第一个AI程序

让我们来编写一个最简单的AI程序 - 线性回归:

importtorchimporttorch.nnasnnimportnumpyasnpimportmatplotlib.pyplotasplt# 生成训练数据x=torch.linspace(0,10,100).unsqueeze(1)y=2*x+1+torch.randn(100,1)*0.5# 定义模型classLinearModel(nn.Module):def__init__(self):super(LinearModel,self).__init__()self.linear=nn.Linear(1,1)defforward(self,x):returnself.linear(x)# 创建模型实例model=LinearModel()# 定义损失函数和优化器criterion=nn.MSELoss()optimizer=torch.optim.SGD(model.parameters(),lr=0.01)# 训练模型epochs=100forepochinrange(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}')# 测试模型withtorch.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("Hello World! AI模型训练完成")

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

2.1 神经网络基础

线性回归是最简单的AI模型,而神经网络则是更复杂的模型。让我们来构建一个简单的神经网络:

importtorchimporttorch.nnasnnimporttorch.optimasoptim# 生成非线性数据x=torch.linspace(-1,1,100).unsqueeze(1)y=x.pow(2)+0.2*torch.randn(100,1)# 定义神经网络模型classNeuralNet(nn.Module):def__init__(self):super(NeuralNet,self).__init__()self.hidden=nn.Linear(1,10)self.output=nn.Linear(10,1)defforward(self,x):x=torch.relu(self.hidden(x))x=self.output(x)returnx# 创建模型实例model=NeuralNet()# 定义损失函数和优化器criterion=nn.MSELoss()optimizer=optim.SGD(model.parameters(),lr=0.01)# 训练模型epochs=1000forepochinrange(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}')# 测试模型withtorch.no_grad():predicted=model(x)# 可视化结果importmatplotlib.pyplotasplt 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 理解神经网络的工作原理

神经网络的基本原理是通过多层神经元的组合,学习数据中的复杂模式:

  1. 输入层:接收原始数据
  2. 隐藏层:提取数据特征
  3. 输出层:产生预测结果
  4. 激活函数:引入非线性,使网络能够学习复杂模式

三、图像分类入门

3.1 数据准备

我们将使用MNIST数据集进行图像分类:

importtorchimporttorchvisionimporttorchvision.transformsastransforms# 数据预处理transform=transforms.Compose([transforms.ToTensor(),transforms.Normalize((0.5,),(0.5,))])# 加载MNIST数据集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)# 查看数据importmatplotlib.pyplotaspltimportnumpyasnp# 函数:显示图像defimshow(img):img=img/2+0.5# 反归一化npimg=img.numpy()plt.imshow(np.transpose(npimg,(1,2,0)))plt.show()# 获取一批训练数据dataiter=iter(trainloader)images,labels=next(dataiter)# 显示图像imshow(torchvision.utils.make_grid(images))print('标签:',' '.join(f'{labels[j]}'forjinrange(4)))

3.2 构建图像分类模型

现在我们来构建一个用于图像分类的卷积神经网络:

importtorch.nnasnnimporttorch.nn.functionalasFclassNet(nn.Module):def__init__(self):super(Net,self).__init__()# 卷积层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)defforward(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)returnx# 创建模型实例net=Net()print(net)

3.3 训练图像分类模型

importtorch.optimasoptim# 定义损失函数和优化器criterion=nn.CrossEntropyLoss()optimizer=optim.SGD(net.parameters(),lr=0.001,momentum=0.9)# 训练模型epochs=5forepochinrange(epochs):running_loss=0.0fori,datainenumerate(trainloader,0):# 获取输入inputs,labels=data# 清零梯度optimizer.zero_grad()# 前向传播outputs=net(inputs)# 计算损失loss=criterion(outputs,labels)# 反向传播loss.backward()# 更新参数optimizer.step()# 统计损失running_loss+=loss.item()ifi%100==99:print(f'[{epoch+1},{i+1}] loss:{running_loss/100:.3f}')running_loss=0.0print('训练完成')

3.4 测试模型

# 测试模型correct=0total=0withtorch.no_grad():fordataintestloader: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)# 显示图像imshow(torchvision.utils.make_grid(images))print('真实标签:',' '.join(f'{labels[j]}'forjinrange(4)))# 预测outputs=net(images)_,predicted=torch.max(outputs,1)print('预测标签:',' '.join(f'{predicted[j]}'forjinrange(4)))

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

4.1 与Rust的对比

作为一个Rust开发者,学习Python AI有以下感受:

  • 开发效率:Python的开发效率比Rust高,尤其是在AI开发中
  • 生态系统:Python的AI生态系统非常丰富,有大量成熟的库
  • 性能:Python的性能虽然不如Rust,但在AI开发中,PyTorch等库已经做了很多优化
  • 类型系统:Python的动态类型与Rust的静态类型有很大不同,需要适应

4.2 学习建议

对于Rust开发者学习Python AI,我有以下建议:

  • 利用系统思维:Rust的系统级编程经验有助于理解AI模型的底层实现
  • 注重代码质量:保持Rust的代码风格,写出清晰、可维护的Python代码
  • 实践项目:通过实际项目巩固学习成果
  • 跨语言学习:将Rust和Python结合起来,发挥各自的优势

五、总结

通过从Hello World到图像分类的学习,我已经初步掌握了Python AI的基本概念和使用方法。作为一个Rust开发者,我发现Python AI的学习过程既有挑战也有机遇。

挑战在于Python的动态类型和内存管理与Rust有很大不同,需要适应新的思维方式。机遇在于Python的AI生态系统非常丰富,开发效率高,能够快速实现AI模型。

http://www.cnnetsun.cn/news/1410602.html

相关文章:

  • TensorFlow-v2.15环境搭建:无需复杂配置,镜像开箱即用,即刻开始编码
  • Qwen3-ASR-1.7B效果对比:在Mandarin-English Switching Test Set上准确率+31.6%
  • 软萌拆拆屋惊艳案例:婚纱复杂结构拆解图(蕾丝/珠片/衬裙分层)
  • 5分钟解锁付费墙:Bypass Paywalls Clean终极免费阅读指南
  • SSD1308 OLED驱动库:I²C接口128×64单色屏嵌入式实战指南
  • 隐私安全!本地离线部署Qwen3-4B写作大师,数据不出门
  • SEO_详解SEO核心关键词研究与布局策略
  • Win11Debloat开源工具:Windows系统优化实用指南
  • ModbusTool深度技术解析:工业协议测试平台架构解密
  • 避坑指南:antd表头提示文字不生效的5个常见原因及解决方案
  • 效率直接起飞!风靡全网的AI论文软件 —— 千笔·专业学术智能体
  • 计算机毕业设计springboot香格里拉幼儿园捐赠物资分配一体化管理系统 基于SpringBoot的迪庆藏区学前教育机构爱心物资流转智能平台 SpringBoot框架下高原地区幼儿园公益捐赠资源协同
  • 突破视觉局限:多光谱目标检测如何重塑AI感知能力
  • 造相-Z-Image-Turbo 作品生成与分享平台构建:全栈技术实践(Vue+ .NET)
  • IMU传感器在无人机飞控中的实战应用:从加速度计校准到陀螺仪数据融合
  • 达梦数据库实战:如何高效管理用户权限与表空间(附常见问题解决方案)
  • 3秒出图!Nunchaku FLUX.1-dev量化版,16GB显卡也能玩转AI绘画
  • MiniCPM-o-4.5-nvidia-FlagOS开源可部署:Apache 2.0许可下二次开发与私有化定制指南
  • 3步打造ESP32物联网环境监测系统:嵌入式开发者的终极指南
  • MS17-010 永恒之蓝漏洞渗透实验|Kali+Windows实操全步骤
  • Blender 3MF插件深度解析:解锁3D打印工作流的5大核心能力
  • 云原生时代必知:Overlay网络在Kubernetes中的5种实战用法(附配置示例)
  • 如何免费扩展显示器:开源虚拟显示器完整教程
  • 嵌入式系统中高效安全的memcpy实现原理与优化
  • Arducam OV5642嵌入式摄像头驱动开发指南
  • 微信聊天记录安全备份与智能应用:一站式解决方案
  • VSCode插件包实战:从零发布一个属于自己的“Java增强包”到官方市场
  • 2026冲刺用!全场景通用降AI率网站 —— 千笔·降AI率助手
  • Qwen2.5-VL-7B-Instruct快速入门:基于Streamlit的可视化界面,图文交互超简单
  • 【笔试真题】- 得物-2026.03.21