保姆级教程:用PyTorch手写Transformer的Causal Mask(附完整代码与逐行解析)

在自然语言处理领域,Transformer架构已经成为现代语言模型的基石。而其中关键的Causal Mask技术,则是确保模型在生成文本时不会"偷看"未来信息的关键机制。本文将带你从零开始,用PyTorch实现一个完整的Causal Mask生成函数,并深入解析每一行代码背后的设计思想。

1. Causal Mask的核心原理

Causal Mask,又称因果掩码,是Transformer解码器中不可或缺的组成部分。它的核心作用是限制模型在每个时间步只能访问当前位置及之前的信息,确保预测过程符合时间因果关系。

想象你正在教一个孩子造句:当他说出第一个词"我"时,你不能让他提前知道后面要说"爱编程",否则就失去了语言学习的意义。同理,Causal Mask就是让模型在生成每个词时,只能基于已经生成的上下文进行预测。

这种掩码通常表现为一个上三角矩阵(主对角线及以下为0,以上为负无穷),例如对于长度为4的序列:

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

在实际应用中,这个矩阵会与注意力分数相加,使得未来位置的注意力权重趋近于0,从而被有效屏蔽。

2. 环境准备与基础配置

在开始编码前,我们需要确保开发环境配置正确。以下是推荐的配置方案:

import torch
import torch.nn as nn

# 检查PyTorch版本
print(torch.__version__)  # 推荐1.12+版本

# 设置随机种子保证可复现性
torch.manual_seed(42)

# 选择设备
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

对于本教程,我们将创建一个适用于大多数现代Transformer模型的通用掩码生成函数。该函数需要处理以下关键参数:

  • input_shape : 输入张量的形状(batch_size, seq_length)
  • dtype : 掩码的数据类型(如torch.float32, torch.bfloat16)
  • device : 计算设备(cpu或cuda)
  • past_key_values_length : 过去键值缓存的长度(用于增量解码)

3. 完整代码实现与逐行解析

下面是我们精心设计的 make_causal_mask 函数实现,每一行都附有详细注释:

def make_causal_mask(
    input_shape: torch.Size,
    dtype: torch.dtype,
    device: torch.device,
    past_key_values_length: int = 0
) -> torch.Tensor:
    """
    生成因果注意力掩码
    
    参数:
        input_shape: 输入张量形状(batch_size, seq_length)
        dtype: 输出掩码的数据类型
        device: 输出设备
        past_key_values_length: 过去键值对的长度
        
    返回:
        (batch_size, 1, seq_length, seq_length + past_key_values_length)形状的掩码
    """
    batch_size, seq_length = input_shape
    
    # 创建初始掩码矩阵,填充为对应数据类型的最小值
    mask = torch.full(
        (seq_length, seq_length),
        torch.tensor(torch.finfo(dtype).min, device=device),
        device=device
    )
    
    # 生成位置索引序列
    seq_ids = torch.arange(seq_length, device=device)
    
    # 构造掩码条件:当列索引 <= 行索引时为True
    mask_cond = seq_ids[:, None] >= seq_ids[None, :]
    
    # 应用掩码条件,将满足条件的位置设为0
    mask.masked_fill_(mask_cond, 0)
    
    # 如果有过去键值,在左侧拼接零矩阵
    if past_key_values_length > 0:
        mask = torch.cat([
            torch.zeros(
                seq_length, 
                past_key_values_length,
                dtype=dtype,
                device=device
            ),
            mask
        ], dim=-1)
    
    # 扩展维度以匹配注意力头的形状
    mask = mask[None, None, :, :].expand(
        batch_size, 1, seq_length, seq_length + past_key_values_length
    )
    
    return mask

3.1 关键步骤深度解析

步骤1:初始化掩码矩阵

mask = torch.full(
    (seq_length, seq_length),
    torch.tensor(torch.finfo(dtype).min, device=device),
    device=device
)

这里使用 torch.full 创建一个正方形矩阵,填充值为该数据类型能表示的最小值。对于float32约为-3.4e38,bfloat16约为-3.39e38。这个极小的值在后续softmax计算中会产生接近0的概率。

步骤2:生成位置条件

seq_ids = torch.arange(seq_length, device=device)
mask_cond = seq_ids[:, None] >= seq_ids[None, :]

通过广播机制创建位置比较矩阵,其中 mask_cond[i,j] 为True表示位置j应该对位置i可见。例如对于seq_length=3:

[[ True, False, False],
 [ True,  True, False],
 [ True,  True, True]]

步骤3:应用掩码条件

mask.masked_fill_(mask_cond, 0)

将满足条件的位置设为0,其余保持极小值。这样在注意力计算中:

  • 0值位置会保留原始注意力分数
  • 极小值位置会被softmax压制

步骤4:处理历史缓存

if past_key_values_length > 0:
    mask = torch.cat([
        torch.zeros(seq_length, past_key_values_length, dtype=dtype, device=device),
        mask
    ], dim=-1)

