批量广播张量矩阵乘法:求各批次响应与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
相关产品推荐
相关产品推荐

