首页/新闻资讯/正文详情

手撕 Decoder 生成:因果掩码 + KV Cache,70 行 PyTorch 看懂流式输出为什么快

发布时间:2026/9/27 23:16:59 来源:云帆数科 栏目:资讯中心
手撕 Decoder 生成:因果掩码 + KV Cache,70 行 PyTorch 看懂流式输出为什么快
手撕 Decoder 生成因果掩码 KV Cache70 行 PyTorch 看懂流式输出为什么快上一篇手撕了 Transformer BlockFFN/残差/LayerNorm这篇接着往下走Decoder 怎么用这个 Block 逐 token 生成文本以及 KV Cache 为什么能让流式输出一路变快。同样的风格完整代码直接跑带数值自检不编数据。结论先放这儿方式每步计算总计算量一句话朴素生成每步重算全部历史O(n³)每生成一个字前面的字全部白算一遍KV Cache每步只算 1 个新 tokenO(n²)历史 k/v 存下来新 token 只拼上去两句话铁律因果掩码保证训练时看不到未来KV Cache 保证生成时不重算过去。一、因果掩码三行代码训练时整个序列一次前向但每个位置只能看左边T, S q.shape[2], k.shape[2] # 本段长度, 总可见长度 mask torch.triu(torch.ones(T, S, dtypetorch.bool), diagonalS - T 1) att (q k.transpose(-2, -1) / k.shape[-1] ** 0.5).masked_fill(mask, float(-inf)).softmax(-1)diagonalS-T1让位置 i 只能看到 0 到 S-Ti整段输入时ST就是标准下三角带 cache 逐 token 生成时T1掩码全 False——新 token 本来就该看到全部历史。二、完整代码单文件直接跑# decoder_gen.py — 朴素生成 vs KV Cache 生成 # 依赖pip install torch import time import torch import torch.nn as nn class Block(nn.Module): def __init__(self, d128, h4): super().__init__() self.h h self.ln1, self.ln2 nn.LayerNorm(d), nn.LayerNorm(d) # Pre-LN接上一篇 self.qkv nn.Linear(d, 3 * d) self.proj nn.Linear(d, d) self.ffn nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d)) def forward(self, x, cacheNone): B, T, D x.shape q, k, v self.qkv(self.ln1(x)).chunk(3, dim-1) q q.view(B, T, self.h, -1).transpose(1, 2) # (B, h, T, dh) k k.view(B, T, self.h, -1).transpose(1, 2) v v.view(B, T, self.h, -1).transpose(1, 2) if cache is not None: # KV Cache新 k/v 拼到历史后 if cache.get(k) is not None: k torch.cat([cache[k], k], dim2) v torch.cat([cache[v], v], dim2) cache[k], cache[v] k, v S k.shape[2] mask torch.triu(torch.ones(T, S, dtypetorch.bool), diagonalS - T 1) att q k.transpose(-2, -1) / k.shape[-1] ** 0.5 att att.masked_fill(mask, float(-inf)).softmax(-1) x x self.proj((att v).transpose(1, 2).reshape(B, T, D)) return x self.ffn(self.ln2(x)) class TinyModel(nn.Module): def __init__(self, vocab500, d128): super().__init__() self.emb nn.Embedding(vocab, d) self.pos nn.Embedding(512, d) self.blocks nn.ModuleList([Block(d) for _ in range(2)]) self.ln nn.LayerNorm(d) self.head nn.Linear(d, vocab, biasFalse) def forward(self, idx, caches): T idx.shape[1] S caches[0][k].shape[2] if caches[0].get(k) is not None else 0 x self.emb(idx) self.pos.weight[S:S T] # cache 模式下位置从 S 起算 for blk, c in zip(self.blocks, caches): x blk(x, c) return self.head(self.ln(x)) def generate_naive(model, prompt, n): idx prompt for _ in range(n): logits model(idx, [None] * len(model.blocks)) # 每步全部历史重算 idx torch.cat([idx, logits[:, -1].argmax(-1, keepdimTrue)]) return idx def generate_cached(model, prompt, n): caches [{} for _ in model.blocks] logits model(prompt, caches) # prefillprompt 的 k/v 一次算完 idx logits[:, -1].argmax(-1, keepdimTrue) out [idx] for _ in range(n - 1): logits model(idx, caches) # 每步只算 1 个新 token idx logits[:, -1].argmax(-1, keepdimTrue) out.append(idx) return torch.cat([prompt] out, dim1) if __name__ __main__: torch.manual_seed(0) model TinyModel().eval() prompt torch.randint(0, 500, (1, 5)) with torch.no_grad(): assert torch.equal(generate_naive(model, prompt, 30), generate_cached(model, prompt, 30)) # 两路输出逐 token 一致 t0 time.perf_counter(); generate_naive(model, prompt, 100) t1 time.perf_counter(); generate_cached(model, prompt, 100) print(f朴素: {t1 - t0:.2f}s KV Cache: {t2 - t1:.2f}s) print(self-check ok)自检两条都有含义torch.equal验证 KV Cache 没算错两路必须逐 token 一致计时的差距你自己跑一下就能看到——模型越长差距越大这就是流式输出能一个字一个字蹦的原因。三、三个踩坑自己实现生成循环都会遇到位置编码偏移cache 模式下新 token 的位置编码必须从 S已有长度起算从 0 重取会错位——掩码对了位置错了输出悄悄变差还不报错prefill 没做直接从第一个新 token 开始逐个喂prompt 部分被拆成一堆单步调用首个 token 延迟翻几倍。prompt 一次前向算完 k/v 才是 prefillcache 与 dropout带 cache 生成是推理路径模型必须.eval()否则 dropout 噪声让两路输出对不上自检直接失败四、和真实推理框架的差距这个玩具缺的是GQA/MQAk/v 头数比 q 少显存省几倍、滑动窗口、投机解码小模型起草大模型验收、连续批处理。但主干你已经有了因果掩码 KV Cache prefill所有推理框架都是在这个骨架上加工程优化。总结铁律压成三句因果掩码管训练时看不到未来KV Cache 管生成时不重算过去位置编码从已有长度起算prefill 一次算完 prompt两路输出逐 token 一致是 KV Cache 实现正确性的硬标准写完先跑这条 assert

