大模型技术与应用课件 第8章 基于RWKV模型和DeepSeek框架的AI小说写作_第1页
大模型技术与应用课件 第8章 基于RWKV模型和DeepSeek框架的AI小说写作_第2页
大模型技术与应用课件 第8章 基于RWKV模型和DeepSeek框架的AI小说写作_第3页
大模型技术与应用课件 第8章 基于RWKV模型和DeepSeek框架的AI小说写作_第4页
大模型技术与应用课件 第8章 基于RWKV模型和DeepSeek框架的AI小说写作_第5页
已阅读5页,还剩27页未读 继续免费阅读

下载本文档

版权说明:本文档由用户提供并上传,收益归属内容提供方,若内容存在侵权,请进行举报或认领

文档简介

第8章

基于RWKV模型和DeepSeek框架的AI小说写作CONTENTS目录01

概述02

案例介绍03

详细步骤04

改进后的DeepSeek模型效果05

总结与展望01概述生成式人工智能发展热点

生成式AI的核心特征生成式人工智能是当前AI发展的热点领域,其核心能力在于能够创造出全新的内容,涵盖文字、图片、视频等多种形式,突破了传统AI以分析和识别为主的应用边界。

本章案例主题本章以AI小说写作作为实践案例,通过具体操作展示大模型在创作领域的应用特点,帮助读者直观体验生成式AI的文本创作流程与效果。学习目标

了解文本生成技术脉络掌握文本生成技术的发展历程,包括从早期序列模型到现代大语言模型的演进,理解不同技术阶段的核心突破。

理解RWKV机制原理深入理解RWKV机制结合循环与注意力机制的改进方法,掌握其在长序列处理中的实现逻辑与优势。

应用AI小说写作方法学会基于RWKV模型的AI小说写作思路,能够独立完成从数据准备到文本生成的全流程操作。

分析案例实现方式通过本章案例,掌握基于RWKV模型的AI小说写作技术细节,具备分析和复现类似文本生成项目的能力。基于生成式文本摘要的领域信息提取

模型框架核心组件模型框架包含局部编码器、混合层级全局编码器(MHT)、全局门控单元、指针生成器及LSTM解码器,协同完成文本语义提取与生成任务。

编码器功能局部编码器负责提取文本局部特征,MHT全局编码器捕捉全局依赖关系,两者共同读取输入文本并提取语义信息。

全局门控与解码机制全局门控单元筛选语义信息并生成上下文向量,指针生成器利用该向量直接复制未登录词,LSTM解码器结合上下文向量生成输出文本。

特征提取逻辑通过Pgen参数控制特征关注频率(Pgen表示特征出现,1-Pgen表示未出现),实现对正向、反向特征出现频率的精准提取。案例创新RWKV机制的融合创新

采用RWKV(recurrentweightedkey-value)机制,结合循环与注意力机制优势,相比传统Transformer更高效处理长序列数据,提升模型效率与性能。自定义分词与预处理

基于SentencePiece库训练分词模型,将文本编码为TokenID,适配中文等复杂语言特性,提升数据处理精度。高效训练与生成策略

引入层归一化、自定义损失函数优化训练过程,生成阶段采用核采样(nucleussampling)技术,平衡文本质量与多样性。案例价值与意义

核心价值通过RWKV机制与自定义训练策略提升文本生成连贯性与自然度;改进预处理方法增强模型对多场景任务的泛化能力;优化资源利用,支持CPU高效训练与生成,降低硬件依赖。

技术发展意义创新模型与训练方法推动NLP文本生成领域技术进步,为长序列处理提供新解决方案。

应用普及意义高效文本生成模型可广泛应用于聊天机器人、内容创作、自动摘要等场景,促进人工智能技术的实际落地与普及。

研究创新意义展示改进现有模型解决实际问题的思路,为研究者与开发者提供技术参考,激发NLP领域的进一步探索与创新。02案例介绍案例文件

