资讯动态

【Bug已解决】Difference between src_mask and src_key_padding_mask 解决方案

发布时间:2026/8/28 23:02:54 来源:尧图企业网站定制
【Bug已解决】Difference between src_mask and src_key_padding_mask 解决方案问题描述在 PyTorch 中使用nn.Transformer或nn.MultiheadAttention时开发者经常对src_mask和src_key_padding_mask这两个参数感到困惑。它们都用于遮蔽masking某些位置但用途和行为完全不同。混淆这两个参数会导致模型训练异常、注意力计算错误甚至产生 NaN 梯度。常见的问题表现使用了错误的 mask 类型导致模型性能下降mask 的形状和数值类型不正确触发运行时错误在解码器中混淆了因果掩码和填充掩码mask 的数据类型不对bool vs float导致注意力分数计算异常不理解 mask 的广播规则导致维度不匹配错误复现以下代码演示了混淆src_mask和src_key_padding_mask时出现的问题import torch import torch.nn as nn # 创建 Transformer 编码器层 encoder_layer nn.TransformerEncoderLayer( d_model512, nhead8, batch_firstTrue ) # 模拟输入 (batch_size4, seq_len10, d_model512) src torch.randn(4, 10, 512) # 错误1混淆 mask 类型 # src_key_padding_mask 应该是 (batch_size, seq_len) 的 bool 张量 # src_mask 应该是 (seq_len, seq_len) 或 (batch_size * num_heads, seq_len, seq_len) # 错误把 padding mask 传给了 src_mask padding_mask torch.tensor([ [False, False, False, False, False, True, True, True, True, True], [False, False, False, False, False, False, False, True, True, True], [False, False, False, False, False, False, False, False, False, True], [False, False, False, False, False, False, False, False, False, False], ]) # shape: (4, 10) - True 表示需要被遮蔽 try: # 错误padding mask 的形状是 (batch, seq_len)但 src_mask 期望 (seq_len, seq_len) output encoder_layer(src, src_maskpadding_mask) except RuntimeError as e: print(f错误1: {e}) # RuntimeError: The shape of mask (4, 10) at /... dimension 0 # does not reflect the number of attention heads or sequence length # 错误2mask 数据类型错误 # 使用 float 类型的 mask但没有用正确的填充值 float_mask torch.zeros(4, 10) float_mask[:, 5:] -float(inf) # float mask try: # src_key_padding_mask 期望 bool 类型 output encoder_layer(src, src_key_padding_maskfloat_mask) except Exception as e: print(f错误2: {type(e).__name__}: {e}) # 错误3mask 值的含义混淆 # bool mask 中 True 表示遮蔽不允许注意但有些开发者理解为保留 wrong_mask torch.tensor([ [True, True, True, True, True, False, False, False, False, False], # True 在这里被错误地理解为保留但实际含义是遮蔽 ]) # 错误4因果掩码用错位置 # 因果掩码应该用于解码器tgt_mask但被用在了编码器src_mask seq_len 10 causal_mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() # 这是一个上三角矩阵用于防止解码器看到未来的 token # 在编码器中使用因果掩码通常不需要除非是自回归编码器 output encoder_layer(src, src_maskcausal_mask) # 虽然不会报错但语义上可能不正确根因分析1.src_mask注意力掩码 / Attention Masksrc_mask是一个注意力分数级别的掩码用于控制哪些位置可以 attend 到哪些位置。它的形状是(seq_len, seq_len)或(batch_size * num_heads, seq_len, seq_len)。用途因果掩码Causal Mask在解码器中防止看到未来的 token自回归生成自定义注意力模式限制特定的位置对之间的注意力数据类型bool类型True表示不允许注意被遮蔽False表示允许float类型添加到注意力分数上通常用-inf或很大的负数来遮蔽广播规则(seq_len, seq_len)所有 batch 和所有 head 共享同一个 mask(batch_size * num_heads, seq_len, seq_len)每个 head 和 batch 可以有不同的 mask2.src_key_padding_mask填充掩码 / Padding Masksrc_key_padding_mask是一个序列级别的掩码用于标记哪些位置是填充的padding不应该被注意到。它的形状是(batch_size, seq_len)。用途在变长序列的批处理中短序列被填充到相同长度填充位置不应该参与注意力计算防止模型关注无意义的填充 token数据类型bool类型True表示该位置是填充的需要被遮蔽False表示是有效数据PyTorch 内部会自动将 bool mask 转换为 float mask将 True 位置设为-inf广播规则(batch_size, seq_len)每个样本有自己的 padding maskPyTorch 内部会将其广播到(batch_size * num_heads, seq_len, seq_len)的注意力分数矩阵3. 两者的关键区别特性src_masksrc_key_padding_mask形状(S, S)或(B*H, S, S)(B, S)作用级别注意力分数矩阵Key/Value 序列用途因果掩码、自定义注意力模式填充位置遮蔽语义控制 query 位置 attend 到 key 位置标记哪些 key 位置是填充典型场景解码器的自回归生成变长序列的批处理4. 组合使用在实际的 Transformer 模型中src_mask和src_key_padding_mask经常同时使用。例如在解码器中tgt_mask因果掩码防止看到未来 tokentgt_key_padding_mask填充掩码忽略填充位置memory_key_padding_mask编码器输出的填充掩码PyTorch 内部会将两个 mask 合并对注意力分数矩阵同时应用因果约束和填充约束。解决方案方案一正确创建和使用src_key_padding_maskimport torch import torch.nn as nn def create_padding_mask(sequences, pad_idx0): 创建 padding mask。 Args: sequences: 输入序列 (batch_size, seq_len)整数索引 pad_idx: 填充 token 的索引 Returns: padding_mask: bool 张量 (batch_size, seq_len) True 表示该位置是填充的需要被遮蔽 return sequences pad_idx # 使用示例 # 模拟一个 batch 的 token 序列0 表示 padding sequences torch.tensor([ [1, 2, 3, 4, 5, 0, 0, 0], # 有效长度 5 [1, 2, 3, 0, 0, 0, 0, 0], # 有效长度 3 [1, 2, 3, 4, 5, 6, 7, 8], # 有效长度 8 [1, 2, 0, 0, 0, 0, 0, 0], # 有效长度 2 ]) padding_mask create_padding_mask(sequences, pad_idx0) print(fPadding mask:\n{padding_mask}) # tensor([ # [False, False, False, False, False, True, True, True], # [False, False, False, True, True, True, True, True], # [False, False, False, False, False, False, False, False], # [False, False, True, True, True, True, True, True], # ]) # 在 Transformer 编码器中使用 encoder_layer nn.TransformerEncoderLayer( d_model64, nhead4, batch_firstTrue ) encoder nn.TransformerEncoder(encoder_layer, num_layers2) # 将 token 转为嵌入 embedding nn.Embedding(100, 64, padding_idx0) src embedding(sequences) # (batch, seq_len, d_model) # 使用 padding mask output encoder(src, src_key_padding_maskpadding_mask) print(f输出形状: {output.shape}) # (4, 8, 64)方案二正确创建和使用src_mask因果掩码import torch import torch.nn as nn def create_causal_mask(seq_len): 创建因果掩码下三角矩阵。 用于自回归模型防止当前位置看到未来的 token。 Args: seq_len: 序列长度 Returns: causal_mask: bool 张量 (seq_len, seq_len) True 表示该位置被遮蔽不允许注意 # 生成上三角矩阵不包括对角线 mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() return mask # 使用示例 seq_len 8 causal_mask create_causal_mask(seq_len) print(f因果掩码:\n{causal_mask}) # tensor([ # [False, True, True, True, True, True, True, True], # [False, False, True, True, True, True, True, True], # [False, False, False, True, True, True, True, True], # ... # ]) # 第 0 行只能看到位置 0 # 第 1 行只能看到位置 0, 1 # 第 2 行只能看到位置 0, 1, 2 # ... # 在 Transformer 解码器中使用 decoder_layer nn.TransformerDecoderLayer( d_model64, nhead4, batch_firstTrue ) decoder nn.TransformerDecoder(decoder_layer, num_layers2) batch_size 4 tgt torch.randn(batch_size, seq_len, 64) memory torch.randn(batch_size, seq_len, 64) # 使用因果掩码 output decoder( tgttgt, memorymemory, tgt_maskcausal_mask, # 因果掩码用于 target ) print(f输出形状: {output.shape}) # (4, 8, 64)方案三同时使用两种 maskimport torch import torch.nn as nn def create_masks(src_sequences, tgt_sequences, pad_idx0): 为 Transformer 创建所有需要的 mask。 Returns: src_mask: 编码器的注意力 mask通常为 None tgt_mask: 解码器的因果 mask src_padding_mask: 编码器的 padding mask tgt_padding_mask: 解码器的 padding mask memory_padding_mask: 编码器输出的 padding mask同 src_padding_mask src_seq_len src_sequences.size(1) tgt_seq_len tgt_sequences.size(1) # 因果掩码解码器用 tgt_mask torch.triu( torch.ones(tgt_seq_len, tgt_seq_len), diagonal1 ).bool() # 填充掩码 src_padding_mask (src_sequences pad_idx) tgt_padding_mask (tgt_sequences pad_idx) return tgt_mask, src_padding_mask, tgt_padding_mask # 使用示例 src_sequences torch.tensor([ [1, 2, 3, 4, 5, 0, 0, 0], [1, 2, 3, 4, 5, 6, 7, 8], ]) tgt_sequences torch.tensor([ [10, 20, 30, 0, 0, 0], [10, 20, 30, 40, 50, 0], ]) tgt_mask, src_pad_mask, tgt_pad_mask create_masks( src_sequences, tgt_sequences, pad_idx0 ) print(fsrc_padding_mask 形状: {src_pad_mask.shape}) # (2, 8) print(ftgt_padding_mask 形状: {tgt_pad_mask.shape}) # (2, 6) print(ftgt_mask 形状: {tgt_mask.shape}) # (6, 6) # 完整的 Transformer transformer nn.Transformer( d_model64, nhead4, num_encoder_layers2, num_decoder_layers2, batch_firstTrue, ) # 嵌入 src_embedding nn.Embedding(100, 64, padding_idx0) tgt_embedding nn.Embedding(100, 64, padding_idx0) src src_embedding(src_sequences) tgt tgt_embedding(tgt_sequences) # 前向传播使用所有 mask output transformer( srcsrc, ![配图](https://i-blog.csdnimg.cn/img_convert/ff8584dd62f1dfbfc5f0f93a48bfe670.png) tgttgt, tgt_masktgt_mask, # 因果掩码 src_key_padding_masksrc_pad_mask, # 编码器 padding tgt_key_padding_masktgt_pad_mask, # 解码器 padding memory_key_padding_masksrc_pad_mask, # 交叉注意力 padding ) print(f输出形状: {output.shape}) # (2, 6, 64)完整修复代码以下是一个完整的、正确使用各种 mask 的 Transformer 模型实现import torch import torch.nn as nn import torch.nn.functional as F import math from typing import Optional, Tuple class MaskFactory: Transformer mask 工厂类。 统一创建和管理各种类型的 mask。 staticmethod def create_padding_mask(sequences: torch.Tensor, pad_idx: int 0) - torch.Tensor: 创建 padding mask。 Args: sequences: (batch_size, seq_len) token 索引序列 pad_idx: padding token 的索引值 Returns: (batch_size, seq_len) bool 张量 True 该位置是 padding需要被遮蔽 return sequences pad_idx staticmethod def create_causal_mask(seq_len: int, device: Optional[torch.device] None) - torch.Tensor: 创建因果掩码下三角。 用于自回归解码器防止看到未来 token。 Returns: (seq_len, seq_len) bool 张量 True 该位置被遮蔽 mask torch.triu( torch.ones(seq_len, seq_len, devicedevice, dtypetorch.bool), diagonal1 ) return mask staticmethod def create_attention_mask(allowed_attentions: torch.Tensor) - torch.Tensor: 根据自定义的注意力模式创建 mask。 Args: allowed_attentions: (seq_len, seq_len) bool 张量 True 允许注意False 不允许 Returns: (seq_len, seq_len) bool 张量取反True 遮蔽 return ~allowed_attentions staticmethod def create_float_mask(bool_mask: torch.Tensor, mask_value: float float(-inf)) - torch.Tensor: 将 bool mask 转为 float mask。 Args: bool_mask: bool 张量True 遮蔽 mask_value: 遮蔽位置的值通常为 -inf Returns: float 张量遮蔽位置为 mask_value其他位置为 0 float_mask torch.zeros_like(bool_mask, dtypetorch.float32) float_mask.masked_fill_(bool_mask, mask_value) return float_mask staticmethod def combine_masks(*masks: Optional[torch.Tensor]) - Optional[torch.Tensor]: 合并多个 mask逻辑或。 任一 mask 中被遮蔽的位置在结果中也被遮蔽。 valid_masks [m for m in masks if m is not None] if not valid_masks: return None combined valid_masks[0].clone() for m in valid_masks[1:]: combined combined | m return combined class MultiHeadAttentionWithMasks(nn.Module): 带完整 mask 支持的多头注意力。 清晰展示 attn_mask 和 key_padding_mask 的区别。 def __init__(self, d_model: int, num_heads: int): super().__init__() assert d_model % num_heads 0 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads self.W_q nn.Linear(d_model, d_model) self.W_k nn.Linear(d_model, d_model) self.W_v nn.Linear(d_model, d_model) self.W_o nn.Linear(d_model, d_model) def forward( self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attn_mask: Optional[torch.Tensor] None, key_padding_mask: Optional[torch.Tensor] None, ) - Tuple[torch.Tensor, torch.Tensor]: Args: query: (batch_size, seq_len_q, d_model) key: (batch_size, seq_len_k, d_model) value: (batch_size, seq_len_v, d_model) attn_mask: (seq_len_q, seq_len_k) 或 (batch * heads, seq_len_q, seq_len_k) 注意力级别的 mask key_padding_mask: (batch_size, seq_len_k) 标记 key 中的 padding 位置 Returns: output: (batch_size, seq_len_q, d_model) attention_weights: (batch_size, num_heads, seq_len_q, seq_len_k) batch_size query.size(0) seq_len_q query.size(1) seq_len_k key.size(1) # 线性变换 Q self.W_q(query) # (B, S_q, D) K self.W_k(key) # (B, S_k, D) V self.W_v(value) # (B, S_k, D) # 分头: (B, S, D) - (B, H, S, d_k) Q Q.view(batch_size, seq_len_q, self.num_heads, self.d_k).transpose(1, 2) K K.view(batch_size, seq_len_k, self.num_heads, self.d_k).transpose(1, 2) V V.view(batch_size, seq_len_k, self.num_heads, self.d_k).transpose(1, 2) # 注意力分数: (B, H, S_q, S_k) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # 应用 attn_mask # attn_mask 是 (S_q, S_k) 或 (B*H, S_q, S_k) 的 if attn_mask is not None: if attn_mask.dim() 2: # (S_q, S_k) - 广播到 (B, H, S_q, S_k) if attn_mask.dtype torch.bool: scores.masked_fill_(attn_mask.unsqueeze(0).unsqueeze(0), float(-inf)) else: scores scores attn_mask.unsqueeze(0).unsqueeze(0) elif attn_mask.dim() 3: # (B*H, S_q, S_k) - reshape 为 (B, H, S_q, S_k) attn_mask attn_mask.view(batch_size, self.num_heads, seq_len_q, seq_len_k) if attn_mask.dtype torch.bool: scores.masked_fill_(attn_mask, float(-inf)) else: scores scores attn_mask # 应用 key_padding_mask # key_padding_mask 是 (B, S_k) 的 if key_padding_mask is not None: # (B, S_k) - (B, 1, 1, S_k) 广播到所有 head 和 query 位置 if key_padding_mask.dtype torch.bool: scores.masked_fill_( key_padding_mask.unsqueeze(1).unsqueeze(2), float(-inf) ) else: scores scores key_padding_mask.unsqueeze(1).unsqueeze(2) # Softmax attn_weights F.softmax(scores, dim-1) # 处理全 -inf 行避免 NaN # 如果某一行全是 -infsoftmax 会产生 NaN # 用 0 替换 NaN attn_weights torch.nan_to_num(attn_weights, nan0.0) # 加权求和: (B, H, S_q, d_k) context torch.matmul(attn_weights, V) # 合并头: (B, S_q, D) context context.transpose(1, 2).contiguous() context context.view(batch_size, seq_len_q, self.d_model) # 输出变换 output self.W_o(context) return output, attn_weights class TransformerModel(nn.Module): 完整的 Transformer 模型序列到序列。 正确使用所有类型的 mask。 def __init__(self, src_vocab_size, tgt_vocab_size, d_model256, nhead8, num_layers4, d_ff512, pad_idx0, max_len512): super().__init__() self.pad_idx pad_idx # 嵌入层 self.src_embedding nn.Embedding(src_vocab_size, d_model, padding_idxpad_idx) self.tgt_embedding nn.Embedding(tgt_vocab_size, d_model, padding_idxpad_idx) self.positional_encoding PositionalEncoding(d_model, max_len) # Transformer self.transformer nn.Transformer( d_modeld_model, nheadnhead, num_encoder_layersnum_layers, num_decoder_layersnum_layers, dim_feedforwardd_ff, dropout0.1, batch_firstTrue, ) # 输出层 self.output_projection nn.Linear(d_model, tgt_vocab_size) def forward(self, src, tgt): Args: src: (batch, src_len) 源序列 token 索引 tgt: (batch, tgt_len) 目标序列 token 索引 # 创建 mask src_padding_mask MaskFactory.create_padding_mask(src, self.pad_idx) tgt_padding_mask MaskFactory.create_padding_mask(tgt, self.pad_idx) tgt_causal_mask MaskFactory.create_causal_mask(tgt.size(1), devicetgt.device) # 嵌入 位置编码 src_emb self.positional_encoding(self.src_embedding(src)) tgt_emb self.positional_encoding(self.tgt_embedding(tgt)) # Transformer 前向传播 output self.transformer( srcsrc_emb, tgttgt_emb, tgt_masktgt_causal_mask, # 因果掩码 src_key_padding_masksrc_padding_mask, # 编码器 padding tgt_key_padding_masktgt_padding_mask, # 解码器 padding memory_key_padding_masksrc_padding_mask, # 交叉注意力 padding ) # 投影到词表大小 logits self.output_projection(output) return logits torch.no_grad() def greedy_decode(self, src, max_len50, bos_idx1, eos_idx2): 贪心解码推理时使用。 self.eval() batch_size src.size(0) device src.device # 编码 src_padding_mask MaskFactory.create_padding_mask(src, self.pad_idx) src_emb self.positional_encoding(self.src_embedding(src)) memory self.transformer.encoder(src_emb, src_key_padding_masksrc_padding_mask) # 初始化解码输入 tgt torch.full((batch_size, 1), bos_idx, dtypetorch.long, devicedevice) for step in range(max_len): # 创建当前步的因果掩码 tgt_len tgt.size(1) tgt_causal_mask MaskFactory.create_causal_mask(tgt_len, devicedevice) tgt_padding_mask MaskFactory.create_padding_mask(tgt, self.pad_idx) # 解码 tgt_emb self.positional_encoding(self.tgt_embedding(tgt)) dec_output self.transformer.decoder( tgt_emb, memory, tgt_masktgt_causal_mask, tgt_key_padding_masktgt_padding_mask, memory_key_padding_masksrc_padding_mask, ) # 预测下一个 token logits self.output_projection(dec_output[:, -1, :]) next_token logits.argmax(dim-1, keepdimTrue) # 拼接 tgt torch.cat([tgt, next_token], dim1) # 如果所有序列都生成了 EOS停止 if (next_token eos_idx).all(): break return tgt class PositionalEncoding(nn.Module): 位置编码 def __init__(self, d_model, max_len512): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) self.register_buffer(pe, pe.unsqueeze(0)) def forward(self, x): return x self.pe[:, :x.size(1)] # 演示 if __name__ __main__: print( * 60) print(Mask 类型对比演示) print( * 60) # 创建模拟数据 batch_size 4 src_len 8 tgt_len 6 src_vocab 100 tgt_vocab 100 src torch.randint(1, src_vocab, (batch_size, src_len)) tgt torch.randint(1, tgt_vocab, (batch_size, tgt_len)) # 添加 padding src[0, 5:] 0 src[1, 3:] 0 tgt[0, 4:] 0 tgt[2, 2:] 0 print(f源序列:\n{src}) print(f目标序列:\n{tgt}) # 创建 mask src_pad MaskFactory.create_padding_mask(src) tgt_pad MaskFactory.create_padding_mask(tgt) tgt_causal MaskFactory.create_causal_mask(tgt_len) print(f\nsrc_padding_mask (shape{src_pad.shape}):) print(f {src_pad}) print(f\ntgt_padding_mask (shape{tgt_pad.shape}):) print(f {tgt_pad}) print(f\ntgt_causal_mask (shape{tgt_causal.shape}):) print(f {tgt_causal}) # 创建模型 model TransformerModel(src_vocab, tgt_vocab, d_model64, nhead4, num_layers2) # 前向传播 logits model(src, tgt) print(f\n输出 logits 形状: {logits.shape}) # (4, 6, 100) # 贪心解码 print(\n--- 贪心解码 ---) decoded model.greedy_decode(src, max_len10) print(f解码结果形状: {decoded.shape}) print(f解码结果:\n{decoded}) # 演示自定义注意力 print(\n--- 自定义多头注意力 ---) mha MultiHeadAttentionWithMasks(d_model64, num_heads4) query torch.randn(2, 5, 64) key torch.randn(2, 8, 64) value torch.randn(2, 8, 64) # 只使用 attn_mask attn_mask MaskFactory.create_causal_mask(5) # (5, 5) output, weights mha(query, key, value, attn_maskattn_mask) print(f使用 attn_mask: output{output.shape}, weights{weights.shape}) # 只使用 key_padding_mask key_padding_mask torch.tensor([ [False, False, False, False, True, True, True, True], [False, False, False, False, False, False, False, False], ]) output, weights mha(query, key, value, key_padding_maskkey_padding_mask) print(f使用 key_padding_mask: output{output.shape}, weights{weights.shape}) # 同时使用两种 mask output, weights mha( query, key, value, attn_maskattn_mask, key_padding_maskkey_padding_mask, ) print(f同时使用两种 mask: output{output.shape}, weights{weights.shape}) print(\n * 60) print(演示完成) print( * 60)常见陷阱与注意事项1. bool mask 中 True 的含义在 PyTorch 的 mask 中True表示遮蔽不允许注意False表示允许。这与很多开发者的直觉相反有些人认为 True 表示保留。2. NaN 问题当某一行注意力分数全为-inf时softmax 会产生 NaN。这通常发生在 padding mask 遮蔽了所有 key 位置时。解决方法确保每个 query 至少有一个有效的 key 位置使用torch.nan_to_num()处理 NaN3. mask 的设备一致性mask 必须与输入张量在同一个设备上CPU 或 CUDA。使用mask.to(device)确保一致性。4.batch_first参数的影响nn.Transformer的batch_firstTrue时输入形状为(batch, seq_len, d_model)batch_firstFalse默认时为(seq_len, batch, d_model)。但 mask 的形状不受此参数影响。5. 3D mask 的 head 维度PyTorch 2.0 支持 3D 的attn_mask形状为(batch, seq_len, seq_len)会自动广播到所有 head。如果需要每个 head 有不同的 mask使用 4D(batch * heads, seq_len, seq_len)。6. float mask 的加法语义float 类型的 mask 是加到注意力分数上的而非替换。因此非遮蔽位置应设为 0而非 1遮蔽位置设为-inf或很大的负数。总结src_mask和src_key_padding_mask是 Transformer 中两个功能完全不同的 mask 参数。核心区别如下src_maskattn_mask注意力分数级别的掩码形状(S, S)用于控制位置间的注意力模式如因果掩码。src_key_padding_mask序列级别的填充掩码形状(B, S)用于标记 padding 位置。bool mask 语义True 遮蔽不允许注意False 允许。float mask 语义加到注意力分数上-inf遮蔽0不影响。组合使用在完整的 Transformer 中因果掩码用于解码器自注意力padding mask 用于编码器和解码器。NaN 处理全-inf行会产生 NaN需要确保至少一个有效位置或使用nan_to_num。正确理解和使用这两种 mask是构建高质量 Transformer 模型的基础。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价