开始使用免费开始使用

为 MultiHeadAttention 类添加方法

在本练习中,您将从零开始完善 MultiHeadAttention 类,定义以下 4 个方法:

  • .split_heads():在各个注意力头之间拆分并变换输入嵌入向量
  • .compute_attention():计算缩放点积注意力权重,并与 values 矩阵相乘
  • .combine_heads():将注意力结果变换回与输入嵌入向量 x 相同的形状
  • .forward():调用以上方法,使输入嵌入向量依次通过这些过程

已将 torch.nnnn 导入,torch.nn.functional 可作为 F 使用,torch 也可用。

本练习是课程的一部分

使用 PyTorch 的 Transformer 模型

查看课程

交互式实操练习

通过完成这段示例代码来试试这个练习。

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super().__init__()
        self.num_heads = num_heads
        self.d_model = d_model
        self.head_dim = d_model // num_heads
        self.query_linear = nn.Linear(d_model, d_model, bias=False)
        self.key_linear = nn.Linear(d_model, d_model, bias=False)
        self.value_linear = nn.Linear(d_model, d_model, bias=False)
        self.output_linear = nn.Linear(d_model, d_model)

    def split_heads(self, x, batch_size):
        seq_length = x.size(1)
        # Split the input embeddings and permute
        x = x.____
        return x.permute(0, 2, 1, 3)
编辑并运行代码