如何高效计算对称张量中按类别分组的行列平均值?
问题描述
我有一个200×200的对称tensor(通过将另一个200×300矩阵与其转置相乘得到)。每一行(及对应的列)都属于一个特定类别,类别及其对应的索引存储在一个字典中,示例代码及输出如下:
print(data.shape) # 存储数据的tensor print(cat_idx) # 类别索引字典
输出:
torch.Size([200, 200]) {'Album': [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10], 'Animal': [11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21], 'Artist': ... ...
上述数据表示,行(及列)0至10中的值是Album类别的距离数据,总共有15个类别。
现在我希望将其转换为一个15×15的tensor,其中每个元素对应类别的平均距离:即第1行第1列是Album类所有点之间的平均距离,第1行第2列是Album类别与Animal类别所有点之间的平均距离,以此类推。
我目前通过两层循环实现了需求(代码如下):
mean_score_by_cat_row = torch.empty((0, cosine_scores.shape[1])) for i, cat in enumerate(categories): cat_scores = data[cat_idx[cat]] cat_row_mean = torch.mean(cat_scores, 0, True) mean_score_by_cat_row = torch.cat((mean_score_by_cat_row, cat_row_mean), 0) mean_score_by_cat = torch.empty((len(cat_idx), 0)) for i, cat in enumerate(categories): cat_scores = mean_score_by_cat_row[:, cat_idx[cat]] cat_col_mean = torch.mean(cat_scores, 1, True) mean_score_by_cat = torch.cat((mean_score_by_cat, cat_col_mean), 1)
但我认为这并非处理tensor的最优方式,请问是否有更高效的实现方法?当前的实现是否正确?
解答
1. 当前实现的正确性
你的实现是正确的。第一层循环对每个类别提取对应行并计算均值,得到每个类别到所有样本的平均距离;第二层循环再对每个类别对应的列计算均值,最终得到符合需求的类别间平均距离矩阵。不过循环结合torch.cat的方式会频繁触发内存分配,效率确实偏低。
2. 更高效的实现方案
方案一:纯PyTorch原生实现(矩阵乘法+广播)
通过构建类别指示矩阵,利用矩阵乘法一次性计算类别间的总距离和,再除以样本数量乘积得到均值,全程无循环:
import torch # 假设categories是类别名称列表,比如list(cat_idx.keys()) categories = list(cat_idx.keys()) num_cats = len(categories) device = data.device dtype = data.dtype # 构建类别指示矩阵:shape [200, 15],indicator[i,c] = 1表示第i个样本属于第c类 indicator = torch.zeros((data.shape[0], num_cats), dtype=dtype, device=device) for c_idx, cat in enumerate(categories): indicator[cat_idx[cat], c_idx] = 1.0 # 计算每个类别的样本数量 counts = indicator.sum(dim=0) # shape [15] # 计算类别间总距离和,再除以样本数乘积得到平均距离 total_dist = indicator.T @ data @ indicator mean_dist = total_dist / (counts.unsqueeze(0) * counts.unsqueeze(1)) # (可选)如果需要排除样本自身的距离(即i=j的情况) # 提取原tensor对角线的自身距离和 diag_self_dist = torch.diag(data).unsqueeze(1) @ indicator.T # shape [1,15] # 总距离减去同类自身的距离和 total_dist_no_self = total_dist - torch.diag(diag_self_dist.squeeze(0)) # 计算有效样本对数量:n*(n-1) count_pairs = counts.unsqueeze(0)*counts.unsqueeze(1) - torch.diag(counts) mean_dist_no_self = total_dist_no_self / count_pairs
方案二:使用torch_scatter库(简洁高效)
如果安装了torch_scatter(PyTorch生态常用库),可以用scatter_mean算子直接完成两次维度上的类别均值计算,代码更简洁且效率更高:
from torch_scatter import scatter_mean categories = list(cat_idx.keys()) num_cats = len(categories) device = data.device # 构建每个样本的类别标签数组:shape [200] labels = torch.zeros(data.shape[0], dtype=torch.long, device=device) for c_idx, cat in enumerate(categories): labels[cat_idx[cat]] = c_idx # 先按类别对行维度取均值,得到[15,200]的矩阵 row_mean = scatter_mean(data, labels, dim=0) # 再按类别对列维度取均值,得到最终的[15,15]类别平均距离矩阵 mean_dist = scatter_mean(row_mean, labels, dim=1)
3. 效率对比
- 原方案:循环中每次
torch.cat都会开辟新内存,时间复杂度为O(N*C)(N为样本数,C为类别数),内存开销较大。 - 矩阵乘法方案:时间复杂度为O(N² + N*C + C²),无需额外依赖,适合样本数规模适中的场景。
- torch_scatter方案:时间复杂度为O(N²),依托底层优化算子,是三种方案中效率最高的,代码也最简洁。
内容的提问来源于stack exchange,提问作者Al Wonder
相关产品推荐
相关产品推荐

