新闻详情

新闻详情

首页 / 资讯中心 / 详情

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

发布时间:2026/9/27 23:17:00来源:尧图网络
手撕 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
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

3个月搞定官网+排名,揭秘网站推广排名公司背后的源码下载陷阱 2026/9/28 0:11:37

3个月搞定官网+排名,揭秘网站推广排名公司背后的源码下载陷阱

3个月搞定官网+排名,揭秘网站推广排名公司背后的源码下载陷阱 改个需求建站公司拖一周,源码下载还要加钱?这种憋屈事,我见得太多了。很多老板花几万块找所谓的 网站推广排名公司…

阅读更多 →
WordPress新版无法保存?5个排查步骤+最佳实践 2026/9/28 0:10:58

WordPress新版无法保存?5个排查步骤+最佳实践

WordPress新版无法保存?5个排查步骤+最佳实践 刚接手一个重庆做建材外贸的站点,后台一进去就发现不对劲。客户急得直跺脚,说刚改完产品描述,点保存没反应,刷新页面全没了。更吓人的是,之前有次更新完插件,首页突然弹出一堆博彩广告,点进去…

阅读更多 →
军博做网站公司从零搭建官网避坑指南 2026/9/28 0:10:51

军博做网站公司从零搭建官网避坑指南

军博做网站公司从零搭建官网避坑指南 网站做好了没人访问,这大概是很多站长最头疼的事。你花大价钱找 军博做网站公司 ,甚至自己 从零搭建 ,结果上线半个月,后台流量还是个位数。别急着怪算法,大概率是底子没打对。…

阅读更多 →
做游戏用什么电脑系统下载网站看3个实战案例 2026/9/28 0:10:32

做游戏用什么电脑系统下载网站看3个实战案例

做游戏用什么电脑系统下载网站看3个实战案例 别再被那些花里胡哨的模板网站忽悠了,看着高大上,实际跑起来卡顿、加载慢,用户留不住。很多刚入行的运营或者小团队老板,总想着找个“做游戏用什么电脑系统下载网站”这种关键词直接套模板,结果上线后发现搜…

阅读更多 →
不会代码做wordpressforum要花多少钱? 2026/9/28 0:10:26

不会代码做wordpressforum要花多少钱?

不会代码做wordpressforum要花多少钱? 手里拿着域名,脑子里全是想法,但打开编辑器就头大,自己不会代码想做网站,这大概是无数创业者最真实的写照。你想知道搞个 wordpressforum…

阅读更多 →
揭阳百度推广优化避坑指南:5个免费工具提升30%转化 2026/9/28 0:08:50

揭阳百度推广优化避坑指南:5个免费工具提升30%转化

揭阳百度推广优化避坑指南:5个免费工具提升30%转化 改个需求建站公司拖一周?别忍了。很多揭阳老板觉得百度推广难搞,其实是被“黑箱”操作坑了。今天直接甩出5个 免费工具 ,教你自己盯数据,不再当冤大头。 运营目标与指标:别只看点击量…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