替换nn.Linear为nn.Parameter后Transformer注意力性能骤降原因咨询
我在PyTorch中做Transformer相关实验时,为了单独查看不同权重矩阵,将自注意力计算中的nn.Linear模块替换为nn.Parameter(torch.tensor()),但发现模型性能出现显著下降。以下是两种实现:
第一种(使用nn.Linear)
class Attention(nn.Module): def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.): super().__init__() inner_dim = dim_head * heads project_out = not (heads == 1 and dim_head == dim) self.heads = heads self.scale = dim_head ** -0.5 self.attend = nn.Softmax(dim = -1) self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False) self.to_out = nn.Sequential( nn.Linear(inner_dim, dim), nn.Dropout(dropout) ) if project_out else nn.Identity() def forward(self, x): qkv = self.to_qkv(x).chunk(3, dim = -1) q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = self.heads), qkv) dots = torch.matmul(q, k.transpose(-1, -2)) * self.scale attn = self.attend(dots) out = torch.matmul(attn, v) out = rearrange(out, 'b h n d -> b n (h d)') return self.to_out(out)
第二种(使用nn.Parameter(torch.tensor()))
class Attention(nn.Module): def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.): super().__init__() inner_dim = dim_head * heads project_out = not (heads == 1 and dim_head == dim) self.dim_head = dim_head self.heads = heads self.scale = dim_head ** -0.5 self.attend = nn.Softmax(dim = -1) self.to_q = nn.Parameter(torch.randn(dim, inner_dim)) self.to_k = nn.Parameter(torch.randn(dim, inner_dim)) self.to_v = nn.Parameter(torch.randn(dim, inner_dim)) self.projection = nn.Parameter(torch.randn(inner_dim, dim)) self.dropout = nn.Dropout(dropout) def forward(self, x): q,k,v = x @ self.to_q, x @ self.to_k, x @ self.to_v q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = self.heads), (q,k,v)) dots = torch.matmul(q, k.transpose(-1, -2)) * self.scale attn = self.attend(dots) out = torch.matmul(attn, v) out = rearrange(out, 'b h n d -> b n (h d)') out = out @ self.projection out = self.dropout(out) return out
性能差异的核心原因
1. 权重初始化方式完全不同
PyTorch的nn.Linear模块默认使用Kaiming均匀初始化(公式:init.kaiming_uniform_(self.weight, a=math.sqrt(5))),这种初始化专门针对线性层设计,能保证初始输出的尺度合理,避免训练初期梯度爆炸或消失。而手动用torch.randn初始化的参数是标准正态分布(均值0,方差1),权重尺度远大于Kaiming初始化的结果,会导致初始阶段模型输出幅值异常,Softmax后的梯度计算出现问题,直接影响模型收敛速度和最终性能。
2. 输出投影逻辑不一致
第一种实现中,当heads == 1且dim_head == dim时(即project_out=False),会直接返回输入特征(用nn.Identity()),不做额外的线性变换和dropout;但第二种实现不管这个条件,始终执行投影和dropout操作,这会破坏输入输出维度一致的特性,引入不必要的特征变换,导致模型学习到的特征被破坏,性能下降。
3. PyTorch内置模块的优化优势
nn.Linear是PyTorch官方优化过的模块,不仅自动处理设备迁移(如CPU/GPU切换),还在自动微分、内存使用等方面做了优化。手动用矩阵乘法x @ self.to_q虽然数学上等价,但反向传播时的梯度计算效率更低,长期训练会导致收敛速度变慢,最终性能不如内置模块。
修复建议
对齐初始化方式:如果坚持用
nn.Parameter,手动使用Kaiming初始化权重:import math from torch.nn import init self.to_q = nn.Parameter(init.kaiming_uniform_(torch.empty(dim, inner_dim), a=math.sqrt(5))) self.to_k = nn.Parameter(init.kaiming_uniform_(torch.empty(dim, inner_dim), a=math.sqrt(5))) self.to_v = nn.Parameter(init.kaiming_uniform_(torch.empty(dim, inner_dim), a=math.sqrt(5))) self.projection = nn.Parameter(init.kaiming_uniform_(torch.empty(inner_dim, dim), a=math.sqrt(5)))对齐输出投影逻辑:先把
project_out设为实例变量self.project_out = project_out,再在forward中根据条件执行操作:def forward(self, x): # ... 前面的计算逻辑不变 out = rearrange(out, 'b h n d -> b n (h d)') if self.project_out: out = out @ self.projection out = self.dropout(out) return out更优方案:保留nn.Linear同时查看权重:其实不需要替换成
nn.Parameter,直接用nn.Linear模块,想要查看权重时访问module.weight.data即可,比如:self.to_q = nn.Linear(dim, inner_dim, bias=False) # 查看权重 print(self.to_q.weight.data)这种方式既保留了
nn.Linear的所有优化优势,又能满足单独查看权重的需求。
内容的提问来源于stack exchange,提问作者new2java