数据文件包含bird_shooter.txt(训练分词模型的原始文本)、test.dat和train.dat(测试集与训练集数据)、ds_model.pth(训练好的模型)、wangwen-2024-04-20.json(改进后的字典)、wangwen-2024-04-20.pth(预处理模型)。

程序文件包括data_set.py(文本数据处理函数)、ds_model.py(定义基于Transformer架构的DeepSeek模型)、train.py(训练DeepSeek模型的脚本)、4_generate.py(加载预训练模型生成文本)、model.py(定义基于RWKV机制的DeepSeek模型)、utils.py(处理模型输出与设置随机种子)、run.py(生成文本)、ui.py(Tkinter构建的GUI应用)、handle.py(GUI事件处理函数)。核心模型——RWKV模型模型特点融合RNN与Transformer优势的序列模型,训练时采用类Transformer架构实现并行化,推理时类似RNN通过状态存储高效处理长序列,兼顾训练效率与推理内存优化。输入与线性投影输入X(维度E×T)经4个独立线性层与权重矩阵W、R、K、V相乘,投影至中间维度M×M,将原始输入转换为模型可处理的内部表示。核心计算通过门控函数σ(R)控制信息传递比例,结合元素级相乘实现加权键值交互,既保留Transformer长程关联捕捉能力,又具备RNN序列处理效率。时间维度残差前向通过残差连接在时间维度传递信息,缓解梯度消失问题,支持超长序列(如万字文本)高效处理。脉冲神经元与输出引入脉冲神经网络(SNN)机制,膜电位积累至阈值触发脉冲输出,将RWKV连续值输出转换为二进制脉冲信号,兼具低功耗与类脑特性。使用的库

PyTorch开源机器学习库,提供张量计算与自动微分功能,用于构建和训练DeepSeek模型。

SentencePiece文本处理库,支持无语言依赖的分词,用于训练分词模型及文本编码为TokenID。

NumPy科学计算库,专注多维数组与矩阵运算,用于数据预处理及数值计算。

torchinfo轻量级工具,打印PyTorch模型详细信息,包括各层输出形状与参数数量。使用的库

onnxruntime加载和运行ONNX格式模型,实现模型跨框架与平台部署。

HuggingFaceTransformers提供预训练模型与NLP工具,支持文本生成等任务。

Python标准库json模块用于JSON文件解析,os模块处理文件路径与环境变量。03详细步骤预处理文本数据(分词)导入核心库引入numpy(数值计算)、os(文件操作)、sentencepiece(分词处理)、sys(系统交互)、torch(深度学习框架)等关键库,为文本预处理提供工具支持。训练分词模型通过train_model函数,以bird_shooter.txt为输入,训练生成bird_shooter.model分词模型,词汇表大小设为16000,采用Unigram模型类型,字符覆盖率达99.95%。生成数据集利用gen_dataset函数,按9:1比例将文本分割为训练集(train.dat,523773tokens)和测试集(test.dat,59378tokens),并编码为TokenID存储。生成批次样本get_batch函数从数据中随机抽取批次样本,批次大小4,序列长度16,返回输入序列x和目标序列y,用于模型训练的批量输入。测试样本与执行脚本test_samples函数验证样本生成效果,执行data_set.py脚本后得到分词模型、词汇表及日志文件,日志记录训练参数、token数量等关键信息。配置DeepSeek模型的基本结构01定义DeepSeek配置类dsConfig类存储模型核心参数,包括词汇表大小(16000)、序列长度(128)、模型维度(128)、层数(4)、头数(4)、偏置(True)及dropout比率(0.0)。02实现位置编码SinusoidPE类采用正弦余弦函数生成位置编码,添加至输入序列以提供位置信息,公式为pe[:,0::2]=sin(position*div_term),pe[:,1::2]=cos(position*div_term)。03构建自注意力机制SelfAttention类实现多头注意力,通过线性层生成Q、K、V,计算注意力分数并应用因果掩码,确保模型仅关注序列左侧部分,输出经投影层和dropout处理。04设计前馈神经网络FeedFoward类包含两层线性层,中间使用GELU激活函数(GeLU(x)=x*Φ(x),Φ为正态分布累积函数),相比ReLU更平滑,缓解神经元死亡问题。05组装解码器块与模型主体Block类由层归一化、自注意力层、前馈网络层组成,采用残差连接;dsModel类整合嵌入层、位置编码、解码器块序列及输出层,实现文本生成的端到端架构。训练DeepSeek模型初始化与环境配置创建DeepSeekConfig实例,设置批次大小32、dropout0.1,模型加载至GPU(优先)或CPU,优化器选用AdamW,学习率1e-3,最大迭代次数12000次。数据准备使用np.memmap加载train.dat和test.dat数据集,以内存映射方式高效处理大型数据,避免内存溢出。批次数据获取get_batch函数根据训练/测试模式,随机抽取序列长度为seq_len的样本,返回输入x和目标y,GPU模式下使用pin_memory加速数据传输。训练过程迭代12000次,每次迭代含前向传播(计算loss)、反向传播(梯度计算)、参数更新,打印迭代次数与loss值,训练结束保存模型权重至ds_model.pth。训练效果观察loss值随迭代逐步下降,表明模型学习有效,但需关注过拟合/欠拟合问题,可通过验证集监控泛化能力。通过DeepSeek模型生成文本

