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:13 阅读更多 →
共享屏幕怎么操作 异地共享屏幕的方法

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

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

2026/7/27 6:57:13 阅读更多 →
高速PCB设计实战:从传输线到电源平面,TI KeyStone II布线指南

高速PCB设计实战:从传输线到电源平面,TI KeyStone II布线指南

1. 项目概述:高速PCB设计的核心战场在处理器主频动辄突破GHz、数据速率向数十Gbps迈进的今天,硬件工程师面临的挑战早已超越了“连通即可”的初级阶段。信号在PCB走线上不再是理想的“瞬时”到达,而是以电磁波的形式,在由导体和介…

2026/7/27 6:56:13 阅读更多 →

最新新闻

2026 年程序员找工作,为什么 AI 写出的代码反而成了简历的“硬伤”?

2026 年程序员找工作,为什么 AI 写出的代码反而成了简历的“硬伤”?

聊《一份看似完整的程序员就业方案,为什么投递时没效果?》之前,先说一句实在的:别急着背概念,先看它在真实项目里到底解决什么问题。摘要摘要:2026 年,AI 编程工具已不是新鲜事,但能…

2026/7/28 2:03:26 阅读更多 →
研究生论文写作必备的10款AI工具与效率提升方案

研究生论文写作必备的10款AI工具与效率提升方案

1. 研究生论文写作的AI工具革命去年帮导师带研一新生时,有个场景让我印象深刻:凌晨两点的实验室里,五个学生围着电脑屏幕,反复修改着论文第三版的参考文献格式。这种场景在高校里太常见了——90%的研究生把30%的论文时间浪费在格式…

2026/7/28 2:03:26 阅读更多 →
ClickHouse merge引擎详解以及应用

ClickHouse merge引擎详解以及应用

一、理解 Merge引擎 (通常用于系统表,非用户数据) 用途:​ Merge引擎本身不存储数据,它的主要作用是提供对多个底层表(通常是结构相同的 MergeTree表)的统一查询视图,可以将它看作一个逻辑上的联合查询器 工作机制: 指定一个数据库和一个用于匹配表名的正则表达式&…

2026/7/28 2:03:26 阅读更多 →
ARM架构挑战:在Android设备上运行Windows应用的完整技术框架

ARM架构挑战:在Android设备上运行Windows应用的完整技术框架

ARM架构挑战:在Android设备上运行Windows应用的完整技术框架 【免费下载链接】winlator Android application for running Windows applications with Wine and Box86/Box64 项目地址: https://gitcode.com/GitHub_Trending/wi/winlator 当你在Android手机上…

2026/7/28 2:03:26 阅读更多 →
LangGraph实战:构建高效多智能体协作系统

LangGraph实战:构建高效多智能体协作系统

1. 项目概述:大模型协作开发的新范式三年前我第一次尝试用GPT-3构建客服机器人时,整整两周都困在单线程对话的泥潭里——用户问天气、转人工、查订单这三个简单需求,就需要反复重写prompt逻辑。直到发现LangChain的Agent机制才恍然大悟&#…

2026/7/28 2:03:26 阅读更多 →
AI办公升级:腾讯WorkBuddy领跑智能体赛道,企业级落地指南

AI办公升级:腾讯WorkBuddy领跑智能体赛道,企业级落地指南

易观分析近期发布的一份报告,在科技圈激起了不小的涟漪。 数据显示,腾讯推出的效率类AI智能体服务WorkBuddy,在6月份访问量突破2097万次,跃居行业第一,甚至超过了同期字节Trae与阿里QoderWork的访问量总和。 这不仅是一…

2026/7/28 2:02:26 阅读更多 →

日新闻

告别臃肿!3步让你的暗影精灵笔记本重获新生

告别臃肿!3步让你的暗影精灵笔记本重获新生

告别臃肿!3步让你的暗影精灵笔记本重获新生 【免费下载链接】OmenSuperHub Control Omen laptop performance, fan speeds, and keyboard lighting, and unlock power limits. 项目地址: https://gitcode.com/gh_mirrors/om/OmenSuperHub 你是否也曾为官方Om…

2026/7/28 0:00:43 阅读更多 →
RAG必踩坑!财报法规检索不准?这款开源工具让答案浮出水面,准确率飙升98.7%!

RAG必踩坑!财报法规检索不准?这款开源工具让答案浮出水面,准确率飙升98.7%!

做 RAG 的人应该都踩过这个致命的坑:把几百页的财报、法规、技术手册扔给向量库,问一个具体问题,搜出来的全是沾边但没用的内容 —— 关键信息要么被硬切块拆碎了,要么藏在几十条结果的最下面。语义相似≠真正相关,这个…

2026/7/28 0:00:43 阅读更多 →
抖音视频文案提取工具全指南:免费2026版、手机App、在线工具一网打尽

抖音视频文案提取工具全指南:免费2026版、手机App、在线工具一网打尽

2026年做短视频运营,从抖音上扒文案早就不是偷偷抄笔记的事了。我刚开始做内容的时候,每天刷半小时抖音,手动把爆款视频的口播敲进备忘录,一条2分钟的视频得花十来分钟,碰到语速快的还要反复回听。后来试了一圈工具&am…

2026/7/28 0:00:43 阅读更多 →

周新闻

深度学习道路桥梁裂缝检测系统 道路桥梁裂缝检测数据集 道路桥梁病害识别检测数据集

深度学习道路桥梁裂缝检测系统 道路桥梁裂缝检测数据集 道路桥梁病害识别检测数据集

深度学习道路桥梁裂缝检测系统 数据集6000张 完整源码已标注数据集训练好的模型环境配置教程程序运行说明文档,可以直接使用!系统支持图片、视频、摄像头等多种方式检测裂缝,功能强大实用。 1数据集6000张 8各类别

2026/7/27 4:33:59 阅读更多 →
深度学习YOLO模型如何训练 PUBG 绝地求生目标检测数据集

深度学习YOLO模型如何训练 PUBG 绝地求生目标检测数据集

pubg数据集 精选原图1.42万数据 1.49万标签 无任何重复、算法增强或冗余图像! pubg绝地求生目标检测数据集 1分类:e_body,14905个标签,txt格式 共计14244张图,99%为640*640尺寸图像 适合yolo目标检测、AI训练关键词&am…

2026/7/27 6:31:56 阅读更多 →
Apex英雄目标检测数据集 深度学习框架YOLO如何训练APEX数据集

Apex英雄目标检测数据集 深度学习框架YOLO如何训练APEX数据集

Apex检测数据集数据集详情检测类别: allies enemy tag图片总量:7247张训练集:5139张验证集:1425张测试集:683张标注状态:全部已标注,即拿即用数据格式:支持YOLO格式及其他格式&#…

2026/7/27 4:01:12 阅读更多 →

月新闻