You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 14:48:03