加载模型与分词器加载预训练模型权重ds_model.pth至指定设备(CPU/GPU),设置为评估模式;通过load_tokenizer加载bird_shooter.model分词器,确保文本编码一致性。

文本生成流程用户输入(如“你若安好便是晴天”)经分词器编码为TokenID,调用model.generate()方法,设置max_new_tokens=50生成文本,输出解码为自然语言。

生成效果展示模型尝试根据输入生成连贯文本,但部分输出内容存在逻辑混乱、语义不连贯问题,需进一步优化(如图8.18所示)。效果不佳原因分析与改进方法效果不佳原因训练数据量不足或质量不高、模型超参数(序列长度、层数等)未优化、训练迭代次数不够、生成策略(温度、采样方法)不当、中文处理复杂等。改进方法增加高质量多样化训练数据;调整模型配置(如增大层数、头数);延长训练时间或采用更优优化器;优化生成策略(调整top_p、temperature);针对中文特性优化预处理与模型结构。通过RWKV机制改进模型RWKV机制核心优势结合RNN与Transformer优点,通过时间混合(TimeMix)和通道混合(ChannelMix)捕捉长序列依赖,时间复杂度O(Td)、空间复杂度O(d),高效处理超长文本。时间混合模块RWKV_TimeMix类通过门控函数(σ(R))、加权键值交互(H=K⊙σ(R))及时间维度残差连接,替代传统注意力机制,实现高效序列信息传递。通道混合模块RWKV_ChannelMix类通过时间偏移(ZeroPad2d)、MISH激活函数(F.mish(k)*v)及线性变换,混合通道维度特征,增强模型表达能力。模型架构组装Block类整合层归一化、TimeMix、ChannelMix,DS类包含嵌入层、Block序列、输出层,支持复制机制(copy_mask),提升长序列处理性能。处理PyTorch模型的输出与设置随机数种子

模型输出处理函数to_float()函数将PyTorch张量转换为浮点数;sample_logits()函数支持核采样(top_p)和温度调整,从概率分布中采样下一个token,控制生成文本多样性。

随机数种子设置set_seed()函数统一设置random、numpy、torch(含GPU)的随机种子,确保实验可重复性,避免随机因素影响结果对比。04改进后的DeepSeek模型效果环境设置运行设备选择设置运行设备为CPU,适用于小规模或测试任务,无需依赖GPU加速。文件路径配置模型文件路径为'model/wangwen-2024-04-20',词汇表文件路径与之相同,确保模型与词汇表匹配。生成参数设定生成文本次数(NUM_OF_RUNS)设为999次,每次生成长度(LENGTH_OF_EACH)为512字符;top_p参数控制随机性,top_p=1时考虑所有可能词汇,top_p_newline=0.9用于换行符概率筛选。上下文准备

