ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

Transformer Demo 准不代表可用:检查掩码、长度与注意力开销

Transformer Demo 准不代表可用:检查掩码、长度与注意力开销 Transformer Demo 准不代表可用检查掩码、长度与注意力开销一个短序列 Demo 输出正常只能证明那组张量形状走通了。换成长序列、Padding 或混合精度后掩码广播和显存开销都可能改变。1. 先写清张量与掩码契约注意力机制的讨论需要同时说明张量形状、掩码语义和数值类型。实验集应在运行前冻结切分规则并用去重与来源隔离检查训练、验证和测试之间的交集。逐层记录 Q、K、V 与 Mask 的形状、dtype 和有效 Token 数。Padding Mask、Causal Mask 和业务可见性 Mask 不应混成一个模糊的布尔数组。2. 用边界输入验证实现解释实现时先核对维度变换和归一化位置再检查长序列、填充和混合精度等边界。模型输出只在对应数据与度量定义下才有解释力。至少覆盖全 Padding、无 Padding、单 Token 和达到服务长度上限的输入并比较参考实现与优化实现。下面的尺寸只是构造用例不能据此推断真实模型的性能。3. 形状检查与稀疏路径[DEBUG] Batch Input Shapes: Q(4, 1, 64), K(4, 128, 64) [DEBUG] Mask Shape: (4, 1, 128) - 包含大量 Padding 零值 [WARN] Attention Softmax Output[0, 0, 96:]: tensor([0.0078, 0.0078, 0.0078, 0.0078 ...]) -- 本应为 0.0 的 Padding 位置分到了概率权重!import torch import torch.nn as nn import torch.nn.functional as F import math from typing import Optional, Tuple class ReproducibleMultiHeadAttention(nn.Module): 可复现、带单步断点校验的 Multi-Head Attention 模块。 精确拦截 Mask 广播异常与 Softmax 权重泄露。 def __init__(self, d_model: int, num_heads: int): super().__init__() assert d_model % num_heads 0, d_model 必须能被 num_heads 整除 self.d_model d_model self.num_heads num_heads self.head_dim d_model // num_heads self.q_proj nn.Linear(d_model, d_model) self.k_proj nn.Linear(d_model, d_model) self.v_proj nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) def forward( self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attn_mask: Optional[torch.Tensor] None ) - Tuple[torch.Tensor, torch.Tensor]: Input Shape: query/key/value: (Batch, Seq_Len, d_model) attn_mask: (Batch, 1, Target_Len, Source_Len) 或 (Batch, Source_Len) batch_size, seq_len, _ query.size() # 1. 线性投影并维度重塑为 (Batch, Num_Heads, Seq_Len, Head_Dim) q self.q_proj(query).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) k self.k_proj(key).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) v self.v_proj(value).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) # 2. Scaled Dot-Product 计算: (Batch, Num_Heads, Seq_Len, Seq_Len) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) # 3. 校验并施加 Mask if attn_mask is not None: # 自动修复 2D Padding Mask 到 4D 广播维度 if attn_mask.dim() 2: # (Batch, Seq_Len) - (Batch, 1, 1, Seq_Len) attn_mask attn_mask.unsqueeze(1).unsqueeze(2) # 使用负无穷值填补掩码位置 scores scores.masked_fill(attn_mask 0, -1e9) # 4. 计算 Softmax 概率矩阵 attn_weights F.softmax(scores, dim-1) # 断点防线断言校验 Padding 位置的权重必须无限接近于 0 if attn_mask is not None: masked_weights_sum (attn_weights * (attn_mask 0)).sum().item() if masked_weights_sum 1e-4: raise ValueError(f[CRITICAL] 检测到 Mask 泄漏! 泄露权重之和: {masked_weights_sum:.6f}) # 5. 聚合 Value 并输出 context torch.matmul(attn_weights, v) context context.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) output self.out_proj(context) return output, attn_weights[TEST 1] 单条运行 (SeqLen12): Attention[0, 0, 5, 5] 0.142857 [TEST 2] Batch 运行 (SeqLen12 与 SeqLen32 拼接): Attention[0, 0, 5, 5] 0.142857 [CHECK] 对应非 PAD 位置注意力数值绝对误差: 0.000000e00 (完全一致) [CHECK] Padding 区域注意力权重最大值: 0.000000e00 (精准掩码)4. 复核清单Q、K、V 与 Mask 的形状和 dtype 是否一致。Padding、Causal 与业务 Mask 是否分开验证。长序列与全 Padding 输入是否覆盖。优化实现是否与参考实现输出对齐。总结“别让演示效果骗了你”应以清晰的条件和脚本复核。先记录边界再解释结果。
返回列表