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

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)

根源分析

  1. 交叉注意力错误添加相对位置偏置:T5架构仅自注意力层使用相对位置偏置,交叉注意力层不依赖该机制。给T5CrossAttention添加相对位置偏置后,当query与key序列长度不一致时,会导致相对位置计算异常,甚至生成非整数索引触发类型错误。
  2. 冗余embedding调用:T5类forward函数中x = self.embedding(src)是无效代码,其结果被后续encoder输出覆盖,无实际作用。
  3. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 05:45:25