Keras转PyTorch实现Transformer参数量差异过大问题排查
参数量差异核心原因
参数量差距和Sequential层无关,完全是Keras与PyTorch的MultiHeadAttention层参数语义不匹配、传参未对齐导致的,计算结果和你贴出的参数统计完全吻合:
- PyTorch的
nn.MultiheadAttention入参embed_dim指注意力层整体的输入/输出总嵌入维度,要求该值能被num_heads整除,单头维度自动计算为head_dim = embed_dim // num_heads。以你代码中embed_dim=768、常规12头配置为例,单头维度为64,Q/K/V三个投影层+最终输出投影层总参数量约236万,加上FFN、LayerNorm的参数,刚好对应你PyTorch模型370万左右的总参数量。 - Keras的
layers.MultiHeadAttention入参key_dim指单个注意力头的维度,不是总嵌入维度。如果仅传入num_heads和key_dim,层内部会将总投影维度设为num_heads * key_dim,默认value_dim与key_dim相等。你当前代码写的是key_dim=embed_dim=768,按12头计算总投影维度就变成了12*768=9216,仅Q/K/V+输出投影四个全连接层参数量就达到约2833万,加上FFN的314.8万、LayerNorm的3072参数,总和刚好是你Keras模型TransformerBlock的31491584参数量,和统计值完全一致。
修正方式
要对齐两边模型结构,只需修改Keras侧MultiHeadAttention的初始化参数,传入单头维度而非总嵌入维度即可:
# 需保证embed_dim能被num_heads整除,为Transformer标准配置 self.att = layers.MultiHeadAttention( num_heads=num_heads, key_dim=embed_dim // num_heads )
修改后两边参数量会完全一致。
Sequential层逻辑差异说明
Keras和PyTorch对注意力层后的Sequential/FFN层处理没有本质区别:
- 你两侧实现的FFN结构完全一致,均为「全连接升维到ff_dim → ReLU激活 → 全连接降维回embed_dim」,从参数统计来看这部分两侧数值完全相等(PyTorch侧两个线性层参数和为3148544,与Keras侧FFN参数量匹配),不存在实现偏差。
- Keras的
keras.Sequential与PyTorch的nn.Sequential行为完全一致,都是按顺序执行容器内的层,前一层输出直接作为后一层输入,无额外隐式处理。
额外注意事项
你PyTorch代码中最后注释掉了Softmax层,Keras侧最后一层则使用了Softmax激活,这部分不影响参数量但会影响训练:PyTorch的CrossEntropyLoss自带log_softmax计算,不需要手动在网络末尾加Softmax;Keras侧如果使用SparseCategoricalCrossentropy,设置from_logits=False时才需要末尾加Softmax,训练时注意对齐损失函数逻辑即可。
内容的提问来源于stack exchange,提问作者Lajos Muzsai
相关产品推荐
相关产品推荐

