LLM令牌遮蔽技术全解析:从原理到PyTorch实战

发布时间:2026/8/27 6:52:57
LLM令牌遮蔽技术全解析:从原理到PyTorch实战 1. 项目概述为什么LLM需要令牌遮蔽在大型语言模型LLM的训练和应用中我们常常会遇到一个看似简单却至关重要的问题模型应该如何“看”待和处理输入序列如果让模型一次性看到所有信息它可能会在训练时“作弊”或者在推理时产生不符合预期的输出。这就引出了“令牌遮蔽”技术的核心价值。简单来说令牌遮蔽就是有选择性地隐藏输入序列中的部分令牌Token让模型只能基于剩余可见的部分进行预测或理解。这不仅是预训练阶段如BERT的掩码语言模型的基石也是微调、推理乃至提升模型鲁棒性的关键技巧。我最初接触这个概念是在处理文本分类任务时模型总是过拟合训练数据中的某些特定词汇。后来在构建生成式模型时又遇到了如何防止模型在解码时“偷看”未来信息的问题。这些实际痛点让我深入研究了各种遮蔽策略。本文将聚焦于五种在LLM实践中高频出现、效果显著的令牌遮蔽技术并手把手带你用PyTorch实现它们。无论你是刚入门的新手还是想系统梳理这块知识的老手这些内容都能帮你构建更清晰、更实用的模型处理流程。2. 五种核心令牌遮蔽技术深度解析令牌遮蔽远不止是随机盖住几个词那么简单。不同的任务目标和模型架构需要截然不同的遮蔽哲学。下面这五种技术基本覆盖了从预训练到推理的各大核心场景。2.1 随机遮蔽预训练的基石随机遮蔽是掩码语言模型MLM的标配也是BERT系列模型预训练的核心。其思想非常简单在输入序列中随机选择一定比例例如15%的令牌并将它们替换为一个特殊的[MASK]令牌。模型的任务就是根据上下文预测这些被遮蔽位置原来的令牌是什么。为什么是15%这个数字是经验与权衡的结果。比例太低模型学习效率低下成本高昂比例太高上下文信息不足导致预测任务过于困难可能损害模型的语言理解能力。15%是一个在多数语料上能达到较好平衡的点。此外在BERT的原始实现中对于这15%的被选令牌并非全部替换为[MASK]其中80%用[MASK]替换10%随机替换为词表中的其他词10%保持原词不变。这种技巧旨在让模型不那么依赖[MASK]这个特定符号增强其泛化能力。实操心得比例调整对于领域特定文本如医学、法律由于专业术语密集可以适当降低遮蔽比例如10%以免破坏关键实体信息。动态遮蔽原始BERT对每个序列只做一次遮蔽。更优的做法是“动态遮蔽”即在每个训练周期Epoch都对同一数据生成不同的遮蔽模式能有效增加数据多样性提升模型鲁棒性。这在预训练中尤为重要。2.2 因果遮蔽生成式模型的纪律因果遮蔽也叫序列遮蔽或前瞻遮蔽是自回归模型如GPT系列的“生命线”。它的规则非常严格在预测当前位置的令牌时模型只能看到它之前左侧的令牌绝对不能看到当前及之后右侧的令牌。这完美模拟了人类逐字生成文本的过程。想象一下你在写文章你只能根据已经写出的内容来思考下一个词而不能参考还未写出来的部分。因果遮蔽在注意力机制中通过一个“掩码矩阵”来实现。这个矩阵是一个上三角矩阵对角线及以下的元素为0允许注意力对角线以上的元素为负无穷经过Softmax后变为0即遮蔽。例如对于序列长度为4的情况掩码矩阵如下[0, -inf, -inf, -inf] [0, 0, -inf, -inf] [0, 0, 0, -inf] [0, 0, 0, 0]这样第一个位置只能关注自己在有些实现中甚至不允许关注自己即对角线也是-inf第二个位置可以关注第一、二个位置依此类推。注意事项训练与推理的一致性因果遮蔽确保了训练教师强制和推理自回归生成过程的一致性这是生成模型能正常工作的根本。缓存机制在自回归推理时由于每次只新增一个令牌可以利用键值缓存KV Cache来避免重复计算之前所有令牌的注意力这是提升推理效率的关键优化而因果遮蔽的结构天然支持这种优化。2.3 填充遮蔽处理变长序列的标尺在实际应用中我们几乎不可能把所有句子都补成一样长那样会引入大量无意义的填充符浪费计算资源。更常见的做法是将一个批次Batch内的序列填充到该批次内的最大长度。填充遮蔽就是为了告诉模型“哪些是真实内容哪些是凑长度的填充符请忽略它们。”具体实现是生成一个与输入序列形状相同的布尔掩码Mask其中真实令牌的位置为False或0不遮蔽填充令牌[PAD]的位置为True或1遮蔽。在注意力计算中将这个掩码加到注意力权重上通常将填充位置加一个极大的负值如-1e9使得这些位置的注意力分数经过Softmax后趋近于0。核心要点注意力屏蔽与损失忽略填充遮蔽主要应用在注意力层防止模型关注无意义的[PAD]。此外在计算损失函数如交叉熵时也需要忽略填充位置对应的损失通常通过ignore_index参数实现。影响位置编码一些位置编码如旋转位置编码RoPE本身会处理序列长度但填充部分仍然需要遮蔽以确保位置信息的正确性。2.4 滑动窗口遮蔽长序列的局部视野Transformer的自注意力机制复杂度是序列长度的平方级O(n²)这严重限制了模型处理长文本的能力。滑动窗口遮蔽是一种高效的近似方法它让每个令牌只能关注其前后一定窗口大小例如512个令牌内的邻居而不是整个序列。这就像你读一本很厚的书在理解某一页的内容时你主要参考的是当前页及其前后几页而不是整本书从头到尾重新读一遍。滑动窗口遮蔽通过一个带状掩码矩阵来实现只有中心对角线附近的一个带状区域是“可见”的。优势与权衡优势将计算复杂度从 O(n²) 降低到 O(n * w)其中w是窗口大小使得处理超长序列如长达数万token的文档成为可能。Longformer、BigBird等模型就采用了这种思想。权衡牺牲了全局上下文信息。对于需要高度理解全文结构的任务如文章主旨归纳纯滑动窗口可能不够。因此这类模型通常会混合使用局部窗口注意力少量全局注意力如对某些特殊令牌如[CLS]赋予全局注意力。2.5 任务特定遮蔽定制化的信息控制这是最灵活的一类遮蔽其模式完全由下游任务的需求决定。它不再是固定的规则而是一种设计模式。句子对任务在诸如文本蕴含、句子相似度任务中我们有两个句子A和B。模型需要同时编码它们并理解其关系。通常的做法是在输入时用[SEP]分隔并在注意力层使用“分段遮蔽”让句子A内部的令牌可以互相关注句子B内部的令牌也可以互相关注但A和B之间的注意力可以被限制或赋予不同的模式例如让[CLS]令牌能关注所有令牌以聚合信息。知识遮蔽在信息抽取或问答任务中我们可能明确知道输入文本中的某些实体或答案是关键。为了鼓励模型更深入地理解上下文而非简单记忆可以在训练时对这些关键信息进行有针对性的遮蔽迫使模型从周边描述进行推理。结构化数据遮蔽当处理表格、代码等结构化文本时可以设计遮蔽单元如整行、整列、一个函数块而非单个令牌以学习结构化的表示。3. PyTorch实现详解与代码实战理论说得再多不如一行代码。下面我们将用PyTorch逐一实现上述五种遮蔽技术并嵌入到一个简化的Transformer编码器中进行演示。我们假设你已了解PyTorch和Transformer的基本使用。3.1 基础环境与数据准备首先我们需要一些模拟数据和一个简单的模型骨架。import torch import torch.nn as nn import torch.nn.functional as F import math # 模拟参数 batch_size 2 seq_len 10 d_model 512 # 模型隐藏层维度 num_heads 8 vocab_size 30522 # BERT词表大小示例 # 模拟输入token ids 和 attention mask (用于填充遮蔽) # 假设序列实际长度分别为8和6后面用0填充到10 input_ids torch.tensor([ [101, 2023, 2003, 1037, 2307, 102, 1045, 2003, 0, 0], # [CLS] hello world [MASK] [SEP] i am [PAD] [PAD] [101, 1045, 2003, 1037, 102, 0, 0, 0, 0, 0] # [CLS] i am [MASK] [SEP] [PAD] ... ]) attention_mask (input_ids ! 0).long() # 填充掩码非0位置为10位置为0 print(Input IDs:\n, input_ids) print(Attention Mask:\n, attention_mask)3.2 实现一随机遮蔽这里我们实现一个完整的随机遮蔽函数包含BERT风格的80-10-10策略。def create_random_mask(input_ids, mask_token_id103, vocab_size30522, mask_prob0.15): 创建随机掩码遵循BERT的80-10-10策略。 参数: input_ids: 原始token id张量形状为 [batch, seq_len] mask_token_id: 词表中[MASK]对应的id vocab_size: 词表大小 mask_prob: 随机遮蔽的概率 返回: masked_input_ids: 被遮蔽后的输入id labels: 用于计算MLM损失的标签未被遮蔽的位置为-100在交叉熵中忽略 labels input_ids.clone() # 初始化一个随机矩阵用于决定每个位置是否被选中 probability_matrix torch.full(labels.shape, mask_prob) # 特殊令牌如[CLS], [SEP], [PAD]不应该被遮蔽 special_tokens_mask (input_ids 101) | (input_ids 102) | (input_ids 0) # [CLS], [SEP], [PAD] probability_matrix.masked_fill_(special_tokens_mask, 0.0) # 根据概率矩阵采样被遮蔽的索引 masked_indices torch.bernoulli(probability_matrix).bool() # 将被遮蔽位置的label保留用于计算损失 # 未被遮蔽的位置label设为-100PyTorch的CrossEntropyLoss忽略索引 labels[~masked_indices] -100 # 80% 替换为 [MASK] indices_replaced torch.bernoulli(torch.full(labels.shape, 0.8)).bool() masked_indices input_ids[indices_replaced] mask_token_id # 10% 随机替换为其他词 indices_random torch.bernoulli(torch.full(labels.shape, 0.5)).bool() masked_indices ~indices_replaced random_words torch.randint(vocab_size, labels.shape, dtypetorch.long) input_ids[indices_random] random_words[indices_random] # 剩下的10% 保持原样 (masked_indices中既非replaced也非random的部分) # 此时input_ids已被修改labels已设置 return input_ids, labels # 测试随机遮蔽 masked_input, mlm_labels create_random_mask(input_ids.clone()) print(Masked Input IDs:\n, masked_input) print(MLM Labels (ignore index -100):\n, mlm_labels)3.3 实现二因果遮蔽因果遮蔽通常在注意力计算时实现。我们将其集成到一个简化的多头注意力层中。def causal_attention_mask(seq_len, devicecpu): 生成因果注意力掩码上三角矩阵为1/-inf # 创建一个形状为 [1, 1, seq_len, seq_len] 的掩码方便后续广播 mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() # 将上三角未来位置设为True需要被遮蔽在注意力计算中通常用 -inf 填充 return mask.to(device) class CausalMultiHeadAttention(nn.Module): 一个集成了因果遮蔽的简化多头注意力模块 def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0 self.d_k d_model // num_heads self.num_heads num_heads self.q_linear nn.Linear(d_model, d_model) self.k_linear nn.Linear(d_model, d_model) self.v_linear nn.Linear(d_model, d_model) self.out_linear nn.Linear(d_model, d_model) def forward(self, q, k, v, maskNone): # mask: 可选的填充掩码形状 [batch, 1, 1, seq_len] batch_size, seq_len, _ q.size() # 线性变换并分头 q self.q_linear(q).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) k self.k_linear(k).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) v self.v_linear(v).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) # 计算注意力分数 scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) # 应用因果掩码 causal_mask causal_attention_mask(seq_len, q.device).view(1, 1, seq_len, seq_len) scores scores.masked_fill(causal_mask, float(-inf)) # 应用填充掩码如果有 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # mask为0的位置是PAD # Softmax和注意力加权 attn_weights F.softmax(scores, dim-1) context torch.matmul(attn_weights, v) # 合并多头输出 context context.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) return self.out_linear(context) # 测试因果注意力 causal_attn CausalMultiHeadAttention(d_model, num_heads) # 假设我们有一些随机初始化的嵌入 embedding nn.Embedding(vocab_size, d_model) x embedding(input_ids) # [batch, seq_len, d_model] # 注意在解码器中q, k, v 通常都来自解码器自身自注意力 output causal_attn(x, x, x, maskattention_mask.unsqueeze(1).unsqueeze(2)) print(Causal Attention Output Shape:, output.shape)3.4 实现三填充遮蔽填充遮蔽通常作为注意力掩码的一部分与因果掩码结合使用。上面的CausalMultiHeadAttention已经包含了处理填充掩码的逻辑mask参数。关键是如何生成这个掩码。通常我们从input_ids或attention_mask张量开始。def get_padding_mask(input_ids, pad_token_id0): 根据input_ids生成填充注意力掩码。 返回形状为 [batch_size, 1, 1, seq_len] 的掩码便于广播到注意力头。 非填充位置为1填充位置为0。 # input_ids形状: [batch, seq_len] mask (input_ids ! pad_token_id).unsqueeze(1).unsqueeze(2) # - [batch, 1, 1, seq_len] return mask # 测试 padding_mask get_padding_mask(input_ids) print(Padding Mask Shape:, padding_mask.shape) print(Padding Mask for first sequence:, padding_mask[0]) # 在注意力计算中使用 scores.masked_fill(padding_mask 0, float(-inf))3.5 实现四滑动窗口遮蔽滑动窗口遮蔽的实现核心是构造一个带状掩码矩阵。def sliding_window_mask(seq_len, window_size, devicecpu): 生成滑动窗口注意力掩码。 参数: seq_len: 序列长度 window_size: 单侧窗口大小。每个位置可以关注前后各window_size个token总共2*window_size1。 返回: mask: 布尔张量形状 [1, 1, seq_len, seq_len]True表示需要被遮蔽。 # 创建一个全为True遮蔽的矩阵 mask torch.ones(seq_len, seq_len, dtypetorch.bool, devicedevice) # 将对角线带状区域设为False不遮蔽 for i in range(seq_len): start max(0, i - window_size) end min(seq_len, i window_size 1) mask[i, start:end] False return mask.view(1, 1, seq_len, seq_len) # 测试滑动窗口掩码 window_size 2 sw_mask sliding_window_mask(seq_len, window_size) print(fSliding Window Mask (window_size{window_size}) for first 5x5 positions:) print(sw_mask[0, 0, :5, :5]) # 在注意力计算中scores.masked_fill_(sw_mask, float(-inf))3.6 实现五任务特定遮蔽示例——句子对遮蔽以句子对分类任务为例我们构造一个掩码让句子A和句子B内部的token可以互相关注但限制跨句子的注意力除了[CLS]等特殊token。def sentence_pair_mask(input_ids, sep_token_id102): 为句子对任务生成分段注意力掩码。 假设输入格式为: [CLS] SentA [SEP] SentB [SEP] [PAD]... 返回的掩码允许句子内部全连接但屏蔽句子间的注意力。 这是一个简化示例实际中[CLS]可能被特殊处理。 batch_size, seq_len input_ids.shape mask torch.zeros(batch_size, 1, seq_len, seq_len, dtypetorch.bool) for b in range(batch_size): # 找到第一个[SEP]的位置 sep_positions (input_ids[b] sep_token_id).nonzero(as_tupleTrue)[0] if len(sep_positions) 2: # 如果没有两个[SEP]则退化为全连接无遮蔽 continue sep1, sep2 sep_positions[0], sep_positions[1] # 句子A区间: [1, sep1) (假设索引0是[CLS]) # 句子B区间: [sep11, sep2) sent_a_end sep1 sent_b_start sep1 1 sent_b_end sep2 # 创建遮蔽不允许句子A看句子B也不允许句子B看句子A # 但允许句子内部互看也允许所有位置看[CLS]索引0[CLS]看所有位置根据任务设计 for i in range(seq_len): for j in range(seq_len): if i 0 or j 0: # [CLS]可以关注所有位置所有位置也可以关注[CLS]可选 continue # 如果i在句子A且j在句子B或者i在句子B且j在句子A则遮蔽 if (1 i sent_a_end and sent_b_start j sent_b_end) or \ (sent_b_start i sent_b_end and 1 j sent_a_end): mask[b, 0, i, j] True return mask # 测试句子对掩码使用我们模拟的第二个序列它包含[SEP] sp_mask sentence_pair_mask(input_ids) print(Sentence Pair Mask Shape:, sp_mask.shape) print(Mask for second seq (first head, first 6x6):\n, sp_mask[1, 0, :6, :6])4. 综合应用与模型集成示例现在我们将上述几种遮蔽技术整合到一个简化的Transformer编码器块中模拟一个可能的应用场景一个同时支持MLM预训练随机遮蔽和下游任务微调填充遮蔽的模型。class TransformerEncoderLayerWithFlexibleMask(nn.Module): 一个支持多种掩码的Transformer编码器层 def __init__(self, d_model, num_heads, dim_feedforward2048, dropout0.1): super().__init__() self.self_attn nn.MultiheadAttention(d_model, num_heads, dropoutdropout, batch_firstTrue) self.linear1 nn.Linear(d_model, dim_feedforward) self.dropout nn.Dropout(dropout) self.linear2 nn.Linear(dim_feedforward, d_model) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, src, src_maskNone, src_key_padding_maskNone, is_causalFalse): src: 输入序列 [batch, seq_len, d_model] src_mask: 自定义的注意力掩码 [seq_len, seq_len] 或 [batch, seq_len, seq_len] src_key_padding_mask: 填充掩码 [batch, seq_len]True/1表示需要被遮蔽PAD is_causal: 是否为因果掩码 # 自注意力 attn_output, _ self.self_attn( src, src, src, attn_masksrc_mask, key_padding_masksrc_key_padding_mask, is_causalis_causal ) src src self.dropout1(attn_output) src self.norm1(src) # 前馈网络 ff_output self.linear2(self.dropout(F.relu(self.linear1(src)))) src src self.dropout2(ff_output) src self.norm2(src) return src # 演示如何使用 encoder_layer TransformerEncoderLayerWithFlexibleMask(d_model, num_heads) # 场景1: 预训练MLM模式 (随机遮蔽 填充遮蔽) # 假设我们已经有了 masked_input (应用了随机遮蔽) 和 padding_mask # src_mask 为 None (不使用额外的结构化掩码) # is_causalFalse output_mlm encoder_layer( srcembedding(masked_input), src_key_padding_mask(attention_mask 0) # PyTorch的key_padding_mask: True位置被遮蔽 ) print(MLM Mode Output Shape:, output_mlm.shape) # 场景2: 类似GPT的解码模式 (因果遮蔽 填充遮蔽) # src_mask 可以传入一个因果掩码或者直接设置 is_causalTrue (PyTorch 1.9.0) causal_mask causal_attention_mask(seq_len) output_causal encoder_layer( srcembedding(input_ids), src_maskcausal_mask, # 传入因果掩码 src_key_padding_mask(attention_mask 0), is_causalFalse # 因为已经传入了src_mask这里设为False ) print(Causal Mode Output Shape:, output_causal.shape) # 场景3: 长文档处理模式 (滑动窗口遮蔽 填充遮蔽) window_mask sliding_window_mask(seq_len, window_size2).squeeze(0) # [1,1,L,L] - [L,L] output_window encoder_layer( srcembedding(input_ids), src_maskwindow_mask, src_key_padding_mask(attention_mask 0) ) print(Sliding Window Mode Output Shape:, output_window.shape)5. 常见问题、调试技巧与性能考量在实际编码和调试过程中你肯定会遇到各种问题。下面是我踩过的一些坑和总结的经验。5.1 掩码形状与广播机制这是最常见的错误来源。注意力掩码的形状必须能与注意力分数矩阵进行广播。PyTorchnn.MultiheadAttentionattn_mask期望的形状是(L, L)或(N*num_heads, L, L)其中N是batch size。key_padding_mask期望的形状是(N, L)。自定义实现为了支持多头我们通常将掩码扩展为[batch_size, 1, 1, seq_len, seq_len]或[batch_size, num_heads, seq_len, seq_len]以便与[batch_size, num_heads, seq_len, seq_len]的注意力分数相加。关键检查点始终打印你的掩码和注意力分数的形状确保它们能够正确广播。一个典型的模式是scores mask.unsqueeze(0).unsqueeze(1)。5.2 遮蔽值的选择-inf 还是 很大的负数在应用掩码到注意力分数后我们需要遮蔽的位置在Softmax后权重为0。通常有两种做法在Softmax前加一个极大的负数如-1e9这是最稳妥和常见的做法。-inf在理论上更干净但在某些硬件或软件环境下可能引发问题如NaN梯度。使用masked_fill并将掩码位置设为Truetorch.masked_fill(scores, mask, -1e9)。确保你的掩码是布尔类型。建议统一使用-1e9作为遮蔽值。在混合精度训练时这个值也足够大。5.3 填充遮蔽与损失函数的协同遮蔽有两个层面注意力层面和损失计算层面。注意力层面通过key_padding_mask防止模型关注[PAD]。损失层面在计算交叉熵损失时需要通过ignore_index参数忽略掉[PAD]位置对应的标签。通常[PAD]的token id是0所以设置loss_fn nn.CrossEntropyLoss(ignore_index0)。必须两者同时做否则模型仍会从[PAD]的位置产生输出并贡献损失干扰学习。5.4 因果遮蔽在训练与推理中的差异训练通常使用“教师强制”一次性输入完整目标序列并应用因果掩码。计算的是整个序列的并行损失。推理自回归生成每次生成一个token。此时因果掩码是动态变化的。更重要的是要利用键值缓存。在生成第t个token时前t-1个token的键值对已被缓存只需计算当前token的查询与缓存键的注意力并更新缓存。这能极大提升推理速度。实现提示在自定义注意力层中实现KV缓存需要维护两个缓存张量并在每次前向传播时进行拼接和切片。5.5 滑动窗口遮蔽的效率陷阱滑动窗口遮蔽虽然降低了理论复杂度但标准的Transformer实现如PyTorch的nn.Transformer可能无法利用这种稀疏性来实际加速计算因为底层的矩阵乘法仍然是稠密的。要获得真正的效率提升需要使用支持稀疏注意力或块状注意力如xformers库、torch.nn.functional.scaled_dot_product_attention的is_causal和自定义attn_mask的优化库。性能建议如果你的序列真的非常长2048考虑使用FlashAttention如果硬件支持或xformers库它们对长序列和特定掩码模式有深度优化。5.6 调试工具可视化注意力权重当模型行为不符合预期时可视化注意力权重是强大的调试手段。import matplotlib.pyplot as plt def plot_attention_weights(attention_weights, layer_idx0, head_idx0): attention_weights: 从模型中获取的注意力权重列表或张量。 假设形状为 [batch, num_heads, seq_len, seq_len] if isinstance(attention_weights, list): attn attention_weights[layer_idx] # 取某一层 else: attn attention_weights # 取第一个样本指定头 attn_map attn[0, head_idx].detach().cpu().numpy() plt.figure(figsize(10, 8)) plt.imshow(attn_map, cmapviridis, interpolationnearest) plt.colorbar() plt.title(fAttention Weights - Layer {layer_idx}, Head {head_idx}) plt.xlabel(Key Positions) plt.ylabel(Query Positions) plt.tight_layout() plt.show() # 假设你从模型的某个注意力层提取了权重 attn_weights # plot_attention_weights(attn_weights, layer_idx0, head_idx0)通过观察热力图你可以清晰地看到模型到底在关注哪些位置从而判断你的遮蔽是否生效以及模型的学习模式是否健康。例如在因果遮蔽下你应该看到一个严格的下三角热图如果出现了上三角的非零权重那就说明遮蔽出了问题。