PyTorch实现LLM令牌遮蔽:从因果掩码到Span遮蔽的5种核心技术

📅 2026/8/26 1:39:08
PyTorch实现LLM令牌遮蔽:从因果掩码到Span遮蔽的5种核心技术
1. 为什么我们需要“遮蔽”LLM的令牌在大型语言模型LLM的训练和应用中我们常常会遇到一个看似矛盾的需求模型需要处理和理解输入的全部信息但有时我们又希望它“看不见”或“忽略”某些部分。这种“选择性失明”的技术就是令牌遮蔽。它远不止是简单地把几个词涂黑那么简单而是一套精细控制模型注意力与信息流的工程艺术。想象一下你正在训练一个翻译模型。输入是“我喜欢吃苹果apple”目标是翻译成英文。在训练时模型需要学会将“苹果”与“apple”对应。但如果模型在编码端理解输入时就提前“偷看”到了解码端生成输出时的“apple”那它根本不需要学习任何复杂的语义映射直接抄答案就行了。这会导致模型无法泛化遇到新词就束手无策。为了防止这种“作弊”我们必须在训练时遮蔽掉解码器未来要生成的令牌迫使模型真正去学习语言间的内在联系。这就是因果遮蔽是Transformer解码器架构的核心。再比如我们做文本分类或情感分析。句子“这部电影的剧情很棒但特效实在太烂了。”包含转折。如果我们想增强模型对局部特征的鲁棒性或者进行数据增强可能会随机遮蔽掉“特效”或“很棒”让模型根据剩余上下文去预测被遮蔽的词或者判断整体情感。这能锻炼模型不依赖于某个特定关键词而是理解整体语义结构。这就是随机令牌遮蔽类似于NLP领域的“完形填空”。因此令牌遮蔽的核心目的至少有三个一是确保自回归生成过程的严谨性防止数据泄露二是作为一种强大的正则化手段提升模型的鲁棒性和泛化能力三是构造特定的预训练任务让模型学习更深层的语言表征。在PyTorch中实现这些遮蔽逻辑是我们驾驭LLM的必备技能。下面我将抛开理论教科书直接进入实战用代码和实例拆解5种最常用、也最具代表性的令牌遮蔽技术。2. 基石因果遮蔽与注意力掩码的实现因果遮蔽也叫前瞻遮蔽是Transformer解码器和自回归模型如GPT系列的“生命线”。它的规则很简单在生成第i个令牌时模型只能看到第1到第i-1个令牌而不能看到第i个及之后的令牌。在PyTorch中这并非通过修改输入数据实现而是通过一个关键的张量——注意力掩码来控制的。对于注意力机制中的QK^T计算我们通过掩码将非法位置的得分置为一个极大的负值如-1e9这样在后续的softmax操作中这些位置的权重就会趋近于0。实现一个通用的因果注意力掩码import torch def generate_causal_mask(seq_len, devicecpu): 生成一个下三角矩阵形式的因果掩码。 对角线及以下为0允许看以上为1遮蔽。 用于加在注意力分数上之前需要将1的位置转换为极小的负值。 # 创建一个上三角矩阵主对角线及以上为1以下为0 mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() # 我们希望被遮蔽的位置为True所以直接返回这个bool矩阵 # 使用时attention_scores.masked_fill_(mask, -1e9) return mask.to(device) # 示例序列长度为5 mask generate_causal_mask(5) print(“因果掩码True表示需要被遮蔽的位置”) print(mask)输出因果掩码True表示需要被遮蔽的位置 tensor([[False, True, True, True, True], [False, False, True, True, True], [False, False, False, True, True], [False, False, False, False, True], [False, False, False, False, False]])在多头自注意力层中的集成在实际的Transformer模块中我们需要处理批次batch和多头head维度。通常我们会生成一个形状为(1, 1, seq_len, seq_len)的掩码利用广播机制应用到所有批次和所有头上。class CausalSelfAttention(nn.Module): def __init__(self, embed_dim, num_heads): super().__init__() self.attn nn.MultiheadAttention(embed_dim, num_heads, batch_firstTrue) def forward(self, x): # x: (batch_size, seq_len, embed_dim) batch_size, seq_len, _ x.shape # 生成因果掩码 causal_mask generate_causal_mask(seq_len, devicex.device) # MultiheadAttention 期望的掩码形状是 (L, S) 或 (N*num_heads, L, S) # 这里我们使用 (L, S) 并让广播生效 attn_mask causal_mask # (seq_len, seq_len) # 前向传播传入注意力掩码 output, _ self.attn(x, x, x, attn_maskattn_mask) return output注意nn.MultiheadAttention的attn_mask参数如果是一个2D张量L, S它会被广播到所有批次和所有头。如果是一个3D张量N*num_heads, L, S则可以指定每个头不同的掩码。布尔类型的掩码中True位置会被遮蔽。也可以使用float类型-inf位置被遮蔽。一个容易踩的坑填充遮蔽与因果遮蔽的结合。在实际任务中我们还有填充遮蔽Padding Mask用于忽略序列中无效的填充位置pad。我们需要将两者合并。def create_combined_mask(tgt_seq, pad_token_id0): 为解码器创建结合了填充遮蔽和因果遮蔽的掩码。 tgt_seq: 目标序列令牌ID形状 (batch_size, seq_len) batch_size, seq_len tgt_seq.shape # 1. 创建填充掩码 (batch_size, 1, 1, seq_len) - 用于遮蔽key中的pad位置 pad_mask (tgt_seq pad_token_id).unsqueeze(1).unsqueeze(2) # 扩张维度以适配注意力头 # 2. 创建因果掩码 (1, 1, seq_len, seq_len) causal_mask generate_causal_mask(seq_len, devicetgt_seq.device).unsqueeze(0).unsqueeze(0) # 3. 合并任何位置如果是pad或者在未来都需要被遮蔽 # 这里我们生成一个最终的注意力掩码张量在计算注意力分数前加上 # 更常见的做法是在注意力函数内部分别处理。这里演示逻辑。 combined_mask pad_mask | causal_mask # 逻辑或运算 # 最终 combined_mask 形状: (batch_size, 1, seq_len, seq_len) # 其中 True 表示该位置需要被遮蔽 return combined_mask这里的合并逻辑是一个位置只要满足是“填充令牌”或者“未来的令牌”中的任意一个条件就应该被遮蔽。在实现时需要根据你使用的注意力层API来调整掩码的维度和类型。3. 预训练的引擎随机令牌遮蔽与BERT-style实现随机令牌遮蔽是BERT等掩码语言模型MLM预训练任务的基石。其核心思想是随机选择输入序列中一定比例如15%的令牌将其替换为特殊的[MASK]令牌然后训练模型根据上下文来预测这些被遮蔽的原始令牌。这个过程看似简单但实现细节直接影响预训练效果。BERT论文中采用了一种更精细的策略80%的时间用[MASK]替换选中的令牌。10%的时间用一个随机词表中的其他令牌替换。10%的时间保持原令牌不变。这样做的目的是为了缓解预训练大量见到[MASK]与微调从未见过[MASK]之间的不一致性让模型学会不仅依赖于“这里有个掩码”的提示而是真正理解上下文。PyTorch实现步骤详解import torch import torch.nn as nn import numpy as np class MLMDataCollator: 用于动态生成MLM训练样本的数据整理器。 在DataLoader中配合使用每个batch实时生成遮蔽。 def __init__(self, tokenizer, mlm_probability0.15): self.tokenizer tokenizer self.mlm_probability mlm_probability self.mask_token_id tokenizer.mask_token_id self.vocab_size len(tokenizer) def __call__(self, features): features: 一个batch的样本列表每个样本是包含‘input_ids’等的字典。 返回: 一个批处理后的字典包含遮蔽后的‘input_ids’和对应的‘labels’。 batch self.tokenizer.pad(features, return_tensors“pt”) input_ids batch[“input_ids”].clone() # 原始输入 labels batch[“input_ids”].clone() # 标签就是原始ID # 1. 创建遮蔽概率矩阵 probability_matrix torch.full(labels.shape, self.mlm_probability) # 特殊令牌如[CLS], [SEP], [PAD]不应该被遮蔽 special_tokens_mask [ self.tokenizer.get_special_tokens_mask(val, already_has_special_tokensTrue) for val in labels.tolist() ] special_tokens_mask torch.tensor(special_tokens_mask, dtypetorch.bool) probability_matrix.masked_fill_(special_tokens_mask, value0.0) # 2. 根据概率矩阵决定哪些位置被选中进行遮蔽 masked_indices torch.bernoulli(probability_matrix).bool() labels[~masked_indices] -100 # 只在被遮蔽的位置计算损失-100在CrossEntropyLoss中会被忽略 # 3. 处理被遮蔽的位置80% - [MASK], 10% - random, 10% - original indices_replaced torch.bernoulli(torch.full(labels.shape, 0.8)).bool() masked_indices input_ids[indices_replaced] self.mask_token_id indices_random torch.bernoulli(torch.full(labels.shape, 0.5)).bool() masked_indices ~indices_replaced random_words torch.randint(self.vocab_size, labels.shape, dtypetorch.long) input_ids[indices_random] random_words[indices_random] # indices_unchanged masked_indices ~indices_replaced ~indices_random (剩下的10%) # 这部分input_ids保持不变 batch[“input_ids”] input_ids batch[“labels”] labels return batch # 使用示例 from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(“bert-base-uncased”) collator MLMDataCollator(tokenizer) # 假设从数据集中加载了一个batch的原始文本 raw_texts [“Hello, world!”, “This is a test.”] features [tokenizer(text, truncationTrue, padding“max_length”, max_length10) for text in raw_texts] batch collator(features) print(“遮蔽后的 input_ids:”, batch[“input_ids”]) print(“对应的 labels (非-100处为需要预测的原始ID):”, batch[“labels”])关键点与避坑指南labels的设置损失函数如nn.CrossEntropyLoss通常通过设置ignore_index-100来忽略不需要计算损失的位置。因此我们将未被遮蔽位置的标签设为-100。special_tokens_mask必须排除[CLS]、[SEP]、[PAD]等特殊令牌被遮蔽否则会干扰模型对句子结构和边界的理解。tokenizer.get_special_tokens_mask方法可以方便地获取这个掩码。动态遮蔽 vs 静态遮蔽上述实现是“动态遮蔽”即每个epoch、每个batch的遮蔽模式都不同能极大增加数据多样性。与之相对的是“静态遮蔽”预处理时固定遮蔽模式动态遮蔽效果通常更好。性能考量在数据加载管道中实时进行遮蔽计算会带来少量开销。对于超大数据集确保你的数据加载器DataLoader的num_workers设置合理以避免成为训练瓶颈。4. 更高级的遮蔽策略N-gram遮蔽与Span遮蔽随机遮蔽单个令牌有时过于简单无法让模型学习到短语或实体级别的语义。于是更高级的遮蔽策略被提出。N-gram遮蔽不是随机选择单个令牌而是随机选择连续的n个令牌n-gram进行整体遮蔽。例如对于句子“The quick brown fox jumps”可能选择“quick brown”这个bigram2-gram一起遮蔽。这迫使模型基于更远的上下文来预测一个连续的片段提升了学习词组表征的能力。实现思路确定要遮蔽的令牌总数如序列长度的15%。从可能的n-gram长度如1, 2, 3中按一定权重分布随机选择一个长度n。在序列中随机选择一个起始位置i遮蔽[i:in]的令牌。重复步骤2-3直到被遮蔽的令牌总数达到目标。def create_ngram_mask(seq_len, target_mask_ratio0.15, max_ngram3): 创建一个n-gram遮蔽位置的布尔张量。 简化版不考虑特殊令牌和重叠仅演示逻辑。 mask torch.zeros(seq_len, dtypetorch.bool) target_masked_count int(seq_len * target_mask_ratio) masked_count 0 ngram_weights [0.7, 0.2, 0.1] # 假设1-gram概率70%2-gram 20%3-gram 10% ngram_dist torch.distributions.Categorical(torch.tensor(ngram_weights)) while masked_count target_masked_count: n ngram_dist.sample().item() 1 # 采样n-gram长度 (1,2,3) start torch.randint(0, max(1, seq_len - n), (1,)).item() # 随机起始位置 # 检查选中的位置是否已被遮蔽避免重叠这里简单跳过 if not mask[start:startn].any(): mask[start:startn] True masked_count n if masked_count target_masked_count: # 如果超出了需要解除最后一部分遮蔽 overflow masked_count - target_masked_count if overflow 0: # 简单处理将最后遮蔽的n-gram尾部取消遮蔽 mask[startn-overflow:startn] False break return maskSpan遮蔽如SpanBERT、T5这是N-gram遮蔽的推广和系统化。它不再固定n-gram长度而是先根据几何分布倾向于短span随机采样一个span的长度l然后随机选择起始位置遮蔽这个连续的span。SpanBERT发现同时预测整个被遮蔽span内所有令牌的边界起始和结束位置而不仅仅是逐个预测能显著提升模型对片段的理解能力。PyTorch Span遮蔽核心逻辑def create_span_mask(seq_len, target_mask_ratio0.15, max_span_length10, geometric_p0.2): 使用几何分布生成span遮蔽。 geometric_p: 几何分布参数p越小生成长span的概率越低。 from torch.distributions.geometric import Geometric mask torch.zeros(seq_len, dtypetorch.bool) target_masked_count int(seq_len * target_mask_ratio) masked_count 0 geometric_dist Geometric(probstorch.tensor(geometric_p)) while masked_count target_masked_count: # 采样span长度并限制最大值 span_length min(int(geometric_dist.sample().item()) 1, max_span_length) start_pos torch.randint(0, max(1, seq_len - span_length), (1,)).item() # 检查重叠生产环境需更复杂的处理逻辑 if not mask[start_pos:start_posspan_length].any(): mask[start_pos:start_posspan_length] True masked_count span_length if masked_count target_masked_count: overflow masked_count - target_masked_count if overflow 0: # 调整最后一个span的长度 new_span_length span_length - overflow mask[start_posnew_span_length:start_posspan_length] False break return mask选择与权衡随机令牌遮蔽实现简单计算高效是BERT的基础适合通用语料预训练。N-gram/Span遮蔽能更好地捕捉局部依赖和短语语义在需要理解实体、短语的任务如问答、指代消解上表现更优但实现稍复杂且可能因为长span的遮蔽增加训练难度。实际建议对于大多数从零开始的预训练可以从标准的BERT式随机遮蔽开始。如果你有领域特定的数据如生物医学文本其中长实体名词很重要尝试Span遮蔽可能会带来惊喜。在Hugging Face的DataCollatorForLanguageModeling中可以通过mlm_probability参数控制遮蔽比例但其默认是随机令牌遮蔽。要实现N-gram或Span遮蔽通常需要自定义数据整理器。5. 适配特定任务的动态遮蔽填充遮蔽与序列对遮蔽前面的遮蔽多用于预训练。在具体的下游任务微调或推理中我们还需要其他遮蔽技术来处理实际数据中的不规则性。填充遮蔽这是处理变长序列批处理时的标配。为了将不同长度的句子组成一个批次我们需要用特殊的[PAD]令牌将较短的句子填充到同一长度。在计算注意力时这些填充位置必须被遮蔽防止模型关注无意义的[PAD]。def create_padding_mask(seq, pad_token_id0): 创建填充掩码。 seq: 令牌ID张量形状 (batch_size, seq_len) 返回: 布尔张量True表示对应位置是pad需要被遮蔽。 padding_mask (seq pad_token_id) # 为了适配注意力机制我们通常将其扩展为 (batch_size, 1, 1, seq_len) # 这样在应用时可以广播到所有注意力头和所有query位置 return padding_mask.unsqueeze(1).unsqueeze(2) # 形状: (batch_size, 1, 1, seq_len) # 在注意力计算中的应用 def scaled_dot_product_attention_with_mask(Q, K, V, padding_maskNone): Q, K, V: 形状 (batch_size, num_heads, seq_len, head_dim) padding_mask: 形状 (batch_size, 1, 1, seq_len) 或 (batch_size, 1, seq_len, seq_len) d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k) # (batch_size, num_heads, seq_len, seq_len) if padding_mask is not None: # 将padding_mask调整为与scores相同的形状通常key的pad需要被遮蔽 # 假设padding_mask是针对key的 (batch_size, 1, 1, key_len) # 我们需要将其广播到 (batch_size, num_heads, query_len, key_len) # 更简单的方式是直接加到scores上 scores scores.masked_fill(padding_mask, float(‘-inf’)) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, V) return output, attn_weights序列对遮蔽如Encoder-Decoder架构在机器翻译、文本摘要等序列到序列任务中模型如BART、T5的编码器处理源序列解码器自回归地生成目标序列。这里的遮蔽更为复杂编码器自注意力对源序列使用双向注意力可以看到整个序列但同样需要填充遮蔽。解码器自注意力对目标序列使用因果遮蔽只能看过去和当前防止信息泄露同时也要填充遮蔽。编码器-解码器交叉注意力解码器在生成每一个目标令牌时可以关注整个编码器输出完整的源序列。这里只需要对编码器输出应用填充遮蔽如果源序列有填充的话。# 以自定义一个简化的Transformer Seq2Seq模型为例 class SimpleSeq2SeqTransformer(nn.Module): def __init__(self, encoder, decoder, src_pad_idx, tgt_pad_idx): super().__init__() self.encoder encoder self.decoder decoder self.src_pad_idx src_pad_idx self.tgt_pad_idx tgt_pad_idx def make_src_mask(self, src): # src: (batch_size, src_len) src_mask (src ! self.src_pad_idx).unsqueeze(1).unsqueeze(2) # (batch_size, 1, 1, src_len) return src_mask # True表示有效位置 def make_tgt_mask(self, tgt): # tgt: (batch_size, tgt_len) batch_size, tgt_len tgt.shape # 创建因果掩码 causal_mask torch.tril(torch.ones(tgt_len, tgt_len)).bool().to(tgt.device) # 下三角为True causal_mask causal_mask.unsqueeze(0).unsqueeze(0) # (1, 1, tgt_len, tgt_len) # 创建填充掩码 tgt_pad_mask (tgt ! self.tgt_pad_idx).unsqueeze(1).unsqueeze(2) # (batch_size, 1, 1, tgt_len) # 合并一个位置需要被遮蔽如果它是pad无效或者它是未来的令牌 # 注意causal_mask中False表示未来位置需要被遮蔽True表示允许看的位置 # 我们需要一个最终掩码其中True表示需要被遮蔽 tgt_mask ~(causal_mask tgt_pad_mask) # 逻辑稍微绕核心是只有“非pad且非未来”的位置才有效 # 更清晰的做法是分别处理在注意力函数中合并 return tgt_mask, tgt_pad_mask def forward(self, src, tgt): src_mask self.make_src_mask(src) tgt_mask, tgt_pad_mask self.make_tgt_mask(tgt) enc_src self.encoder(src, src_key_padding_mask~src_mask.squeeze(1).squeeze(1)) # 注意API对mask格式的要求可能不同 output self.decoder(tgt, enc_src, tgt_masktgt_mask, tgt_key_padding_mask~tgt_pad_mask.squeeze(1).squeeze(1)) return output重要提示不同的Transformer库如nn.Transformer,nn.MultiheadAttention, Hugging FaceTransformers对掩码的格式是布尔掩码还是加性掩码、含义True是遮蔽还是保留要求可能不同。上述代码是概念演示实际使用时务必查阅你所使用模块的官方文档。例如nn.Transformer的src_mask和tgt_mask是加性掩码-inf用于遮蔽而src_key_padding_mask和tgt_key_padding_mask是布尔掩码True表示需要被遮蔽的key。6. 实战整合在自定义训练循环中应用多种遮蔽理解了各种遮蔽的原理和独立实现后我们需要将它们整合到一个真实的模型训练循环中。下面我将演示如何在一个简化的、基于Hugging Facetransformers库的BERT掩码语言模型微调任务中结合填充遮蔽和MLM遮蔽。场景我们有一个分类任务但想通过继续MLM预训练来领域适应。我们将使用动态MLM遮蔽并在模型前向传播时处理填充遮蔽。import torch from torch.utils.data import DataLoader from transformers import AutoTokenizer, AutoModelForMaskedLM, AdamW from transformers import DataCollatorForLanguageModeling # 1. 准备数据和分词器 tokenizer AutoTokenizer.from_pretrained(“bert-base-uncased”) model AutoModelForMaskedLM.from_pretrained(“bert-base-uncased”) # 假设我们有一个文本列表 domain_texts from datasets import Dataset dataset Dataset.from_dict({“text”: domain_texts}) def tokenize_function(examples): return tokenizer(examples[“text”], truncationTrue, padding“max_length”, max_length128) tokenized_dataset dataset.map(tokenize_function, batchedTrue, remove_columns[“text”]) # 2. 使用内置的动态MLM数据整理器 # 它会自动处理80/10/10的遮蔽策略和标签生成 data_collator DataCollatorForLanguageModeling( tokenizertokenizer, mlmTrue, mlm_probability0.15 ) train_dataloader DataLoader(tokenized_dataset, batch_size16, shuffleTrue, collate_fndata_collator) # 3. 训练循环 device torch.device(“cuda”) if torch.cuda.is_available() else torch.device(“cpu”) model.to(device) optimizer AdamW(model.parameters(), lr5e-5) model.train() for epoch in range(3): total_loss 0 for batch_idx, batch in enumerate(train_dataloader): # 将数据移动到设备 batch {k: v.to(device) for k, v in batch.items()} # 前向传播 outputs model(**batch) loss outputs.loss # 反向传播 loss.backward() optimizer.step() optimizer.zero_grad() total_loss loss.item() if batch_idx % 100 0: print(f“Epoch {epoch}, Batch {batch_idx}, Loss: {loss.item():.4f}”) print(f“Epoch {epoch} Average Loss: {total_loss / len(train_dataloader):.4f}”)在这个流程中遮蔽是如何工作的动态MLM遮蔽DataCollatorForLanguageModeling在每个batch加载时实时调用torch.bernoulli等函数按照15%的比例和80/10/10的策略生成遮蔽模式并准备好labels。填充遮蔽AutoModelForMaskedLM内部的BERT模型在其注意力层中会自动处理填充遮蔽。它是通过attention_mask参数实现的。在我们使用tokenizer(..., padding“max_length”)时生成的attention_mask就已经标识了哪些是真实令牌1哪些是填充令牌0。模型的前向传播逻辑会利用这个attention_mask来生成最终的注意力掩码遮蔽掉填充位置。模型内部模型接收input_ids可能包含[MASK]、attention_mask和labels。它首先将input_ids转换为嵌入向量然后在每一层Transformer中计算注意力分数时会组合自注意力掩码对于BERT是双向的所以每个位置都可以看到所有其他非填充位置。这通过attention_mask实现将填充位置对应的分数设为负无穷。MLM任务最后的输出层对每个位置进行词汇表分类但损失只计算labels非-100的位置即被遮蔽的位置。如果你想手动控制或实现更复杂的遮蔽逻辑比如加入Span遮蔽你可以继承DataCollatorForLanguageModeling并重写其__call__方法用我们前面写的create_span_mask函数替换其内部的随机遮蔽逻辑。然后在创建DataLoader时使用你这个自定义的collator。7. 遮蔽技术的高级应用与避坑经验掌握了基础遮蔽后我们来看看一些进阶场景和容易出错的地方。1. 处理批次中不同长度的序列变长序列这是填充遮蔽的主要用武之地。但要注意使用过长的max_length会浪费大量计算在无效的填充令牌上。解决方案是使用动态填充。from transformers import DataCollatorWithPadding # 在tokenize时不要padding def tokenize_function_no_pad(examples): return tokenizer(examples[“text”], truncationTrue) # 移除了 padding“max_length” tokenized_dataset dataset.map(tokenize_function_no_pad, batchedTrue, remove_columns[“text”]) # 使用DataCollatorWithPadding进行动态填充到批次内最大长度 data_collator DataCollatorWithPadding(tokenizer, padding“longest”) train_dataloader DataLoader(tokenized_dataset, batch_size16, shuffleTrue, collate_fndata_collator)DataCollatorWithPadding会收集一个batch内的所有样本然后将其填充到该batch内最长的序列长度而不是一个固定的全局最大长度这能显著提升训练效率尤其是当序列长度差异较大时。2. 解码器推理时的缓存与掩码在自回归生成如GPT文本生成时为了效率我们会缓存之前时间步的键值对KV Cache。这时因果掩码需要与缓存正确配合。每次生成新令牌时我们只计算新令牌对于所有历史令牌包括缓存的注意力。# 伪代码演示推理时因果掩码的增量生成 past_key_values None # 初始化缓存 generated [start_token_id] for step in range(max_length): # 准备当前步的输入通常是上一步生成的令牌 input_ids torch.tensor([generated[-1]]).unsqueeze(0).to(device) # (1, 1) # 注意力掩码需要看到所有已生成的令牌 # 假设past_key_values不为None说明已有历史 if past_key_values is not None: attention_mask torch.ones(1, len(generated)).to(device) # 所有位置都可见 else: attention_mask torch.ones(1, 1).to(device) # 模型前向传入past_key_values和attention_mask outputs model(input_idsinput_ids, attention_maskattention_mask, past_key_valuespast_key_values, use_cacheTrue) next_token_logits outputs.logits[:, -1, :] next_token_id torch.argmax(next_token_logits, dim-1).item() generated.append(next_token_id) past_key_values outputs.past_key_values # 更新缓存关键在于每一步的attention_mask都是一个全1的向量长度等于当前已生成序列的长度包括当前新令牌。模型内部的因果自注意力机制会通过缓存和当前输入自动处理“只能看过去”的逻辑。现代的Transformer库如transformers的generate()方法已经完美封装了这些细节。3. 一个常见的坑掩码张量的设备不一致在PyTorch中确保你的掩码张量和模型/输入数据在同一个设备上CPU或GPU。# 错误示例掩码在CPU数据在GPU input_ids input_ids.to(‘cuda’) attention_mask attention_mask # 还在CPU上 outputs model(input_ids, attention_maskattention_mask) # 会报错 # 正确做法 attention_mask attention_mask.to(input_ids.device) outputs model(input_ids, attention_maskattention_mask)4. 理解掩码的广播规则注意力分数矩阵的形状是(batch_size, num_heads, query_len, key_len)。你的掩码需要能广播到这个形状。形状为(batch_size, 1, 1, key_len)的掩码可以广播对所有query和所有head遮蔽相同的key位置。常用于填充掩码。形状为(batch_size, 1, query_len, key_len)的掩码可以对每个批次、每个query位置指定不同的遮蔽模式。形状为(batch_size, num_heads, query_len, key_len)的掩码最灵活可以为每个注意力头指定不同的遮蔽模式如定制化注意力。务必根据你的库如nn.Transformer,nn.MultiheadAttention, Hugging Face的文档要求来准备正确形状和类型的掩码。令牌遮蔽是LLM工程中的“沉默的守护者”它默默地在数据流中划定了可见与不可见的边界确保了模型的正确学习与推理。从最基本的因果遮蔽到预训练核心的随机遮蔽再到提升模型理解力的Span遮蔽以及处理实际数据的填充遮蔽每一种技术都对应着模型能力的一个关键拼图。在PyTorch中实现它们关键在于理解注意力机制中QK^T矩阵与掩码相加或填充这一核心操作并小心处理张量的形状、设备和广播规则。当你下次调试模型发现效果不佳时不妨检查一下你的掩码——很可能问题的答案就藏在那些被设置为-inf或True的矩阵元素之中。