固定尺寸张量的Mahalanobis距离函数简化及C++改写需求
固定尺寸马氏距离最简实现方案
问题定位
你遇到的mat1 and mat2 shapes cannot be multiplied (42x42 and 45x42)错误,核心是矩阵乘法维度不匹配:42×42的协方差逆矩阵bL无法直接与45×42的样本矩阵bx相乘,马氏距离的计算需要遵循维度匹配的运算逻辑。
最简PyTorch实现(适配固定尺寸)
假设bx是45×42的中心化样本矩阵(每行对应一个42维样本),bL是42×42的协方差逆矩阵,通过两步运算即可得到长度为45的1D张量:
import torch def fixed_mahalanobis(bL: torch.Tensor, bx: torch.Tensor) -> torch.Tensor: # 第一步:样本矩阵与协方差逆矩阵相乘,得到(45,42)中间张量 intermediate = bx @ bL # 第二步:每行与原样本行逐元素相乘后求和,得到每个样本的平方马氏距离 mahalanobis_sq = torch.sum(intermediate * bx, dim=1) return mahalanobis_sq
该实现完全规避维度错误,且因尺寸固定,无需处理通用批量逻辑,比PyTorch源码中的_batch_mahalanobis更简洁。
对应C++实现思路(固定尺寸优化)
由于尺寸固定为42和45,可手写循环实现,避免通用矩阵库的额外开销:
- 定义输入数组:
float bL[42][42](协方差逆)、float bx[45][42](样本) - 定义输出数组:
float result[45] - 逐样本计算:先求样本与协方差逆的乘积,再计算该乘积与原样本的点积,得到平方马氏距离
核心代码片段:
void fixed_mahalanobis(float bL[42][42], float bx[45][42], float result[45]) { for (int i = 0; i < 45; ++i) { float intermediate[42] = {0.0f}; // 计算bx[i]与bL的矩阵乘积 for (int j = 0; j < 42; ++j) { for (int k = 0; k < 42; ++k) { intermediate[j] += bx[i][k] * bL[k][j]; } } // 计算中间张量与原样本的点积,得到平方马氏距离 float dist_sq = 0.0f; for (int j = 0; j < 42; ++j) { dist_sq += intermediate[j] * bx[i][j]; } result[i] = dist_sq; } }
若需实际马氏距离,只需在最后对dist_sq取平方根即可。
内容的提问来源于stack exchange,提问作者Ant
相关产品推荐
相关产品推荐

