You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在因果语言模型中如何结合缓存键值使用scaled_dot_product_attention?

在PyTorch的scaled_dot_product_attention中结合缓存键值对实现正确因果注意力

你遇到的问题本质是PyTorch的is_causal=True参数是为同长度的自注意力场景设计的,并不适配增量推理中缓存键值对(k/v长度大于q)的情况。

为什么is_causal=True不符合需求?

当q的长度L小于k/v的长度S时,is_causal=True生成的L×S掩码会强制每个q的第i个token只能关注k/v的前i+1个位置(也就是左下三角区域),这就导致新的q token无法访问完整的历史缓存,只能看到和自己索引对应的前几个k/v位置,完全不符合增量推理中“每个新token可以看到所有历史+当前及之前的新token”的需求。

正确的解决方案:手动构造适配缓存的掩码

你需要自己构造一个L×S的矩形掩码,让每个q的第i个token可以访问所有历史缓存 + 当前及之前的新token。以你给出的例子(q长度3,k/v长度6,其中前3个是历史缓存,后3个是新的键值对)为例,对应的掩码构造代码可以这样写:

import torch

# 假设参数
n_heads = 8
head_dim = 64

q = torch.rand((1, n_heads, 3, head_dim))
k = torch.rand((1, n_heads, 6, head_dim))
v = torch.rand((1, n_heads, 6, head_dim))

curr_seq_len = q.size(2)
total_seq_len = k.size(2)
past_seq_len = total_seq_len - curr_seq_len  # 这里是3

# 构造掩码:每个q的第i个token可以访问前(past_seq_len + i + 1)个k/v位置
max_access_indices = past_seq_len + torch.arange(curr_seq_len) + 1
# 生成布尔掩码,再转换为-inf形式
attn_mask = torch.arange(total_seq_len).unsqueeze(0) < max_access_indices.unsqueeze(1)
attn_mask = attn_mask.float().masked_fill(~attn_mask, float('-inf'))

# 调用注意力函数,注意不要设置is_causal=True
attn_output = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)

这段代码生成的掩码正好是你期望的形式:

[[0, 0, 0, 0, -inf, -inf]
 [0, 0, 0, 0,    0, -inf]
 [0, 0, 0, 0,    0,    0]]

关于增量推理中缓存的正确使用逻辑

在Transformer的增量推理场景中,通常我们会维护一个缓存的k/v张量,每次新输入q后,会把新生成的k/v拼接到缓存中。此时每个新的q token必须能访问所有历史缓存 + 到当前位置为止的新token,这种场景下is_causal=True的预设逻辑不适用,必须手动构造掩码才能实现正确的因果注意力行为。

内容的提问来源于stack exchange,提问作者turboderp

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.22 19:42:40