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

如何高效计算对称张量中按类别分组的行列平均值?

问题描述

我有一个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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 01:32:58