PyTorch如何实现(N,*,in_feature)矩阵乘得到(N,*,out_feature)结果
PyTorch手动复现nn.Linear运算逻辑的正确实现
input @ weight.T本身是符合张量广播规则的,正常计算得到的输出形状就是你需要的(N, *, out_feature),如果结果不符合预期,通常是两个原因:漏加偏置项、维度数值不匹配。
完整实现逻辑
nn.Linear的官方运算逻辑为:输出 = 输入 × 权重转置 + 偏置,你可以选择以下两种写法实现:
写法1:矩阵乘法写法(最简洁)
# weight形状为(out_feature, in_feature),bias形状为(out_feature,) output = input @ weight.T + bias # 如果不需要偏置,直接写 output = input @ weight.T
写法2:einsum显式指定维度匹配(可读性更高,避免维度错位)
# 显式指定输入最后一维与权重第二维匹配,保留所有中间维度 output = torch.einsum('...i,oi->...o', input, weight) + bias # 如果不需要偏置,直接去掉+bias部分即可
验证与对齐官方结果示例
你可以用以下代码验证自行实现的结果与官方nn.Linear完全一致:
import torch import torch.nn as nn # 配置超参数 in_features = 4 out_features = 5 # 初始化官方Linear层 official_linear = nn.Linear(in_features, out_features) weight = official_linear.weight bias = official_linear.bias # 构造带任意中间维度的输入 input = torch.randn(2, 3, 6, in_features) # 形状为(2,3,6,4),对应N=2,*为3、6两个维度 # 官方输出 official_out = official_linear(input) # 自行实现输出 my_out = input @ weight.T + bias # 验证形状与数值都匹配 print(my_out.shape) # 输出torch.Size([2, 3, 6, 5]),符合(N,*,out_feature)要求 print(torch.allclose(official_out, my_out)) # 输出True,数值完全对齐
常见报错排查
如果你的input @ weight.T结果不符合预期,优先检查以下两点:
- 确认输入张量的最后一维大小,与权重张量的第二维大小相等,即
input.shape[-1] == weight.shape[1] - 确认你是否开启了nn.Linear的偏置(默认是开启状态),如果开启偏置必须加上偏置项才能对齐结果
内容的提问来源于stack exchange,提问作者Maxwell Albert
相关产品推荐
相关产品推荐

