PyTorch中实现Pandas groupby.apply(list)的等价方法
问题描述
我有如下PyTorch张量long_format:
tensor([[ 1., 1.], [ 1., 2.], [ 1., 3.], [ 1., 4.], [ 0., 5.], [ 0., 6.], [ 0., 7.], [ 1., 8.], [ 0., 9.], [ 0., 10.]])
希望按第一列分组,将第二列存储为张量(各分组大小可不同),预期结果如下:
[tensor([ 1., 2., 3., 4., 8.]), tensor([ 5., 6., 7., 9., 10.])]
需要仅使用PyTorch算子的优雅实现,出于可追踪性考虑,要避免for循环。此前用for循环+空张量列表的实现会导致追踪错误(不同输入得到相同结果),代码如下:
n_groups = 2 inverted = [torch.empty([0]) for _ in range(n_groups)] for index, value in long_format: value = value.unsqueeze(dim=0) index = index.int() if type(inverted[index]) != torch.Tensor: inverted[index] = value else: inverted[index] = torch.cat((inverted[index], value))
解决方案
可以通过排序+分割点定位+张量分割的纯PyTorch算子组合实现,完全避免循环且支持自动微分/追踪:
# 提取分组键和对应的值 group_keys = long_format[:, 0] values = long_format[:, 1] # 按分组键排序,确保同组元素连续 sorted_keys, sorted_indices = torch.sort(group_keys) sorted_values = values[sorted_indices] # 找到分组之间的分割点 split_mask = sorted_keys[1:] != sorted_keys[:-1] split_points = torch.where(split_mask)[0] + 1 # 分割点位置后移一位 split_sizes = torch.diff(torch.cat([torch.tensor([0]), split_points, torch.tensor([len(sorted_values)])])) # 按分割点拆分张量得到分组结果 groups = torch.split(sorted_values, split_sizes.tolist())
执行后groups就是预期的分组张量列表,且所有操作均为PyTorch原生算子,可被正确追踪。
如果已知分组键只有0和1,也可以用更简洁的布尔索引方式(同样可追踪):
groups = [ values[group_keys == 1], values[group_keys == 0] ]
内容的提问来源于stack exchange,提问作者Trevor Waite
相关产品推荐
相关产品推荐

