PyTorch MultiheadAttention RuntimeError:形状不匹配问题原因排查
PyTorch MultiheadAttention输入形状错误原因分析
问题详情
初始化代码:
attention = MultiheadAttention(embed_dim=1536, num_heads=4)
输入张量形状:
- query.shape:
torch.Size([1, 1, 1536]) - key.shape 和 value.shape:
torch.Size([1, 23, 1536])
运行时出现以下错误:
RuntimeError Traceback (most recent call last) Cell In[15], line 1 ----> 1 _ = cal_attn_weight_embedding(attention, top_j_sim_video_embeddings_list) File ~/main/reproduct/choi/make_embedding.py:384, in cal_attn_weight_embedding(attention, top_j_sim_video_embeddings_list) 381 print(embedding.shape) 383 # attention --> 384 output, attn_weights = attention(thumbnail, embedding, embedding) 385 # attn_weight shape: (1, 1, j+1) 387 attn_weights = attn_weights.squeeze(0).unsqueeze(-1) # shape: (j+1, 1) File ~/anaconda3/envs/choi_venv/lib/python3.8/site-packages/torch/nn/modules/module.py:1501, in Module._call_impl(self, *args, **kwargs) 1496 # If we don't have any hooks, we want to skip the rest of the logic in 1497 # this function, and just call forward. 1498 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks 1499 or _global_backward_pre_hooks or _global_backward_hooks 1500 or _global_forward_hooks or _global_forward_pre_hooks): --> 1501 return forward_call(*args, **kwargs) 1502 # Do not call functions when jit is used 1503 full_backward_hooks, non_full_backward_hooks = [], [] File ~/anaconda3/envs/choi_venv/lib/python3.8/site-packages/torch/nn/modules/activation.py:1205, in MultiheadAttention.forward(self, query, key, value, key_padding_mask, need_weights, attn_mask, average_attn_weights, is_causal) 1191 attn_output, attn_output_weights = F.multi_head_attention_forward( 1192 query, key, value, self.embed_dim, self.num_heads, ... 5281 # TODO finish disentangling control flow so we don't do in-projections when statics are passed 5282 assert static_k.size(0) == bsz * num_heads, \ 5283 f"expecting static_k.size(0) of {bsz * num_heads}, but got {static_k.size(0)}" RuntimeError: shape '[1, 4, 384]' is invalid for input of size 35328
运行环境:
- Ubuntu 20.04
- Anaconda 1.7.2
- Python 3.8.5
- VSCode 1.87.2
- PyTorch 2.0.1
错误原因
PyTorch原生MultiheadAttention默认要求输入张量的维度顺序为**[序列长度(seq_len), 批次大小(batch_size), 嵌入维度(embed_dim)]**,但你传入的张量是[batch_size, seq_len, embed_dim]的顺序,导致维度解析错误。
具体验证:
每个注意力头的维度为embed_dim / num_heads = 1536 /4 =384。代码期望将key张量处理为[batch_size*num_heads, seq_len, head_dim]的中间形状(即[1*4, 23, 384]),但因为输入顺序错误,代码错误地把第一维1当成序列长度、第二维23当成批次大小,尝试将总元素数为1*23*1536=35328的张量reshape为[1*4, 1, 384](总元素数1536),两者不匹配,因此抛出形状错误。
解决办法
有两种修正方式:
- 交换输入张量的前两维,转换成模型默认要求的维度顺序:
# 转换query、key、value的维度顺序 query = query.transpose(0, 1) key = key.transpose(0, 1) value = value.transpose(0, 1) # 再传入attention计算 output, attn_weights = attention(query, key, value)
- 初始化时设置
batch_first=True,让模型直接支持[batch_size, seq_len, embed_dim]的输入顺序:
# 初始化时添加batch_first参数 attention = MultiheadAttention(embed_dim=1536, num_heads=4, batch_first=True) # 直接传入原形状的张量即可 output, attn_weights = attention(query, key, value)
内容的提问来源于stack exchange,提问作者ララララ
相关产品推荐
相关产品推荐

