新闻详情

新闻详情

首页 / 资讯中心 / 详情

循环神经网络详解(RNN)

发布时间:2026/9/29 21:15:40来源:尧图网络
循环神经网络详解(RNN)
循环神经网络是一类专门处理序列数据的神经网络循环神经网络的每个神经元具有两个输入一个当前时间步的输入和一个前一时刻的隐藏状态通过这两个输入产生当前时刻的隐藏状态。RNN的核心优势在于可以处理变长序列参数共享和具备记忆能力。1、RNN参数详解import torch from torch import nn torch.nn.RNN(input_size,hidden_size,num_layers,nonlinerity,bias,batch_first,dropout,didirectional)1、input_size每个时间步的输入向量维度即单个输入大小2、hidden_size隐藏状态维度即RNN中神经元的个数3、num_layersRNN堆叠层数4、nonlineraity计算隐藏状态的激活函数tanh或relutiona5、bias是否使用偏置6、batch_first输入是否为batch,time_step,input_size)7、dropout是否使用丢弃学习8、bidirectional是否使用双向RNN2、示例代码import torch from torch import nn from torchvision import transforms,datasets import torch.utils.data as Data torch.cuda.empty_cache() devicetorch.device(cuda:0 if torch.cuda.is_available() else cpu) BATCH_SIZE50 TIME_SIZE50 INPUT_SIZE50 transformtransforms.Compose([ transforms.Resize((50,50)), transforms.ToTensor(), transforms.Normalize((0.1307,),( 0.3081,)) ]) train_datadatasets.MNIST( rootD:/mypython/MNISTdataset, trainTrue, downloadTrue, transformtransform ) train_loaderData.DataLoader(datasettrain_data,batch_sizeBATCH_SIZE,shuffleTrue) traintest,labeltestnext(iter(train_loader)) #print(traintest.shape) #print(labeltest.shape) test_datadatasets.MNIST( rootD:/mypython/MNISTdataset, trainFalse, transformtransform ) test_loaderData.DataLoader(datasettest_data,batch_size50,shuffleFalse) test_x,test_ynext(iter(test_loader)) #print(test_x.size()) #print(test_y.size()) class RNN(nn.Module): def __init__(self): super(RNN,self).__init__() self.rnnnn.GRU( input_sizeINPUT_SIZE, hidden_size50, num_layers1, batch_firstTrue, bidirectionalTrue ) self.outnn.Linear(5000,10) def forward(self,x): r_out,(h_n,h_c)self.rnn(x,None) r_outr_out.reshape(r_out.size(0),-1) outputself.out(r_out) return output modelRNN() modelmodel.to(devicedevice) optimizertorch.optim.Adam(model.parameters(),lr0.01) loss_funcnn.CrossEntropyLoss() for step,(x,y) in enumerate(train_loader): xx.squeeze(1) b_xx.to(devicedevice) b_yy.to(devicedevice) outputmodel(b_x) lossloss_func(output,b_y) optimizer.zero_grad() loss.backward() optimizer.step() if step%1000: test_xtest_x.squeeze(1) t_xtest_x.to(devicedevice) t_ytest_y.to(devicedevice) test_outputmodel(t_x) pred_ytorch.max(test_output,1)[1].data.squeeze() accuracy(pred_yt_y).sum().item()/float(test_y.size(0)) print(train loss%.4f %loss.data,|test accuracy:%.2f %accuracy)
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

OpenClaw实战:Dashboard-v2 军团化管理配置全解析,Agents 批量接入 TaoToken 实战 2026/9/29 21:15:26

OpenClaw实战:Dashboard-v2 军团化管理配置全解析,Agents 批量接入 TaoToken 实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
TensorFlow+Flask失物招领平台:从模型训练到本地部署 2026/9/29 21:15:25

TensorFlow+Flask失物招领平台:从模型训练到本地部署

最近帮学校学生会做了个失物招领平台,核心功能是“发一条丢东西的帖子,系统自动帮你从一堆招领信息里找出最像的那几条”,整个过程走了一遍TensorFlow网页端开发,从模型训练到Flask部署再到本地跑通。这篇东西不是教科书&#xff…

阅读更多 →
I2C通信故障排查全攻略:从万用表到逻辑分析仪的完整链路 2026/9/29 21:15:25

I2C通信故障排查全攻略:从万用表到逻辑分析仪的完整链路

1. 为什么I2C排查值得单独写一篇I2C这玩意儿,说简单是真简单,两根线一挂,上拉电阻一焊,代码里调个库函数就能读写。但说难也是真难,多少人卡在“设备没反应”这四个字上,一卡就是一整天。我见过太多人一上来…

阅读更多 →
芯片烧录自研还是外包?从成本、良率到风险的全套决策指南 2026/9/29 21:15:25

芯片烧录自研还是外包?从成本、良率到风险的全套决策指南

做硬件这行,几乎没有人能绕开芯片烧录。小到一颗MCU,大到Flash存储芯片,出厂前都得把固件写进去;这道工序看起来只是“点一下烧录”,真正落地的时候却总是绕回到同一个问题:到底是自己买设备来烧&#xff0…

阅读更多 →
谷歌杀疯了!Gemini 3.8 Flash 突袭发布:代码打平 Opus 5,成本仅需 1/5,TaoToken 统一 Key 接入实测 2026/9/29 21:15:25

谷歌杀疯了!Gemini 3.8 Flash 突袭发布:代码打平 Opus 5,成本仅需 1/5,TaoToken 统一 Key 接入实测

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
webpack方式破解h5st签名:Node.js补环境完整代码链路 2026/9/29 21:15:18

webpack方式破解h5st签名:Node.js补环境完整代码链路

简介:面向爬虫与前端逆向学习者的实战代码包,聚焦某东平台基于webpack方式打包的H5ST签名算法,重点解决采集过程中加密参数生成与校验难题。压缩包共2个文件,包含1个Python脚本与1个JavaScript文件,其中JS文件对应webp…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