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

如何基于另一数组的唯一索引拆分PyTorch多维张量?

基于PyTorch张量的维度拆分问题解决

问题描述

现有两个PyTorch张量a和b:

import torch
torch.manual_seed(0) # 保证可复现

a = torch.rand(size = (5, 10, 1))
b = torch.tensor([3, 3, 1, 5, 3, 1, 0, 2, 1, 2])

想要基于b里的唯一值,拆分a的第2维度(也就是dim=1),预期得到一个列表:列表元素个数等于b中唯一值的数量,每个元素是形状为(5, 对应唯一值的元素个数, 1)的张量。

试过的代码:

# 获取b的唯一值和唯一索引
unique_values, unique_indices = torch.unique(b, return_inverse = True)

# 基于唯一索引拆分a的dim=1维度
l = torch.tensor_split(a, unique_indices, dim = 1)

结果里出现了tensor([], size=(5, 0, 1))这类空张量,需要搞清楚原因,以及怎么实现正确的需求。

问题原因

torch.unique(..., return_inverse=True)返回的unique_indices是原张量b中每个元素对应的唯一值索引,比如这个例子里它是tensor([3, 3, 1, 4, 3, 1, 0, 2, 1, 2])。但torch.tensor_split的第二个参数需要的是切分的位置列表——也就是你要从哪些维度索引处把张量切开。直接把unique_indices传进去,相当于让函数在这些索引位置挨个切分,很多位置连续重复或者顺序混乱,自然就切出空张量了。

正确实现方法

方法1:遍历唯一值,筛选对应索引切片

最直观的方式就是遍历b的每个唯一值,找出该值在b里的所有位置,再用这些位置对a的dim=1维度做切片:

import torch
torch.manual_seed(0)

a = torch.rand(size=(5, 10, 1))
b = torch.tensor([3, 3, 1, 5, 3, 1, 0, 2, 1, 2])

unique_values = torch.unique(b)
result = []
for val in unique_values:
    # 找出b中等于当前唯一值的所有索引
    idx = torch.where(b == val)[0]
    # 按索引切片a的dim=1维度
    sliced_tensor = a[:, idx, :]
    result.append(sliced_tensor)

# 验证结果
for i, tensor in enumerate(result):
    print(f"唯一值{unique_values[i]}对应的张量形状:{tensor.shape}")

输出:

唯一值0对应的张量形状:torch.Size([5, 1, 1])
唯一值1对应的张量形状:torch.Size([5, 3, 1])
唯一值2对应的张量形状:torch.Size([5, 2, 1])
唯一值3对应的张量形状:torch.Size([5, 3, 1])
唯一值5对应的张量形状:torch.Size([5, 1, 1])

方法2:分组排序后拆分(更高效)

如果要处理大规模张量,这个方法更高效:先按unique_indices给a排序,再找出分组的边界位置做拆分:

import torch
torch.manual_seed(0)

a = torch.rand(size=(5, 10, 1))
b = torch.tensor([3, 3, 1, 5, 3, 1, 0, 2, 1, 2])

unique_values, unique_indices = torch.unique(b, return_inverse=True)
# 按unique_indices排序,得到排序后的索引
sorted_idx = torch.argsort(unique_indices)
# 对a做排序,让同一组的元素集中在一起
sorted_a = a[:, sorted_idx, :]
# 找出分组的边界(每个分组结束的下一个位置)
sorted_inverse = unique_indices[sorted_idx]
split_pos = torch.where(sorted_inverse[1:] != sorted_inverse[:-1])[0] + 1
# 拆分得到结果
result = torch.tensor_split(sorted_a, split_pos, dim=1)

# 验证结果
for i, tensor in enumerate(result):
    print(f"唯一值{unique_values[i]}对应的张量形状:{tensor.shape}")

这个方法的结果和方法1完全一致,但避免了循环里的多次切片,处理大张量时速度更快。

内容的提问来源于stack exchange,提问作者bird

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 21:30:33