自定义Multihead Attention类因果注意力存在数据泄露问题排查
自定义Multihead Attention因果注意力的数据泄露问题
我在学习Transformer模型时,用PyTorch实现了一个仅包含基础功能的自定义Multihead Attention类,但发现在因果注意力场景下(token不能关注未来token)存在数据泄露问题——这个结论是通过和torch.nn.MultiheadAttention类对比测试得到的。
我猜测问题出在掩码的应用方式上,但多次排查都没找到根源。已经验证过二维掩码能正确广播到四维张量,也确认了掩码的目标token是对的。
自定义MultiHeadAttention实现代码
class MultiHeadAttention(nn.Module): def __init__(self, n_heads, d_model, dropout=0.1): super().__init__() self.n_heads = n_heads self.d_model = d_model self.dropout = nn.Dropout(dropout) self.query = nn.Linear(d_model, d_model, bias=False) self.key = nn.Linear(d_model, d_model, bias=False) self.value = nn.Linear(d_model, d_model, bias=False) self.att_proj = nn.Linear(d_model, d_model, bias=False) self.register_buffer('mask', torch.triu(torch.ones(block_size, block_size), diagonal=1).bool()) def forward(self, x): q = x k = x v = x B,T,C = x.shape dk = d_model // n_heads # linear projections q = self.query(q) k = self.key(k) v = self.value(v) # add number of heads q = q.view(B,T,n_heads,dk).permute(0,2,1,3) # B,T,h,dk k = k.view(B,T,n_heads,dk).permute(0,2,1,3) v = v.view(B,T,n_heads,dk).permute(0,2,1,3) # attention x = q @ k.transpose(-2,-1) # B,h,T,dk @ B,h,dk,T --> B,h,T,T x = x * dk ** -0.5 # B,h,T,T x = x.masked_fill(self.mask, float('-inf')) # B,h,T,T x = F.softmax(x, dim=(-1)) # B,n_h,T,T x = x @ v # B,h,T,T @ B,T,h,dv --> B,h,T,dv x = x.view(B,T,-1) out = self.att_proj(x) # B,T,C return out
测试结果
测试中自定义类的训练损失为2.307、评估损失为2.278;使用官方类迭代9999次后,训练损失为2.469、评估损失为2.483。
配套Model类实现代码
class Model(nn.Module): def __init__(self, vocab_size, *args, **kwargs) -> None: super().__init__(*args, **kwargs) self.embedding_table = nn.Embedding(vocab_size, d_model) self.mha = MultiHeadAttention(n_heads, d_model) self.out = nn.Linear(d_model, vocab_size, bias=False) def forward(self, x, targets=None): x = self.embedding_table(x) B, T, C = x.shape x = self.mha(x) # B,T,C logits = self.out(x) # B,T,vocab_size if targets is not None: logits = logits.reshape(-1, logits.shape[-1]) targets = targets.reshape(-1) loss = F.cross_entropy(input=logits, target=targets) else: loss = None return logits, loss def generate(self, n_chars, ix): for _ in range(n_chars): logits, loss = self(ix) # B, T, C logits = logits[:,-1,:] # B, C -- we need to reshape to calculate probabilities probs = F.softmax(logits, dim=-1) # B, C next_ix = torch.multinomial(input=probs, num_samples=1) ix = torch.cat((ix, next_ix), dim=1) return ix
已尝试的排查方向
- 更换训练验证拆分方式
- 多种掩码实现方式(如tril填充-inf、triu标记True填充-inf)
- 确保对角线设为1,仅掩码未来token
内容的提问来源于stack exchange,提问作者EdnellaU
相关产品推荐
相关产品推荐

