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

发布时间:2026/9/27 23:17:01

手撕 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
延伸阅读

更多相关文章

2026/9/27 23:17:01

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

/* 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:17:01

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:17:01

创维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/28 0:07:04

阿尔及利亚网站后缀选错?从零搭建外贸站避坑指南

阿尔及利亚网站后缀选错?从零搭建外贸站避坑指南 改个需求建站公司拖一周,这种憋屈事儿谁没遇到过?很多做西南外贸的朋友,本来想自己 从零搭建 个独立站,结果卡在最基础的域名后缀上,稀里糊涂买了个通用的…

2026/9/28 0:07:04

搞定wordpress主题v2ex图解步骤避坑指南

搞定wordpress主题v2ex图解步骤避坑指南 域名服务器配置让人头大?别急。很多老板在部署 wordpress主题v2ex 时,卡在环境搭建这一步,导致后续内容无法展示。 这份 图解步骤…

2026/9/28 0:07:04

群晖wordpress设置避坑指南 用免费工具搞定备案难题

群晖wordpress设置避坑指南 用免费工具搞定备案难题 很多老板刚接触群晖NAS,想在上面搭个WordPress博客或企业站,心里其实七上八下。最头疼的不是代码,而是那个让人头秃的备案流程。看着工信部系统里的步骤,再加上服务器IP变更、…

2026/9/28 0:07:04

西安seo培训机构排名解析:从零搭建官网防黑挂马实战指南

西安seo培训机构排名解析:从零搭建官网防黑挂马实战指南 上周刚帮一个做建材的老总处理完网站危机。他的官网首页被植入了大量赌博链接,后台登录密码被改,SEO排名一夜归零。他问我:“我明明买了最贵的服务器,为什么还是被黑?”更扎心的是,他之前…

2026/9/28 0:07:04

网页设计怎样做才安全 保姆级建站教程防黑客

网页设计怎样做才安全 保姆级建站教程防黑客 网站做好了没人访问,往往不是内容不行,而是安全没过关。很多老板花大钱做站,上线三天就被挂马,百度一查全是违规链接,流量直接归零。这就是典型的“带病上线”。今天这篇 保姆级建站教程…

2026/9/28 0:02:04

建设网站北京市进阶技巧

北京建设网站选错技术栈,流量归零?3个对比评测帮你避坑 网站做好了没人访问,比没做还让人焦虑。很多北京本地的站长,明明代码写得漂亮,UI也在线,结果上线一个月,百度收录寥寥无几,自然流量几乎为零。这时候再回头找开发团队,对方只会甩锅说“内容…

2026/9/27 0:00:45

东莞市品牌网站建设报价常见报错与解决

东莞品牌网站建设报价单背后:一份保姆级建站教程避坑实录 网站做好了没人访问,这大概是很多老板最头疼的事。花了大几万做的品牌站,上线后流量惨淡,比路边摊还冷清。别急着骂外包公司,很多“东莞品牌网站建设报价”里藏着不少猫腻,比如用模板站冒充定制…

2026/9/27 0:00:45

如何划分训练/验证集:Spirula Studio五种eval_mode策略详解

如何划分训练/验证集:Spirula Studio五种eval_mode策略详解 【免费下载链接】spirula-studio Cross-vendor 3D Gaussian Splatting trainer - video to splat to mesh, Vulkan or CUDA. 项目地址: https://gitcode.com/GitHub_Trending/sp/spirula-studio Sp…

2026/9/27 0:00:45

SEO怎么推广速查手册新手避坑实战指南

SEO怎么推广速查手册新手避坑实战指南 模板网站太丑不够用?别急着加滤镜,那是治标不治本。很多老板盯着后台流量掉得眼红,却还在纠结首页Banner的圆角是不是3像素。这就像穿着西装去挖土,姿势不对,努力白费。我整理这份 速查手册…

2026/9/28 0:02:03

广州外贸网站建设推广:从零搭建全流程拆解与真实报价避坑

广州外贸网站建设推广:从零搭建全流程拆解与真实报价避坑 改个需求建站公司拖一周,后台改个文案还得再交一笔“技术维护费”。这种憋屈事儿,做外贸的朋友太熟悉了。很多老板在找广州外贸网站建设推广服务商时,光盯着首页好不好看,却忽略了从零搭建一个能…

2026/9/28 0:02:04

搞懂百度竞价推广价格,网站性能优化别掉链子

搞懂百度竞价推广价格,网站性能优化别掉链子 网站突然打不开,浏览器弹出红色警告“此网站存在安全风险”,后台一看全是乱码代码和奇怪的跳转链接。这种网站被黑挂马的绝望感,很多刚转行做网站的朋友都经历过,尤其是那些为了省几百块钱服务器费用的新手。…

2026/9/25 20:55:38

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

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

2026/9/26 19:58:38

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

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

2026/9/25 18:34:56

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

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

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

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

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