AI基本结构11-rnn实现简单nlp
循环神经网络数据准备和之前差不多不赘述了# 一些超参数 learning_rate 1e-3 # 如果有GPU该脚本将使用GPU进行计算 device cuda if torch.cuda.is_available() else cpu raw_datasets load_dataset(code_search_net, python) datasets raw_datasets[train].filter(lambda x: apache/spark in x[repository_name])其中 lambda x: apache/spark in x[repository_name] 定义了一个匿名函数Lambda 函数它接收数据集中的每一条样本 x 作为输入判断其 repository_name 字段中是否包含字符串 apache/spark若条件成立则返回 True 保留该样本否则返回 False 将其过滤掉。Lambda 函数本质上是一个简洁的匿名函数这里的写法等价于def func(x): return apache/spark in x[repository_name]由于该函数只使用一次因此采用 lambda 写法更加简洁高效也便于直接作为 filter() 的筛选条件。编码器更新#运行不定长所以可以删掉begin字符 class CharTokenizer: def __init__(self, data, end_ind0): # data: list[str] # 得到所有的字符 chars sorted(list(set(.join(data)))) #self.char2ind {s: i 2 for i, s in enumerate(chars)} self.char2ind {s: i 1 for i, s in enumerate(chars)} #self.char2ind[|b|] begin_ind self.char2ind[|e|] end_ind self.ind2char {v: k for k, v in self.char2ind.items()} #self.begin_ind begin_ind self.end_ind end_ind def encode(self, x): # x: str return [self.char2ind[i] for i in x] def decode(self, x): # x: int or list[x] if isinstance(x, int): return self.ind2char[x] return [self.ind2char[i] for i in x] #测试 tokenizer CharTokenizer(datasets[whole_func_string]) test_str def f(x): re tokenizer.encode(test_str) print(re) .join(tokenizer.decode(range(len(tokenizer.char2ind))))定义了一个基于字符级别Character-Level的 CharTokenizer用于建立字符与编号之间的映射实现文本的编码encode和解码decode。与注释中的原始版本相比最大的变化是删除了起始标记Begin Token这是因为当前模型采用不定长运行方式不再需要在每个样本前添加固定数量的起始标记因此可以简化词表结构仅利用结束标记表示文本终止从而减少额外的输入符号使数据预处理和文本生成过程更加简洁。循环神经元定义class RnnCell(nn.Module): def __init__(self,input_size,hidden_size): super().__init__() self.input_size input_size self.hidden_size hidden_size self.inTohinn.Linear(input_sizehidden_size,hidden_size) def forward(self,input,hiddenNone): if hidden is None: hidden self.init_hidden(input.device) combined torch.concat((input,hidden),dim-1) hidden F.relu(self.inTohi(combined)) return hidden def init_hidden(self,device): return torch.zeros((1,self.hidden_size),devicedevice) # 测试 r_model RnnCell(2, 3) data torch.randn(4, 1, 2) hidden None for i in range(data.shape[0]): hidden r_model(data[i], hidden) print(hidden)1RnnCell 类的作用RnnCell 实现了一个最基本的循环神经网络RNN单元用于处理序列数据。与多层感知器MLP不同RNN 在每个时间步都会接收当前输入和上一时刻的隐藏状态Hidden State从而能够保留历史信息实现对序列上下文的建模。2初始化网络结构在 __init__() 中定义了输入维度 input_size 和隐藏层维度 hidden_size并创建了一个全连接层 inTohi。由于每次输入都需要与上一时刻的隐藏状态拼接因此该线性层的输入维度为 input_size hidden_size输出维度为 hidden_size用于计算新的隐藏状态。3前向传播过程forward() 函数首先判断是否存在上一时刻的隐藏状态若没有则调用 init_hidden() 初始化为全零向量。随后将当前输入 input 与上一时刻隐藏状态 hidden 在最后一个维度进行拼接形成包含当前信息和历史信息的特征向量再经过全连接层和 ReLU 激活函数得到当前时刻新的隐藏状态并作为输出返回。循环神经网络定义class CharRNN(nn.Module): def __init__(self, vs): super().__init__() self.emb nn.Embedding(vs, 30) self.rnn RnnCell(30, 50) self.lm nn.Linear(50, vs) def forward(self, x, hiddenNone): # x: (1) # hidden: (1, 50) embeddings self.emb(x) # (1, 30) hidden self.rnn(embeddings, hidden) # (1, 50) out self.lm(hidden) # (1, vs) return out, hidden简单网络不多说啦生成函数更新torch.no_grad() def generate(model, idx, tokenizer, max_new_tokens300): # idx: (1) out idx.tolist() hidden None model.eval() for _ in range(max_new_tokens): logits, hidden model(idx, hidden) probs F.softmax(logits, dim-1) # (1, 98) # 随机生成文本 ix torch.multinomial(probs, num_samples1) # (1, 1) ## 更新背景 #context torch.concat((context[:, 1:], ix), dim-1) out.append(ix.item()) idx ix.squeeze(0) if out[-1] tokenizer.end_ind: break model.train() return out #测试 inputs torch.tensor(tokenizer.encode(d), devicedevice) print(.join(tokenizer.decode(generate(c_model, inputs, tokenizer)))) def process(text, tokenizer): # text: str enc tokenizer.encode(text) inputs enc labels enc[1:] [tokenizer.end_ind] return torch.tensor(inputs, devicedevice), torch.tensor(labels, devicedevice) #测试 print(process(test_str, tokenizer))1generate 函数作用generate 用于利用训练好的循环神经网络RNN生成文本。函数以初始字符 idx 作为输入模型根据当前输入和历史隐藏状态逐步预测下一个字符并不断将预测结果作为下一时刻的输入实现字符级文本的连续生成。2随机采样与更新输入模型输出经过 Softmax 转换为概率分布后利用 torch.multinomial() 依概率随机采样得到下一个字符编号 ix并将其加入输出序列。同时将 ix 作为下一时刻模型的输入idx ix.squeeze(0)形成“预测一个字符再将其作为下一次输入”的自回归生成过程。代码中被注释掉的 context 更新语句是 MLP 固定窗口模型的实现方式而 RNN 利用隐藏状态记录上下文因此不再需要滑动窗口更新背景信息。3process 函数作用新的 process 函数用于构造 RNN 的训练数据。首先将输入文本编码为字符编号序列 enc然后直接将整个序列作为模型输入 inputs而标签 labels 则是输入序列整体向后移动一位并在末尾补充结束标记 end_ind。这种构造方式使模型能够学习“根据当前字符预测下一个字符”符合 RNN 自回归语言模型的训练目标。训练与测试def trainRnn(model,optimizer,epochs2): lossi [] for e in range(epochs): for data in datasets: inputs, labels process(data[whole_func_string], tokenizer) hidden None _loss 0.0 lens len(inputs) for i in range(lens): logits, hidden model(inputs[i].unsqueeze(0), hidden) _loss F.cross_entropy(logits, labels[i].unsqueeze(0)) / lens lossi.append(_loss.item()) optimizer.zero_grad() _loss.backward() optimizer.step() print(_loss) return lossi #测试 epochs 1 optimizer optim.Adam(c_model.parameters(), lrlr) losstrainRnn(c_model,optimizer,epochs) plt.plot(loss) plt.show() inputs torch.tensor(tokenizer.encode(d), devicedevice) print(.join(tokenizer.decode(generate(c_model, inputs, tokenizer))))正常训练的老三样损失优化循环

