ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

手撕 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
返回列表