模型核心参数定义上下文长度(ctx_len)=512,层数(n_layer)=12,注意力头数(n_head)=12,嵌入维度(n_embd)=768(12头×64维度/头),注意力层维度(n_attn)和前馈网络维度(n_ffn)均为768。

上下文文本清理以“三体舰队”为输入上下文,通过strip()去除首尾空白字符,split('n')按行分割,循环去除每行全角空格(u3000),最终重组为规范文本,确保模型仅关注最后512个字符。

输入长度提示打印上下文长度及模型有效上下文范围,提醒用户模型仅处理最后512个字符,保障输入符合模型处理能力。加载词汇表与模型

01词汇表加载以UTF-16编码读取JSON格式词典文件,构建字符到索引(stoi)和索引到字符(itos)的映射,未知字符(UNKNOWN_CHAR)索引设为'ue083',词汇表大小为词典文件元素数量。

02模型加载与配置CPU环境下初始化DeepSeek模型,加载.pth权重文件;针对每层注意力机制,调整time_w、time_alpha、time_beta权重为time_ww矩阵,适配上下文长度,确保模型结构与权重匹配。

03设备适配处理根据RUN_DEVICE参数,GPU环境下将模型移至CUDA设备,DML环境下使用ONNXRuntime加载.onnx模型并配置会话选项,实现跨设备兼容。文本生成循环

上下文向量化将清理后的上下文文本转换为整数索引数组,未知字符用UNKNOWN_CHAR填充,数组长度记为real_len,作为生成起始状态。

模型推理与输出获取循环生成LENGTH_OF_EACH个字符:CPU/GPU环境下输入最近ctx_len个字符张量,通过模型前向传播获取输出;DML环境下对输入补零至ctx_len长度,使用ONNXRuntime会话推理,输出张量维度为[1,ctx_len,vocab_size]。

采样位置确定若real_len≥ctx_len,采样位置pos=-1(最后一个字符);否则pos=real_len-1,确保基于最新上下文生成下一个字符。采样函数与输出处理

条件采样策略若最后一个字符为换行符,使用top_p_newline=0.9采样;否则使用top_p=1采样,通过sample_logits函数实现核采样,控制生成文本的多样性与连贯性。

生成文本拼接将采样得到的字符索引追加至数组x,更新real_len;每2次迭代、生成结束、前10个字符或非GPU环境下,将整数索引转换为字符并打印,确保实时输出生成进度。通过flush=True强制刷新输出缓冲区,保证生成文本即时显示,提升用户体验。效果展示

生成就绪状态如图8.23所示,模型加载完成,上下文预处理完毕,显示“已就绪”状态,等待开始生成文本。

生成过程状态如图8.24所示,文本生成中,实时打印生成内容,展示模型逐字符扩展上下文的动态过程。

生成结束状态如图8.25所示,生成达到指定长度(512字符),输出完整文本,内容连贯且符合输入上下文主题(如“三体舰队”相关情节延续)。05总结与展望项目总

温馨提示

  • 1. 本站所有资源如无特殊说明,都需要本地电脑安装OFFICE2007和PDF阅读器。图纸软件为CAD,CAXA,PROE,UG,SolidWorks等.压缩文件请下载最新的WinRAR软件解压。
  • 2. 本站的文档不包含任何第三方提供的附件图纸等,如果需要附件,请联系上传者。文件的所有权益归上传用户所有。
  • 3. 本站RAR压缩包中若带图纸,网页内容里面会有图纸预览,若没有图纸预览就没有图纸。
  • 4. 未经权益所有人同意不得将文件中的内容挪作商业或盈利用途。
  • 5. 人人文库网仅提供信息存储空间,仅对用户上传内容的表现方式做保护处理,对用户上传分享的文档内容本身不做任何修改或编辑,并不能对任何下载内容负责。
  • 6. 下载文件中如有侵权或不适当内容,请与我们联系,我们立即纠正。
  • 7. 本站不保证下载资源的准确性、安全性和完整性, 同时也不承担用户因使用这些下载资源对自己和他人造成任何形式的伤害或损失。

最新文档

评论

0/150

提交评论