PyTorch实现T5架构:交叉注意力层与解码器问题求助
T5架构PyTorch实现问题:交叉注意力/解码器错误与修复指导
一、RuntimeError错误修复
错误信息
return torch.embedding(weight, input, padding_idx, scale_grad_by_freq, sparse)
RuntimeError: Expected tensor for argument #1 'indices' to have one of the following scalar types: Long, Int; but got torch.FloatTensor instead (while checking arguments for embedding)
根源分析
- 交叉注意力错误添加相对位置偏置:T5架构仅自注意力层使用相对位置偏置,交叉注意力层不依赖该机制。给
T5CrossAttention添加相对位置偏置后,当query与key序列长度不一致时,会导致相对位置计算异常,甚至生成非整数索引触发类型错误。 - 冗余embedding调用:
T5类forward函数中x = self.embedding(src)是无效代码,其结果被后续encoder输出覆盖,无实际作用。 - mask维度不匹配:传入注意力层的mask形状为
(batch_size, seq_len),无法正确广播到注意力相似度矩阵(batch_size, heads, query_len, key_len)的维度,存在潜在广播错误。
具体修复代码
1. 修改T5CrossAttention类(移除相对位置偏置+修正mask维度)
class T5CrossAttention(nn.Module): def __init__( self, *, dim, context_dim = None, heads = 12, dim_head = 64, dropout = 0. ): super().__init__() inner_dim = dim_head * heads context_dim = default(context_dim, dim) self.heads = heads self.scale = dim_head ** -0.5 self.to_q = nn.Linear(dim, inner_dim, bias = False) self.to_k = nn.Linear(context_dim, inner_dim, bias = False) self.to_v = nn.Linear(context_dim, inner_dim, bias = False) self.to_out = nn.Linear(inner_dim, dim) # 移除不必要的相对位置偏置 self.dropout = nn.Dropout(dropout) def forward(self, x, context, mask = None, context_mask = None): b, n, _, h = *x.shape, self.heads kv_input = default(context, x) q, k, v = self.to_q(x), self.to_k(kv_input), self.to_v(kv_input) q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), (q, k, v)) q = q * self.scale sim = torch.einsum('b h i d, b h j d -> b h i j', q, k) mask_value = -torch.finfo(sim.dtype).max # 修正mask维度,适配注意力矩阵形状 if mask is not None: mask = mask[:, None, :, None] sim = sim.masked_fill_(~mask, mask_value) if context_mask is not None: context_mask = context_mask[:, None, None, :] sim = sim.masked_fill_(~context_mask, mask_value) attn = sim.softmax(dim = -1) attn = self.dropout(attn) out = torch.einsum('b h i j, b h j d -> b h i d', attn, v) out = rearrange(out, 'b h n d -> b n (h d)') return self.to_out(out)
2. 修改T5类forward函数(移除冗余代码)
class T5(nn.Module): # __init__部分保持不变 def forward(self, src, tgt, mask = None, context_mask = None): # 移除冗余的embedding调用 x = self.encoder(src, mask = mask) x = self.decoder(tgt, x, mask = mask, context_mask = context_mask) x = self.to_logits(x) return x
3. 修正T5SelfAttention的mask维度
class T5SelfAttention(nn.Module): # __init__部分保持不变 def forward(self, x, mask = None): b, n, _, h = *x.shape, self.heads q, k, v = self.to_q(x), self.to_k(x), self.to_v(x) q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), (q, k, v)) q = q * self.scale sim = torch.einsum('b h i d, b h j d -> b h i j', q, k) sim = self.relative_position_bias(sim) mask_value = -torch.finfo(sim.dtype).max if mask is not None: mask = mask[:, None, None, :] sim = sim.masked_fill_(~mask, mask_value) if self.causal: i, j = sim.shape[-2:] causal_mask = torch.ones((i, j), dtype = torch.bool, device = x.device).triu(j - i + 1) sim = sim.masked_fill(causal_mask, mask_value) attn = sim.softmax(dim = -1) attn = self.dropout(attn) out = torch.einsum('b h i j, b h j d -> b h i d', attn, v) out = rearrange(out, 'b h n d -> b n (h d)') return self.to_out(out)
二、交叉注意力与解码器实现的核心规范
1. 交叉注意力层T5规范
- 不使用相对位置偏置:交叉注意力专注于对齐encoder输出与decoder状态,无需额外位置偏置。
- 正确传递context mask:必须传入encoder输入对应的padding mask,避免模型关注无效padding内容。
2. 解码器模块T5规范
- 层结构顺序:每层严格遵循自注意力→交叉注意力→前馈网络的顺序,自注意力层必须开启因果掩码(
causal=True),防止模型看到未来token。 - 权重共享优化:T5默认共享encoder、decoder的词嵌入权重,以及输出logits层的权重,可补充以下代码实现完整共享:
class T5(nn.Module): def __init__(self, *, tie_token_emb=True, **kwargs): super().__init__() # ... 其他初始化代码 ... if tie_token_emb: self.encoder.token_emb.weight = self.decoder.token_emb.weight self.to_logits.weight = self.decoder.token_emb.weight
三、测试验证
修改后运行测试代码,将输出预期的torch.Size([1, 1024, 512]),错误完全修复。
内容的提问来源于stack exchange,提问作者Enrico Shippole
相关产品推荐
相关产品推荐

