关于HuggingFace Transformer模型注意力头权重及投影矩阵的技术咨询
我用PyTorch结合HuggingFace加载预训练Transformer模型,在Colab中运行以下代码并打印state_dict的键:
model = DistilBertModel.from_pretrained("distilbert-base-uncased") model.state_dict().keys()
输出结果:
odict_keys(['embeddings.word_embeddings.weight', 'embeddings.position_embeddings.weight', 'embeddings.LayerNorm.weight', 'embeddings.LayerNorm.bias', 'transformer.layer.0.attention.q_lin.weight', 'transformer.layer.0.attention.q_lin.bias', 'transformer.layer.0.attention.k_lin.weight', 'transformer.layer.0.attention.k_lin.bias', 'transformer.layer.0.attention.v_lin.weight', 'transformer.layer.0.attention.v_lin.bias', 'transformer.layer.0.attention.out_lin.weight', 'transformer.layer.0.attention.out_lin.bias', 'transformer.layer.0.sa_layer_norm.weight', 'transformer.layer.0.sa_layer_norm.bias', 'transformer.layer.0.ffn.lin1.weight', 'transformer.layer.0.ffn.lin1.bias', 'transformer.layer.0.ffn.lin2.weight', 'transformer.layer.0.ffn.lin2.bias', 'transformer.layer.0.output_layer_norm.weight', 'transformer.layer.0.output_layer_norm.bias', 'transformer.layer.1.attention.q_lin.weight', 'transformer.layer.1.attention.q_lin.bias', 'transformer.layer.1.attention.k_lin.weight', 'transformer.layer.1.attention.k_lin.bias', 'transformer.layer.1.attention.v_lin.weight', 'transformer.layer.1.attention.v_lin.bias', 'transformer.layer.1.attention.out_lin.weight', 'transformer.layer.1.attention.out_lin.bias', 'transformer.layer.1.sa_layer_norm.weight', 'transformer.layer.1.sa_layer_norm.bias', 'transformer.layer.1.ffn.lin1.weight', 'transformer.layer.1.ffn.lin1.bias', 'transformer.layer.1.ffn.lin2.weight', 'transformer.layer.1.ffn.lin2.bias', 'transformer.layer.1.output_layer_norm.weight', 'transformer.layer.1.output_layer_norm.bias', 'transformer.layer.2.attention.q_lin.weight', 'transformer.layer.2.attention.q_lin.bias', 'transformer.layer.2.attention.k_lin.weight', 'transformer.layer.2.attention.k_lin.bias', 'transformer.layer.2.attention.v_lin.weight', 'transformer.layer.2.attention.v_lin.bias', 'transformer.layer.2.attention.out_lin.weight', 'transformer.layer.2.attention.out_lin.bias', 'transformer.layer.2.sa_layer_norm.weight', 'transformer.layer.2.sa_layer_norm.bias', 'transformer.layer.2.ffn.lin1.weight', 'transformer.layer.2.ffn.lin1.bias', 'transformer.layer.2.ffn.lin2.weight', 'transformer.layer.2.ffn.lin2.bias', 'transformer.layer.2.output_layer_norm.weight', 'transformer.layer.2.output_layer_norm.bias', 'transformer.layer.3.attention.q_lin.weight', 'transformer.layer.3.attention.q_lin.bias', 'transformer.layer.3.attention.k_lin.weight', 'transformer.layer.3.attention.k_lin.bias', 'transformer.layer.3.attention.v_lin.weight', 'transformer.layer.3.attention.v_lin.bias', 'transformer.layer.3.attention.out_lin.weight', 'transformer.layer.3.attention.out_lin.bias', 'transformer.layer.3.sa_layer_norm.weight', 'transformer.layer.3.sa_layer_norm.bias', 'transformer.layer.3.ffn.lin1.weight', 'transformer.layer.3.ffn.lin1.bias', 'transformer.layer.3.ffn.lin2.weight', 'transformer.layer.3.ffn.lin2.bias', 'transformer.layer.3.output_layer_norm.weight', 'transformer.layer.3.output_layer_norm.bias', 'transformer.layer.4.attention.q_lin.weight', 'transformer.layer.4.attention.q_lin.bias', 'transformer.layer.4.attention.k_lin.weight', 'transformer.layer.4.attention.k_lin.bias', 'transformer.layer.4.attention.v_lin.weight', 'transformer.layer.4.attention.v_lin.bias', 'transformer.layer.4.attention.out_lin.weight', 'transformer.layer.4.attention.out_lin.bias', 'transformer.layer.4.sa_layer_norm.weight', 'transformer.layer.4.sa_layer_norm.bias', 'transformer.layer.4.ffn.lin1.weight', 'transformer.layer.4.ffn.lin1.bias', 'transformer.layer.4.ffn.lin2.weight', 'transformer.layer.4.ffn.lin2.bias', 'transformer.layer.4.output_layer_norm.weight', 'transformer.layer.4.output_layer_norm.bias', 'transformer.layer.5.attention.q_lin.weight', 'transformer.layer.5.attention.q_lin.bias', 'transformer.layer.5.attention.k_lin.weight', 'transformer.layer.5.attention.k_lin.bias', 'transformer.layer.5.attention.v_lin.weight', 'transformer.layer.5.attention.v_lin.bias', 'transformer.layer.5.attention.out_lin.weight', 'transformer.layer.5.attention.out_lin.bias', 'transformer.layer.5.sa_layer_norm.weight', 'transformer.layer.5.sa_layer_norm.bias', 'transformer.layer.5.ffn.lin1.weight', 'transformer.layer.5.ffn.lin1.bias', 'transformer.layer.5.ffn.lin2.weight', 'transformer.layer.5.ffn.lin2.bias', 'transformer.layer.5.output_layer_norm.weight', 'transformer.layer.5.output_layer_norm.bias'])
疑问点
- 初看输出里似乎没有不同注意力头的权重,这些权重在哪里?
- 是非题:不同注意力头的权重是否已被拼接?观察到投影矩阵为768x768,这是否确实是12个768x64的投影矩阵拼接而成?
- 相关文档在哪里?我在HuggingFace上找不到任何关于这些state_dict键的说明。
补充:我尝试用TensorFlow加载预训练BERT模型,遇到同样问题。Wq和Wk矩阵均为768x768。我猜测12个注意力头的Wq矩阵原本应为64×维度,当前矩阵是按行堆叠这些投影矩阵,但由于PyTorch或TensorFlow均无相关状态定义文档,我无法确认是否反向或转置。
1. 注意力头权重的位置
注意力头的权重并没有单独存储,而是整合在q_lin、k_lin、v_lin这些线性层的权重矩阵里。以DistilBert为例,每个注意力模块的查询、键、值投影都是用单个线性层实现的,而非为每个注意力头单独设置参数。
2. 注意力头权重的拼接方式
是,这些投影矩阵确实是多个注意力头的权重拼接而成。以distilbert-base-uncased为例,它的隐藏层维度是768,注意力头数量是12,每个头的维度为768/12=64。q_lin.weight的形状是(768,768),其实就是12个(64,768)的头投影矩阵按行拼接的结果(PyTorch中线性层权重的形状为(out_features, in_features))。
你可以用代码验证拆分逻辑:
import torch q_weight = model.transformer.layer[0].attention.q_lin.weight split_weights = torch.split(q_weight, 64, dim=0) # 每个split_weights[i]对应第i个注意力头的查询投影权重,形状为(64,768)
对于TensorFlow的BERT模型,逻辑类似,只是TensorFlow中线性层权重的形状为(in_features, out_features),因此768x768的Wq矩阵是12个768x64的头投影矩阵按列拼接的结果。
3. 相关文档说明
HuggingFace Transformers库没有专门针对state_dict键的详细文档,但你可以通过查看模型源码理解参数结构。比如DistilBert的注意力模块源码中,会明确展示q_lin、k_lin、v_lin如何实现多注意力头投影——本质是用单个线性层完成所有头的投影,之后在forward过程中通过view操作把输出拆分为(batch_size, seq_len, num_heads, head_dim)的形状,再转置得到适合计算注意力的维度,这就证明了权重是提前拼接好的。
内容的提问来源于stack exchange,提问作者Paul