相关新闻

AI工作流是什么?为什么比单个工具更重要

AI工作流是什么?为什么比单个工具更重要

在当前数字化办公普及的环境下,绝大多数学生与职场人员的AI使用方式,仍停留在“单次提问、单次解决问题”的浅层阶段。日常工作中遇到写文案、整理表格、简单改错、内容润色等任务,多数人都是临时输入指令、单次获取结果,用完即结…

2026/7/27 6:57:20阅读更多 →
共享屏幕怎么操作 异地共享屏幕的方法

共享屏幕怎么操作 异地共享屏幕的方法

异地对接工作、分隔两地相伴观影时,很多人都会疑惑共享屏幕怎么操作,常规共享软件存在时长限制、画面模糊等短板,普通投屏工具又只局限局域网使用,很难适配远距离场景。共享屏幕怎么操作才能兼顾流畅度与隐私保障?推荐…

2026/7/27 6:57:20阅读更多 →
华为CANN架构下Mul与Div算子的深度优化实践

华为CANN架构下Mul与Div算子的深度优化实践

1. CANN架构与算子基础解析华为CANN(Compute Architecture for Neural Networks)作为全栈AI计算架构,其核心设计理念是通过硬件-软件协同优化来释放Ascend芯片的算力潜能。在深度学习领域,基础算子如Mul(乘法&#xff…

2026/7/27 6:57:20阅读更多 →
HsMod炉石传说插件终极指南:解锁32倍速游戏和200+皮肤定制

HsMod炉石传说插件终极指南:解锁32倍速游戏和200+皮肤定制

HsMod炉石传说插件终极指南:解锁32倍速游戏和200皮肤定制 【免费下载链接】HsMod Hearthstone Modification Based on BepInEx 项目地址: https://gitcode.com/GitHub_Trending/hs/HsMod HsMod是一款基于BepInEx框架开发的炉石传说游戏增强插件,为…

2026/7/27 8:23:26阅读更多 →
Python after-class 包完全指南:功能、安装、语法与实战案例

Python after-class 包完全指南:功能、安装、语法与实战案例

1. after-class 包概述 after-class 是一个轻量级的 Python 工具包,主要用于自动化课后作业的提交、批改与成绩管理。它通常与在线教学平台(如 ClassIn、钉钉、腾讯课堂等)配合使用,帮助教师和学生简化课后流程。该包的核心功能包括:自动提交作业、批量下载作业、成绩统计…

2026/7/27 8:23:26阅读更多 →
Python实战:医疗金融教育三大场景下的差分隐私落地指南

Python实战:医疗金融教育三大场景下的差分隐私落地指南

1. 项目概述:当数据安全成为业务基石在医疗、金融、教育这三大领域摸爬滚打多年,我深刻体会到数据价值的另一面是如履薄冰的安全责任。一份匿名的诊疗记录,通过与其他公开数据的关联,可能被重新识别出患者身份;一组看似…

2026/7/27 8:23:26阅读更多 →
Python afs-file-validator 包:功能详解、安装使用与实战案例

Python afs-file-validator 包:功能详解、安装使用与实战案例

1. 引言 在文件上传、数据导入和批量处理等场景中,文件格式验证是一个常见且重要的需求。Python 的 afs-file-validator 包提供了一套轻量、可扩展的文件验证方案,支持多种文件类型的格式校验、MIME 类型检测、大小限制和内容完整性检查。本文将详细介绍该包的安装、核心功能…

2026/7/27 8:23:26阅读更多 →
线性回归原理与实现:从数学基础到工程实践

线性回归原理与实现:从数学基础到工程实践

1. 线性回归的本质与最小训练闭环 线性回归是机器学习领域最基础也最重要的算法之一,它构建了从数据到预测的桥梁。这个看似简单的模型背后,蕴含着监督学习的核心范式——最小训练闭环。所谓最小训练闭环,指的是一个完整的机器学习流程中最精…

2026/7/27 8:23:26阅读更多 →
TB-RK3399Pro移植ubuntu 26.04及烧录

TB-RK3399Pro移植ubuntu 26.04及烧录

TB-RK3399Pro移植ubuntu 26.04及烧录一、Ubuntu 26.04根文件系统制作1.1 安装必要工具包1.2 获取 Ubuntu Base RootFS1.3 文件系统挂载1.4 换源:1.5 更新系统并安装基础软件1.6 配置中文支持(可选)1.7 设置时区1.8 添加账户1.9 串口 getty&am…

2026/7/27 8:21:26阅读更多 →
覆盖国产 + 海外 + 开源模型,OpenClaw 2.7.9 Windows/Mac 双端部署详解

覆盖国产 + 海外 + 开源模型,OpenClaw 2.7.9 Windows/Mac 双端部署详解

🔹 工具基础介绍 OpenClaw 是开源生态中一款实用性较强的本地智能工具,凭借本地离线运行、可视化图形操作和任务自动化三大核心特性,赢得了众多用户的青睐。与普通在线对话AI工具不同,它属于能够直接操控本机软硬件的智能数字员工…

2026/7/27 1:14:34阅读更多 →
伺服阀焊完微漏毁整机?精密激光焊接三关锁住高压

伺服阀焊完微漏毁整机?精密激光焊接三关锁住高压

所谓液压伺服阀体的精密激光焊接,是用激光束对阀座壳体(通常为不锈钢或铝合金)进行密封焊接,使阀体在21-35MPa的高压液压油或压缩气体中长期运行而不发生介质泄漏。液压伺服阀是高端液压系统的"大脑"。从航空航天飞行控…

2026/7/27 1:14:52阅读更多 →
D2DX:三步实现《暗黑破坏神2》高清宽屏体验的终极指南

D2DX:三步实现《暗黑破坏神2》高清宽屏体验的终极指南

D2DX:三步实现《暗黑破坏神2》高清宽屏体验的终极指南 【免费下载链接】d2dx D2DX is a complete solution to make Diablo II run well on modern PCs, with high fps and better resolutions. 项目地址: https://gitcode.com/gh_mirrors/d2/d2dx 你是否还在…

2026/7/27 1:14:56阅读更多 →
SPI实战指南:从时钟模式到寄存器配置,解决嵌入式通信难题

SPI实战指南:从时钟模式到寄存器配置,解决嵌入式通信难题

1. 项目概述:从寄存器手册到实战指南 如果你手头有一份类似德州仪器(TI)TMS320x240xA系列DSP的SPI模块技术手册,看着里面密密麻麻的寄存器位定义、时序图和公式,是不是感觉头大?这份资料虽然权威&#xff0…

2026/7/27 0:00:24阅读更多 →
【JAVA毕设源码分享】基于springboot的水果购物管理系统的设计与实现(程序+文档+代码讲解+一条龙定制)

【JAVA毕设源码分享】基于springboot的水果购物管理系统的设计与实现(程序+文档+代码讲解+一条龙定制)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于Java、小程序技术领域和毕业项目实战 ✌️技术范围:&am…

2026/7/27 0:00:24阅读更多 →
2007-2023年各市区县生态文明建设示范区DID

2007-2023年各市区县生态文明建设示范区DID

数据简介 自改革开放以来,我国依赖高投入、高资源消耗和高污染等传统发展模式实现了经济短期内的快速增长, 然而这也导致了严重的生态环境危机。因此,国家有力于推动企业高质量经济发展,协同生态保护的方针,从而从201…

2026/7/27 0:00:24阅读更多 →
YOLOv8推理性能优化:从1.2FPS到35FPS的全链路加速实践

YOLOv8推理性能优化:从1.2FPS到35FPS的全链路加速实践

如果你在部署 YOLOv8 时,发现推理速度只有可怜的 1-2 FPS,而别人的演示视频却能跑到 30 FPS 以上,那么问题很可能不在模型本身,而在于你的整个处理链路。很多开发者拿到一个训练好的 YOLOv8 模型后,会直接使用官方示例…

2026/7/25 23:03:25阅读更多 →
Coze与Dify对比指南:低代码AI应用开发从入门到实战

Coze与Dify对比指南:低代码AI应用开发从入门到实战

1. 从零到一:为什么你需要了解 Coze 和 Dify?如果你对 AI 应用开发感兴趣,但一看到“大模型”、“智能体”、“工作流”这些词就头疼,觉得门槛太高,那这篇文章就是为你准备的。很多开发者,包括我自己&#…

2026/7/26 19:05:21阅读更多 →
AI生图工具怎么选?2026年6月版实测对比

AI生图工具怎么选?2026年6月版实测对比

做自媒体的朋友应该都有体会:配图一直是个让人头疼的问题。2026年,AI生图工具已经非常成熟了,但工具太多反而不知道怎么选。以下是截至2026年6月我对主流AI生图工具的实测对比。Midjourney V8.1:速度之王2026年6月11日&#xff0c…

2026/7/26 19:05:21阅读更多 →