相关推荐

嵌入式驱动开发培训如何选?看硬件、内核、调试三要素
嵌入式驱动开发培训如何选?看硬件、内核、调试三要素

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

UFS 3.1协议栈全解析:从UPIU到WriteBooster的工程实践
UFS 3.1协议栈全解析:从UPIU到WriteBooster的工程实践

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

创维HC2910机顶盒强刷海美迪安卓7.0固件教程
创维HC2910机顶盒强刷海美迪安卓7.0固件教程

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

using-lwc - code-graph
using-lwc - code-graph

LWC CodeGraph 索引 使用时机 对已检出代码的结构性问题使用 CodeGraph:符号定义、签名、调用者/被调用者、依赖流、文件拓扑、可达性,或跨符号/文件的变更影响。 跳过时机 对于单文件字面编辑、仅格式化工作、仅文档/配置工作、注释/日志字符串&#xf… · 2026/9/27 23:57:03

Github周刊2026W37:用更少认知负荷换取更高产出效率的五个开发实践
Github周刊2026W37:用更少认知负荷换取更高产出效率的五个开发实践

1. 这期周刊到底在聊什么先说清楚,这不是一篇翻译稿,也不是简单的链接罗列。Github周刊2026W37这一期,我翻来覆去看了三遍,最大的感受是:它把当下开发者圈子里几个看似不搭界的热点,用一条暗线串起来了——… · 2026/9/27 23:56:57

Tencent BrowserSkill:已登录浏览器与编码Agent的本地桥接方案
Tencent BrowserSkill:已登录浏览器与编码Agent的本地桥接方案

1. 这个项目到底在解决什么问题先说结论:Tencent BrowserSkill 做的事情,用一句话概括就是——在“已经登录了各种账号的真实浏览器”和“跑在终端里的编码 Agent”之间,架一座本地桥。让 Agent 不用重新登录、不用重新配置 Cookie、不用去啃… · 2026/9/27 23:56:57

using-lwc - strong-context
using-lwc - strong-context

