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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 17:55:19