AI基本结构11-rnn实现简单nlp

发布时间:2026/9/11 21:16:55

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))))正常训练的老三样损失优化循环
延伸阅读

更多相关文章

2026/9/5 8:49:52

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

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

2026/9/10 23:43:43

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

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

2026/9/11 21:13:35

菜单行为函数:从二级联动到行为树的工程实践

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

2026/9/10 16:39:38

超人会飞不算本事:系统稳定依赖清晰规则与边界设计

开头先不绕弯子。“#斯坦李吐槽dc 所以超人是无缘无故会飞的嘛哈哈哈哈哈哈哈锤哥真是技术人才啊!#雷神 #复联”这类调侃式短标题,第一波冲击力在于它把两个宇宙的角色塞进同一个吐槽箱里,但细想一下就能发现,它真正碰到的根本不是…

2026/9/10 11:16:38

超人VS蜘蛛侠:拆解超级IP的影响力与传播方法论

把“蜘蛛侠 vs 超人”放在 CSDN 上聊,可能很多人第一反应是走错片场了。但如果把这两个角色看成“两个持续运营了 80 多年的文化产品”,你会发现,这场比较本质上是两个不同 IP 策略的长期结果对比:超人赢在定义了整个超级英雄题材…

2026/9/9 16:31:09

基于CNN的调制信号识别:MATLAB实现时频图分类实战

简介:本资源是一套面向通信工程与信号处理方向学习者、研究者的深度学习实践方案,聚焦调制信号自动检测与识别这一典型无线通信任务,解决传统方法依赖人工特征、低信噪比下性能下降等痛点。压缩包共12个文件(10.73MB)&…

2026/9/10 12:32:02

USB Type-C PCB布局分区设计:电源、高速信号与PD协议全攻略

做硬件这行,Type-C接口算是典型的“看着简单,做起来全坑”的东西。光引脚就24个,高低速信号、电源、控制线全部塞在一个小小的连接器里,如果PCB布局不做规划,打样回来基本就是“插上没反应”、“高速掉线”、“静电一打…

2026/9/10 15:19:50

系统编程学习原型如何补齐稳定性边界

系统编程学习原型如何补齐稳定性边界预算有限时&#xff0c;我先优化明显多余的复制&#xff0c;而不是猜测性地换容器。用借用传递只读数据通常就能减少分配&#xff1a; fn parse(line: &str) -> Result<Item, Error> { /* ... */ }用基准确认热点确实在分配&am…

2026/9/10 15:49:53

雨花区哪家财务公司代理记账比较好?

在雨花区&#xff0c;企业处理财税事务常常面临诸多挑战&#xff0c;选择一家靠谱的财务公司至关重要。湖南巨勤财务管理咨询有限公司就是本地正规实体财税服务机构&#xff0c;深耕本地工商财税行业多年&#xff0c;熟悉当地工商局、税务局最新政策与申报流程。主营公司注册、…

还想了解更多?直接咨询顾问

免费诊断 + 免费方案 + 透明报价。

全国咨询热线400-8866-253
免费获取方案
咨询二维码