拓冰建站拓冰建站
首页 / 资讯中心 / 正文

LLM令牌遮蔽技术详解:从因果掩码到滑动窗口的PyTorch实践

1. 项目概述为什么LLM需要“看不见”某些词在大型语言模型LLM的训练和应用中我们常常希望模型能“选择性失明”——不是真的看不见而是有策略地忽略输入序列中的某些部分。这种技术就是令牌遮蔽Token Masking。乍一听你可能觉得这和BERT等模型预训练时的掩码语言模型MLM任务很像都是为了学习上下文。但实际上在LLM特指GPT这类自回归模型的语境下令牌遮蔽的应用场景、技术目标和实现方式都更加多样和精细。简单来说LLM的令牌遮蔽核心目标是控制模型的注意力范围。想象一下你正在写一篇长文但规定自己不能回头看已经写过的某个段落或者被要求必须忽略文章里的所有数字。这种限制会迫使你改变思考和生成的方式。对LLM而言遮蔽技术就是施加这种限制的工具。它能让模型在生成下一个词时不去“注意”某些特定的、我们不想让它参考的令牌Token比如敏感信息、无关上下文、或者答案本身。这直接关系到模型的安全性、可控性、推理能力以及训练效率。从搜索热词来看大家关注点很集中LLM框架如如何搭建、使用、具体实现尤其是PyTorch、以及高级应用如RAG、Agent。这反映出社区已经从“惊叹模型能力”进入到“深入控制与优化模型行为”的实践阶段。掌握几种核心的遮蔽技术就等于拿到了精细调控LLM的“手术刀”。接下来我会结合PyTorch拆解五种最常用、也最具代表性的令牌遮蔽技术不仅告诉你“怎么做”更重点剖析“为什么这么做”以及“实践中会遇到什么坑”。2. 核心需求与场景解析不止于预训练在深入代码之前我们必须厘清在LLM的哪些环节我们需要动用遮蔽技术这决定了我们选择哪种方法。2.1 训练阶段的因果注意力遮蔽这是LLM训练的基石。在标准的自回归语言模型训练中比如训练一个GPT我们必须确保模型在预测位置i的令牌时只能看到位置0到i-1的令牌而不能“偷看”未来的信息。这就是因果遮蔽。它通过一个下三角矩阵对角线及以下为1以上为负无穷来实现强制注意力机制具有因果性。没有它模型就失去了“预测未来”的意义因为答案已经摆在眼前了。2.2 推理与部署阶段的输入控制模型训练好后在推理时我们同样需要遮蔽。例如防止信息泄露在对话系统中当模型生成回复时我们不应让它看到自己即将生成的回复内容。这通常通过动态扩展的因果掩码来实现。上下文管理在处理超长文本时超出模型上下文窗口我们需要有策略地遮蔽掉一部分历史信息比如滑动窗口遮蔽只让模型关注最近N个令牌。安全与合规主动遮蔽输入中的敏感词或特定实体防止模型基于这些信息生成不当内容。2.3 高级任务中的结构化遮蔽对于一些复杂任务遮蔽模式不再是简单的三角或窗口而是根据任务结构定制。填充生成在文本摘要、翻译等任务中我们可能先遮蔽掉待生成的部分让模型根据上下文进行填充。思维链提示为了引导模型进行分步推理我们可能会在提示中遮蔽中间推理步骤让模型自己推导出来。多模态对齐在视觉-语言模型中可能需要遮蔽掉文本序列中的某些视觉标记以研究模态间的依赖关系。理解这些场景后我们就能明白遮蔽不是一个单一的“开关”而是一套用于塑造模型信息流的“模具”。下面我们就从最基础的开始用PyTorch逐一实现。3. 五种核心令牌遮蔽技术详解与PyTorch实现在PyTorch中遮蔽的核心是操作一个与注意力权重矩阵形状相同的布尔张量或浮点数张量mask其中被遮蔽的位置需要忽略的位置通常被设置为True或一个极大的负值如-1e9然后在计算注意力分数后将这个掩码加到分数上scores scores mask。被遮蔽位置的分数经过Softmax后会趋近于0从而使其对应的注意力权重为0。我们先定义一个通用的注意力函数以便后续演示import torch import torch.nn.functional as F def scaled_dot_product_attention(query, key, value, maskNone): 计算缩放点积注意力。 Args: query: [batch_size, num_heads, seq_len_q, head_dim] key: [batch_size, num_heads, seq_len_k, head_dim] value: [batch_size, num_heads, seq_len_v, head_dim] mask: [batch_size, num_heads, seq_len_q, seq_len_k] 或广播兼容的形状。 True/1表示需要遮蔽的位置。 Returns: 注意力输出注意力权重 d_k query.size(-1) scores torch.matmul(query, key.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtypetorch.float32)) if mask is not None: # 将布尔掩码转换为分数掩码True的位置置为负无穷 scores scores.masked_fill(mask, float(-inf)) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, value) return output, attn_weights3.1 基础技术因果遮蔽这是自回归模型的命脉确保生成过程是单向的。原理与实现 因果掩码是一个下三角矩阵形状为[1, 1, target_len, source_len]。对于目标序列中的每个位置i它只能看到源序列中j i的位置。在训练时target_len和source_len通常是相等的都是输入序列的长度。def create_causal_mask(seq_len, devicecpu): 创建因果遮蔽矩阵。 Args: seq_len: 序列长度 Returns: mask: [1, 1, seq_len, seq_len], 下三角为False不遮蔽上三角为True遮蔽 # 创建一个上三角矩阵不包括对角线作为需要遮蔽的区域 mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() # 调整维度以适配注意力头 [1, 1, seq_len, seq_len] mask mask.unsqueeze(0).unsqueeze(0) return mask.to(device) # 使用示例 batch_size 2 num_heads 4 seq_len 10 head_dim 16 query torch.randn(batch_size, num_heads, seq_len, head_dim) key value query # 简化示例 causal_mask create_causal_mask(seq_len, query.device) # 注意这个掩码对所有批次和头都是相同的所以可以广播 output, attn_weights scaled_dot_product_attention(query, key, value, causal_mask) # 验证查看第一个批次第一个头最后一个查询位置的注意力权重 print(attn_weights[0, 0, -1, :]) # 输出应该只有最后几个位置有非零值因为因果遮蔽前面的位置权重应为0。实操心得与注意事项对角线问题torch.triu(..., diagonal1)中的diagonal1是关键。它确保了对角线元素当前位置自身不被遮蔽。在一些严格的定义中预测当前位置时也不应看到当前位置的key所以使用diagonal1。如果你希望包含自身则使用diagonal0但这在标准的自回归预测中不常见。广播机制我们创建的掩码形状是[1, 1, seq_len, seq_len]。当与形状为[batch_size, num_heads, seq_len, seq_len]的注意力分数相加时PyTorch的广播机制会自动将其扩展非常高效。避免为每个批次和头都创建独立的掩码张量。推理时的动态掩码在自回归生成如GPT推理时序列是逐步生成的。常见的做法是缓存之前时间步的键值对KV Cache并为当前新生成的令牌计算注意力。此时掩码需要动态扩展。通常我们会维护一个全局的因果掩码随着生成步骤增加一列一行新的行对应新的查询新的列对应新的键。3.2 实用技术填充遮蔽在实际任务中批次内的序列长度往往不一致。我们需要将短序列填充到同一长度并在计算注意力时遮蔽这些填充位置防止模型从无意义的填充符中学习。原理与实现 填充掩码通常是一个二维张量[batch_size, seq_len]指示哪些位置是真实的令牌False哪些是填充的True。在注意力中我们需要将其扩展为四维[batch_size, 1, 1, seq_len]这样对于每个查询位置都会遮蔽所有填充的键位置。def create_padding_mask(padding_indices, seq_len): 根据填充索引创建掩码。 Args: padding_indices: 一个列表的列表每个内层列表是一个序列中填充位置的索引。 例如batch_size2, seq_len5: [[3,4], [4]] 表示第一个序列索引3,4是填充第二个序列索引4是填充。 seq_len: 序列长度 Returns: mask: [batch_size, 1, 1, seq_len], True对应填充位置。 batch_size len(padding_indices) mask torch.zeros(batch_size, seq_len, dtypetorch.bool) for i, indices in enumerate(padding_indices): if indices: mask[i, indices] True # 调整维度: [batch_size, 1, 1, seq_len] # 这样对于该批次每个序列的所有查询都会遮蔽这些填充键。 mask mask.unsqueeze(1).unsqueeze(2) return mask # 更常见的做法是从tokenizer的attention_mask生成 def create_padding_mask_from_attention_mask(attention_mask): 从标准的attention_mask1表示真实token0表示填充生成用于遮蔽的mask。 Args: attention_mask: [batch_size, seq_len], 1为真实token0为填充。 Returns: mask: [batch_size, 1, 1, seq_len], True对应填充位置。 # 将attention_mask反转真实token为0不遮蔽填充为1遮蔽 mask (attention_mask 0) mask mask.unsqueeze(1).unsqueeze(2) return mask # 使用示例 attention_mask torch.tensor([[1, 1, 1, 0, 0], # 序列1后两个位置是填充 [1, 1, 1, 1, 0]]) # 序列2最后一个位置是填充 padding_mask create_padding_mask_from_attention_mask(attention_mask) print(padding_mask.shape) # torch.Size([2, 1, 1, 5]) print(padding_mask) # 输出第一个序列的索引3,4为True第二个序列的索引4为True。组合使用在训练中我们通常需要同时应用因果遮蔽和填充遮蔽。实现方式是将两个掩码用逻辑或|合并。def create_combined_mask(seq_len, attention_mask, devicecpu): 创建用于训练自回归模型的组合掩码因果填充。 causal_mask create_causal_mask(seq_len, device) # [1,1,seq_len,seq_len] padding_mask create_padding_mask_from_attention_mask(attention_mask).to(device) # [batch,1,1,seq_len] # 扩展因果掩码以匹配批次大小如果需要 if causal_mask.size(0) ! padding_mask.size(0): # 通常因果掩码是[1,1,...]可以直接广播。但为了逻辑清晰我们显式扩展。 causal_mask causal_mask.expand(padding_mask.size(0), -1, -1, -1) # 合并如果一个位置是填充True或者是未来的tokenTrue则遮蔽。 # 注意因果掩码中True表示未来需遮蔽填充掩码中True表示填充需遮蔽。 combined_mask causal_mask | padding_mask return combined_mask避坑指南掩码类型确保你的掩码是布尔类型torch.bool。如果使用float(‘-inf’)初始化的张量直接使用逻辑运算可能会出错。维度对齐padding_mask的形状是[batch_size, 1, 1, seq_len]这意味着它对一个批次内所有序列的所有查询位置遮蔽的键位置是相同的即所有填充位置。这是正确的因为无论查询位置在哪都不应该去关注填充符。解码器的交叉注意力在Seq2Seq模型如T5、BART的解码器中除了自注意力的因果掩码其交叉注意力关注编码器输出通常只需要填充掩码而不需要因果掩码因为解码器可以同时看到编码器的所有输出。3.3 进阶技术滑动窗口注意力遮蔽对于超长序列完全的自注意力计算复杂度是序列长度的平方O(n²)无法承受。滑动窗口注意力限制每个令牌只能关注其前后一定窗口大小内的令牌将复杂度降至O(n * w)其中w是窗口大小。这是像Longformer、BigBird等模型处理长文本的核心。原理与实现 为序列中的每个位置i创建一个掩码仅允许关注位置在[i - window_size, i window_size]范围内的j同时还要结合因果限制在自回归模型中j不能大于i。def create_sliding_window_mask(seq_len, window_size, is_causalTrue, devicecpu): 创建滑动窗口注意力掩码。 Args: seq_len: 序列长度 window_size: 单侧窗口大小。实际关注范围为 [i-window_size, iwindow_size]如果非因果。 is_causal: 是否为因果自回归模型。如果是则j不能i。 Returns: mask: [1, 1, seq_len, seq_len], True表示需要遮蔽的位置。 # 创建全1矩阵然后挖出窗口内的区域设为0/False mask torch.ones(seq_len, seq_len, dtypetorch.bool) for i in range(seq_len): start max(0, i - window_size) end i window_size 1 if not is_causal else i 1 # 因果模式下不能看未来 end min(end, seq_len) mask[i, start:end] False # 如果是因果的还需要确保上三角ji被遮蔽上面循环中的endi1已经实现了这一点。 # 但为了通用性我们可以显式地应用一个因果掩码。 if is_causal: causal_part torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() mask mask | causal_part # 合并窗口外或未来的位置都遮蔽 mask mask.unsqueeze(0).unsqueeze(0).to(device) return mask # 使用示例 seq_len 15 window_size 3 sw_mask create_sliding_window_mask(seq_len, window_size, is_causalTrue) print(sw_mask[0,0]) # 查看掩码矩阵 # 你会发现它是一个带状矩阵只有主对角线附近的带状区域和左下角是False可关注。性能与实现考量高效实现上述循环实现仅用于演示。在实际模型中如Longformer会使用更高效的带状矩阵乘法或自定义CUDA内核来实现滑动窗口注意力避免构建巨大的显式掩码矩阵尤其当seq_len很大时。全局注意力许多滑动窗口注意力模型会为某些特殊位置如序列开头[CLS]、结尾[SEP]或用户指定的位置添加“全局注意力”让这些位置可以看到整个序列反之亦然。这需要更复杂的掩码逻辑。结合填充同样需要与填充掩码结合确保模型不关注填充符。3.4 高级技术随机令牌遮蔽这直接来源于BERT的MLM任务但在LLM训练中也有其用途例如用于数据增强、提高模型鲁棒性或在T5等“编码器-解码器”架构的预训练中。原理与实现 随机选择输入序列中一定比例如15%的令牌将其替换为一个特殊的[MASK]令牌或者随机替换为其他令牌然后训练模型预测被遮蔽的原始令牌。在LLM的自回归训练中我们需要小心地整合这种遮蔽因为标准的自回归损失是预测下一个令牌而MLM是预测被遮蔽的任意位置。def random_token_masking(input_ids, mask_token_id, vocab_size, mask_prob0.15, replace_prob0.1, random_token_prob0.1): 对输入序列进行随机令牌遮蔽遵循BERT的原始策略。 Args: input_ids: [batch_size, seq_len] mask_token_id: 用于替换的[MASK]令牌的id vocab_size: 词表大小用于随机令牌替换 mask_prob: 被选中进行遮蔽的令牌比例 replace_prob: 在被选中的令牌中有多大比例直接替换为[MASK] random_token_prob: 在被选中的令牌中有多大比例替换为随机令牌 Returns: masked_input_ids: 遮蔽后的输入 labels: 用于计算损失的真实标签未被遮蔽的位置通常设为-100在CrossEntropyLoss中忽略 labels input_ids.clone() # 创建概率矩阵 prob_matrix torch.full_like(input_ids, float(mask_prob), dtypetorch.float32) # 决定哪些位置被选中进行遮蔽操作 masked_indices torch.bernoulli(prob_matrix).bool() # 确保特殊令牌如[CLS], [SEP]不被遮蔽这里假设pad_id0也需要排除 # 假设pad_id0, cls_id101, sep_id102 (根据具体tokenizer调整) special_tokens_mask (input_ids 0) | (input_ids 101) | (input_ids 102) masked_indices.masked_fill_(special_tokens_mask, False) # 将被选中的位置在labels中保持不变用于计算损失在input_ids中进行替换 labels[~masked_indices] -100 # 忽略未被遮蔽位置的损失 # 对于被遮蔽的位置决定是替换为[MASK]、随机词还是保持不变 replace_mask torch.bernoulli(torch.full_like(input_ids, float(replace_prob))).bool() masked_indices random_mask torch.bernoulli(torch.full_like(input_ids, float(random_token_prob))).bool() masked_indices ~replace_mask keep_mask masked_indices ~replace_mask ~random_mask # 执行替换 masked_input_ids input_ids.clone() masked_input_ids[replace_mask] mask_token_id # 替换为[MASK] masked_input_ids[random_mask] torch.randint(5, vocab_size-5, random_mask.shape, deviceinput_ids.device) # 替换为随机令牌避免极特殊id # keep_mask的位置input_ids保持不变 return masked_input_ids, labels # 使用示例假设一个简单的环境 batch_size 4 seq_len 20 vocab_size 30522 mask_token_id 103 input_ids torch.randint(100, 30000, (batch_size, seq_len)) masked_input, labels random_token_masking(input_ids, mask_token_id, vocab_size) print(原始输入片段:, input_ids[0, :8]) print(遮蔽后输入:, masked_input[0, :8]) print(损失标签 (忽略-100):, labels[0, :8]) # 可以看到部分位置被替换为103([MASK])或随机idlabels中对应位置为原始id其余为-100。在自回归LLM中的整合 对于纯解码器LLM直接应用上述MLM会破坏自回归性质。一种变通方法是“前缀语言模型”或“跨度遮蔽”随机遮蔽一个连续的令牌跨度然后让模型自回归地预测这个跨度内的所有令牌。这需要更复杂的掩码生成和损失计算逻辑。3.5 策略性技术定制化模式遮蔽这是最灵活的一类掩码模式完全由下游任务定义。例如在文本填充任务中我们遮蔽掉句子中间的一段在文档级翻译中为了保持段落连贯性可能遮蔽其他段落的信息。原理与实现 核心是根据任务规则生成一个二进制矩阵。这里以实现一个简单的“文本中间挖空”任务为例。def create_span_masking_mask(seq_len, mask_span_start, mask_span_length, devicecpu): 创建遮蔽一个连续区间的掩码。 Args: seq_len: 序列长度 mask_span_start: 遮蔽区间开始索引 mask_span_length: 遮蔽区间长度 Returns: mask: [1, 1, seq_len, seq_len], True表示需要遮蔽的位置。 这个掩码用于自注意力确保对于所有查询位置被遮蔽的键位置都不可见。 但更常见的做法是直接修改输入将对应位置替换为[MASK]并调整损失函数。 这里展示如何生成一个注意力掩码来“阻止”关注该区间。 mask torch.zeros(seq_len, seq_len, dtypetorch.bool) # 我们想遮蔽掉键Key中位于[mask_span_start, mask_span_startmask_span_length)的位置 # 对于任何查询Query都不应关注这些键。 mask[:, mask_span_start:mask_span_startmask_span_length] True # 如果是因果模型还需要叠加因果掩码 causal_mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() mask mask | causal_mask mask mask.unsqueeze(0).unsqueeze(0).to(device) return mask # 更实用的直接生成用于损坏输入的掩码和标签 def create_span_masking_inputs(input_ids, mask_span_start, mask_span_length, mask_token_id): 对输入进行区间遮蔽生成损坏的输入和标签。 Args: input_ids: [batch_size, seq_len] mask_span_start: 开始索引列表长度为batch_size mask_span_length: 遮蔽长度列表长度为batch_size mask_token_id: [MASK] token id Returns: masked_input_ids: 遮蔽后的输入 labels: 标签遮蔽区间外为-100 batch_size, seq_len input_ids.shape masked_input_ids input_ids.clone() labels torch.full_like(input_ids, -100) # 默认全部忽略 for i in range(batch_size): start mask_span_start[i] length mask_span_length[i] end min(start length, seq_len) # 保存原始标签 labels[i, start:end] input_ids[i, start:end] # 将输入中的该区间替换为[MASK] masked_input_ids[i, start:end] mask_token_id return masked_input_ids, labels应用场景 这种定制化遮蔽是构建指令微调或特定任务适配数据的关键。例如为了训练模型完成“根据上文填充下文”的任务我们可以随机选择文章的一个位置进行遮蔽。在RAG系统中当模型生成答案时我们可以遮蔽掉检索到的文档中的某些无关段落迫使模型更依赖核心证据。4. 综合应用与避坑实战理解了单项技术如何将它们组合并应用到真实的LLM训练或推理流程中这里以构建一个简单的、支持因果和填充遮蔽的自回归模型训练批次为例。4.1 完整训练步骤中的掩码集成假设我们使用Hugging Face的Transformers库中的GPT-2模型。from transformers import AutoTokenizer, AutoModelForCausalLM import torch model_name gpt2 tokenizer AutoTokenizer.from_pretrained(model_name) model AutoModelForCausalLM.from_pretrained(model_name) # 设置pad_token如果tokenizer没有的话 if tokenizer.pad_token is None: tokenizer.pad_token tokenizer.eos_token # 准备一个批次数据 texts [Hello, how are you?, Im fine, thank you. And you?] inputs tokenizer(texts, return_tensorspt, paddingTrue, truncationTrue, max_length10) input_ids inputs[input_ids] attention_mask inputs[attention_mask] # 这是标准的attention mask (1 for real tokens) # 模型内部已经实现了因果遮蔽。我们只需要传入attention_mask。 # 在transformers库中attention_mask会自动被转换为模型需要的格式。 outputs model(input_ids, attention_maskattention_mask, labelsinput_ids) # 使用labels进行训练 loss outputs.loss关键点transformers库的模型内部已经集成了因果逻辑。我们提供的attention_mask主要用于处理填充。模型内部的forward方法会调用_prepare_decoder_attention_mask函数来合并因果掩码和传入的填充掩码。4.2 自定义注意力中的掩码处理如果你想在自己的注意力层中实现这些掩码一个完整的多头注意力模块可能如下import torch.nn as nn import math class MultiHeadAttentionWithMasking(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0 self.d_model d_model self.num_heads num_heads self.head_dim d_model // num_heads self.wq nn.Linear(d_model, d_model) self.wk nn.Linear(d_model, d_model) self.wv nn.Linear(d_model, d_model) self.wo nn.Linear(d_model, d_model) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 1. 线性投影并分头 Q self.wq(query).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) K self.wk(key).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) V self.wv(value).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) # 2. 计算缩放点积注意力 scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim) if mask is not None: # mask形状应为 [batch_size, 1, 1, seq_len] 或 [batch_size, 1, seq_len_q, seq_len_k] # 确保mask能广播到scores的形状 scores scores.masked_fill(mask, float(-1e9)) # 使用负无穷遮蔽 attn_weights F.softmax(scores, dim-1) context torch.matmul(attn_weights, V) # 3. 合并多头并输出 context context.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) output self.wo(context) return output, attn_weights4.3 常见陷阱与调试技巧掩码形状错误这是最常见的问题。务必记住注意力分数的形状是[batch_size, num_heads, seq_len_q, seq_len_k]。你的掩码必须能广播成这个形状。padding_mask通常为[batch_size, 1, 1, seq_len_k]causal_mask为[1, 1, seq_len_q, seq_len_k]。使用mask.unsqueeze(1).unsqueeze(2)是增加维度的常用技巧。遮蔽值的选择被遮蔽的位置在加到注意力分数上时应设置为一个很大的负数。float(‘-inf’)是理论上的选择但在某些硬件或软件环境下可能不稳定。通常使用-1e9或-1e4等足够大的负数。在应用Softmax之前添加掩码。梯度问题被遮蔽的位置由于Softmax后权重为0其梯度也为0。这通常是我们期望的。但要确保你的掩码操作不会意外地阻断需要梯度的路径。验证掩码有效性编写简单的测试用例来验证掩码是否正确工作。def test_causal_mask(): seq_len 5 mask create_causal_mask(seq_len) print(Causal Mask (True掩碼):) print(mask[0,0]) # 模拟一个均匀分数 scores torch.zeros(1,1,seq_len,seq_len) scores_masked scores.masked_fill(mask, float(-inf)) attn F.softmax(scores_masked, dim-1) print(注意力权重最后一行:) print(attn[0,0,-1]) # 应该只有最后一个元素是1因为前面都是-inf验证因果性。与KV Cache的配合在自回归生成中KV Cache极大地提升了效率。此时因果掩码需要动态更新。通常做法是维护一个全局的掩码每次生成新token时在右侧和下侧扩展一行一列新token不能看到未来的key未来的query也不能看到它。5. 性能优化与高级话题当序列长度很长时显式的全尺寸掩码矩阵[seq_len, seq_len]会消耗大量内存。例如seq_len8192时一个布尔掩码矩阵就要占用 8192*8192/8 ≈ 8MB 内存如果是float则更大。对于批处理和多头这个开销会倍增。优化策略使用带状矩阵对于滑动窗口这类稀疏掩码可以使用PyTorch的带状矩阵函数如torch.band或稀疏张量来隐式表示。Flash Attention现代高效的注意力实现如Flash Attention将掩码计算融合到核函数中避免了在HBM显存中实例化庞大的中间矩阵包括掩码矩阵。如果你的模型支持优先使用这些优化后的注意力实现。自定义CUDA内核对于极其复杂的掩码模式如BigBird中的随机窗口全局注意力可能需要编写自定义的CUDA内核来实现高效的掩码计算和注意力。选择哪种技术训练标准自回归LM因果遮蔽 填充遮蔽是标配。处理长文本考虑滑动窗口遮蔽如Longformer或其变种。数据增强或特定预训练可尝试随机令牌遮蔽但要注意与自回归损失的兼容性。实现特定任务逻辑使用定制化模式遮蔽。令牌遮蔽是连接LLM理论能力与实际可控行为的关键桥梁。从确保模型不乱说因果性到让它忽略无关信息填充、窗口再到引导它完成特定任务定制掩码每一种遮蔽技术都对应着一种对模型“注意力”的约束和引导。理解并熟练运用它们是进行LLM二次开发、模型优化乃至安全对齐的必备技能。在实际编码时多画图理解掩码矩阵的形状和含义从小例子开始测试能帮你避开大多数坑。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门