在自回归生成任务中,模型会缓存之前时间步的键值对。这部分代码在掩码左侧拼接零矩阵,确保模型可以完全访问历史信息。

4. 实战测试与常见问题

让我们通过具体示例验证函数行为:

# 基础测试
basic_mask = make_causal_mask(
    input_shape=(2, 4),
    dtype=torch.float32,
    device="cpu"
)
print("基础掩码形状:", basic_mask.shape)
print("示例掩码值:\n", basic_mask[0, 0])

# 带历史缓存的测试
past_mask = make_causal_mask(
    input_shape=(1, 3),
    dtype=torch.bfloat16,
    device="cpu",
    past_key_values_length=2
)
print("\n带历史缓存的掩码形状:", past_mask.shape)
print("示例掩码值:\n", past_mask[0, 0])

预期输出:

基础掩码形状: torch.Size([2, 1, 4, 4])
示例掩码值:
 tensor([[0., -inf, -inf, -inf],
        [0., 0., -inf, -inf],
        [0., 0., 0., -inf],
        [0., 0., 0., 0.]])

带历史缓存的掩码形状: torch.Size([1, 1, 3, 5])
示例掩码值:
 tensor([[0., 0., 0., -inf, -inf],
        [0., 0., 0., 0., -inf],
        [0., 0., 0., 0., 0.]], dtype=torch.bfloat16)

4.1 常见问题排查

问题1:掩码形状不符合预期

确保输入形状是(batch_size, seq_length)而不是(seq_length, batch_size)。常见错误是维度顺序颠倒。

问题2:数据类型导致的数值问题

当使用bfloat16等低精度类型时,极小值可能不够"小",导致掩码效果不佳。可以显式检查:

print(torch.finfo(dtype).min)  # 确认最小值足够小

问题3:设备不一致错误

如果输入张量在GPU而掩码在CPU,会导致运行时错误。最佳实践是:

mask = mask.to(input_ids.device)  # 确保设备一致

5. 高级应用与性能优化

5.1 内存高效实现

对于超长序列,完整存储N×N掩码矩阵会消耗大量内存。可以采用以下优化策略:

# 动态生成掩码
def efficient_causal_mask(seq_length, device):
    return torch.triu(
        torch.full((seq_length, seq_length), float('-inf'), device=device),
        diagonal=1
    )

# 使用示例
attn_scores = attn_scores + efficient_causal_mask(seq_len, attn_scores.device)

5.2 混合精度训练

当使用自动混合精度(AMP)时,需要注意掩码数据类型与计算类型一致:

with torch.cuda.amp.autocast():
    # 确保掩码与注意力分数类型匹配
    mask = make_causal_mask(input_shape, torch.float16, device)
    attn_output = torch.softmax(attn_scores + mask, dim=-1)

5.3 并行生成优化

在批量生成不同长度序列时,可以通过以下方式创建批量掩码:

def batch_causal_mask(lengths, max_len, dtype, device):
    """ lengths: 各序列实际长度的一维张量 """
    mask = torch.full((len(lengths), max_len), float('-inf'), device=device)
    for i, l in enumerate(lengths):
        mask[i, :l] = 0
    return mask.unsqueeze(1).unsqueeze(2)

6. 完整示例:集成到Transformer层

让我们将实现的掩码集成到一个简化的Transformer解码层中:

class TransformerDecoderLayer(nn.Module):
    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
        super().__init__()
        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
        self.cross_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
        self.ffn = nn.Sequential(
            nn.Linear(d_model, dim_feedforward),
            nn.ReLU(),
            nn.Linear(dim_feedforward, d_model)
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.norm3 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)
        
    def forward(self, tgt, memory, tgt_mask=None):
        # 自注意力层
        attn_output, _ = self.self_attn(
            tgt, tgt, tgt,
            attn_mask=tgt_mask,
            need_weights=False
        )
        tgt = tgt + self.dropout(attn_output)
        tgt = self.norm1(tgt)
        
        # 交叉注意力层
        attn_output, _ = self.cross_attn(
            tgt, memory, memory,
            need_weights=False
        )
        tgt = tgt + self.dropout(attn_output)
        tgt = self.norm2(tgt)
        
        # 前馈网络
        ffn_output = self.ffn(tgt)
        tgt = tgt + self.dropout(ffn_output)
        tgt = self.norm3(tgt)
        
        return tgt

# 使用示例
decoder_layer = TransformerDecoderLayer(d_model=512, nhead=8)
tgt = torch.rand(10, 32, 512)  # (seq_len, batch_size, d_model)
memory = torch.rand(20, 32, 512)
tgt_mask = make_causal_mask(
    input_shape=(32, 10),
    dtype=torch.float32,
    device="cpu"
)

output = decoder_layer(tgt, memory, tgt_mask=tgt_mask)
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