LWC 强上下文与标签 使用时机 对一小部分经过显式审查的核心页面——规则、操作手册、安全策略或运行手册——使用标签,这些页面必须在无相关性搜索的情况下完整加载。 跳过时机 不要将标签用作搜索别名、主题标签、推断关键词或加载广泛语料的方式。如果页面只是松… · 2026/9/27 23:56:57

OpenCV车牌识别实战:从定位到识别的完整Pipeline
OpenCV车牌识别实战:从定位到识别的完整Pipeline

简介:本资源是一套完整的Python毕业设计项目——基于OpenCV的车牌识别系统实现方案,面向计算机、人工智能及相关专业本科生,解决课程设计、毕设选题与图像处理实践中的核心需求。压缩包共2000个文件,含1987张实拍车牌JPG样本&… · 2026/9/27 23:56:51

基于SringBoot的智慧博物馆预约平台的设计与实现(源码+文档+部署讲解等)
基于SringBoot的智慧博物馆预约平台的设计与实现(源码+文档+部署讲解等)

联系博主 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 … · 2026/9/27 23:56:45

MATLAB雷达信号脉冲压缩仿真:LFM线性调频、匹配滤波与距离分辨率实现
MATLAB雷达信号脉冲压缩仿真:LFM线性调频、匹配滤波与距离分辨率实现

简介:这套Matlab仿真工具完整呈现雷达信号脉冲压缩过程,从线性调频(LFM)信号生成、目标回波仿真到匹配滤波压缩处理均有可运行代码支撑,面向电子信息工程、计算机、数学等专业学生,适用于课程设计、期末大作… · 2026/9/27 0:00:01

汕头网站建设制作厂家避坑指南:5大注意事项救急
汕头网站建设制作厂家避坑指南:5大注意事项救急

汕头网站建设制作厂家避坑指南:5大注意事项救急 改个需求建站公司拖一周,这种憋屈事我见得太多了。 很多汕头老板找本地建站团队,签合同前看着方案挺美,一上线就变脸。 今天不聊虚的,直接拆解找 汕头网站建设制作厂家 时的5个核心 注意事项… · 2026/9/27 0:00:01

多模态虚假新闻检测实战:BERT+ResNet双塔与对比学习
多模态虚假新闻检测实战:BERT+ResNet双塔与对比学习

简介:基于PyTorch的多模态虚假新闻检测项目完整代码包,面向自然语言处理与计算机视觉交叉方向的开发者、科研人员及毕业设计选题者,解决社交媒体中文本与图像联合识别虚假新闻的问题。系统以BERT预训练模型提取文本语义特征,以Res… · 2026/9/27 0:00:01

MATLAB雷达信号脉冲压缩仿真:LFM线性调频、匹配滤波与距离分辨率实现
MATLAB雷达信号脉冲压缩仿真:LFM线性调频、匹配滤波与距离分辨率实现

简介:这套Matlab仿真工具完整呈现雷达信号脉冲压缩过程,从线性调频(LFM)信号生成、目标回波仿真到匹配滤波压缩处理均有可运行代码支撑,面向电子信息工程、计算机、数学等专业学生,适用于课程设计、期末大作… · 2026/9/27 0:00:01

汕头网站建设制作厂家避坑指南:5大注意事项救急
汕头网站建设制作厂家避坑指南:5大注意事项救急

汕头网站建设制作厂家避坑指南:5大注意事项救急 改个需求建站公司拖一周,这种憋屈事我见得太多了。 很多汕头老板找本地建站团队,签合同前看着方案挺美,一上线就变脸。 今天不聊虚的,直接拆解找 汕头网站建设制作厂家 时的5个核心 注意事项… · 2026/9/27 0:00:01

多模态虚假新闻检测实战:BERT+ResNet双塔与对比学习
多模态虚假新闻检测实战:BERT+ResNet双塔与对比学习

简介:基于PyTorch的多模态虚假新闻检测项目完整代码包,面向自然语言处理与计算机视觉交叉方向的开发者、科研人员及毕业设计选题者,解决社交媒体中文本与图像联合识别虚假新闻的问题。系统以BERT预训练模型提取文本语义特征,以Res… · 2026/9/27 0:00:01

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码