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

批量广播张量矩阵乘法:求各批次响应与X数据的点积

问题需求

计算每个批次的响应残差(y_true与y_pred的差值)和X数据的点积,期望结果张量形状为[2,5],其中每行对应一个批次。已知各关键张量的形状:

  • yTrue_yHat_allBatches_tensorSub.shape = [2,15]:2个批次,每个批次包含15个响应残差
  • interceptXY_data_allBatches[:, :, :-1].shape = torch.Size([2,15,5]):2个批次,每个批次有15个样本、5个特征

尝试的核心代码行:

y_yhat_allBatches_matmulX_allBatches = torch.matmul(yTrue_yHat_allBatches_tensorSub, interceptXY_data_allBatches[:, :, :-1])

完整可复现代码:

#define dataset
nFeatures_withIntercept = 5
NObservations = 15
miniBatches = 2
interceptXY_data_allBatches = torch.randn(miniBatches, NObservations, nFeatures_withIntercept+1) #+1 Y(response variable)

#random assign beta to work with
beta_holder = torch.rand(nFeatures_withIntercept)

#y_predicted for each mini-batch
y_predBatchAllBatches = torch.matmul(interceptXY_data_allBatches[:, :, :-1], beta_holder)

#y_true - y_predicted for each mini-batch
yTrue_yHat_allBatches_tensorSub = torch.sub(interceptXY_data_allBatches[..., -1], y_predBatchAllBatches)
y_yhat_allBatches_matmulX_allBatches = torch.matmul(yTrue_yHat_allBatches_tensorSub, interceptXY_data_allBatches[:, :, :-1])
解决方案

原代码的核心问题是张量维度不匹配:残差张量是[2,15],X张量是[2,15,5],直接用torch.matmul无法完成批次内的对应计算。需要调整残差的维度,让它和X张量的维度对齐,再执行批次矩阵乘法。

修正后的完整代码:

#define dataset
nFeatures_withIntercept = 5
NObservations = 15
miniBatches = 2
interceptXY_data_allBatches = torch.randn(miniBatches, NObservations, nFeatures_withIntercept+1) #+1 Y(response variable)

#random assign beta to work with
beta_holder = torch.rand(nFeatures_withIntercept)

#y_predicted for each mini-batch
y_predBatchAllBatches = torch.matmul(interceptXY_data_allBatches[:, :, :-1], beta_holder)

#y_true - y_predicted for each mini-batch
yTrue_yHat_allBatches_tensorSub = torch.sub(interceptXY_data_allBatches[..., -1], y_predBatchAllBatches)

# 调整残差维度并执行批次矩阵乘法,得到预期形状
y_yhat_allBatches_matmulX_allBatches = torch.bmm(
    yTrue_yHat_allBatches_tensorSub.unsqueeze(1),  # 转为[2,1,15]
    interceptXY_data_allBatches[:, :, :-1]        # 形状[2,15,5]
).squeeze(1)  # 去掉多余维度,得到[2,5]

# 验证结果形状
print(y_yhat_allBatches_matmulX_allBatches.shape)  # 输出: torch.Size([2, 5])

关键调整说明

  • unsqueeze(1):将残差张量从[2,15]扩展为[2,1,15],让每个批次的残差成为1×15的行向量,匹配X张量的批次维度
  • torch.bmm():专门用于批次矩阵乘法,会对每个批次独立计算1×15与15×5的矩阵乘积,得到每个批次的1×5结果
  • squeeze(1):移除结果中多余的维度,最终得到[2,5]的形状,每行对应一个批次的计算结果

也可以用torch.matmul替代torch.bmm,效果完全一致:

y_yhat_allBatches_matmulX_allBatches = torch.matmul(
    yTrue_yHat_allBatches_tensorSub.unsqueeze(1),
    interceptXY_data_allBatches[:, :, :-1]
).squeeze(1)

内容的提问来源于stack exchange,提问作者sahuno

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 06:22:36