如何在3维PyTorch Tensor中选指定行并保留3维结构
这个问题我之前也碰到过!核心问题是直接用布尔索引a[i]会把所有符合条件的行从整个张量里“拉平”提取,完全丢失了原来按第一个维度分组的信息,所以最终得到的是2维张量。下面给你两种解决方案,根据你的实际需求选就行:
方案1:使用不规则嵌套张量(保留每行数量差异)
如果你的场景允许每个子2维张量的行数不固定,PyTorch的嵌套张量(Nested Tensor)正好能满足需求,它可以保持3维的逻辑结构,同时允许每个子张量的行数灵活变化:
import torch a = torch.rand(4,4,4) i = a[:,:,0] > 0.5 # 逐个处理每个子2维张量,提取符合条件的行 selected_rows = [a[k][i[k]] for k in range(a.size(0))] # 转换成嵌套张量,保持3维结构 result = torch.nested.nested_tensor(selected_rows) # 查看结果结构 print(result.shape) # 输出: torch.Size([4, None, 4]),None表示行数不固定 print(result)
这种方式完美保留了每个原始子张量的选择结果,不会破坏分组信息,而且不需要额外的padding操作。
方案2:Padding为固定形状的标准3维张量
如果你的下游代码必须要求输入是形状固定的标准3维张量,那我们可以把所有子张量的结果padding到同一行数(比如取所有子张量中的最大行数):
import torch a = torch.rand(4,4,4) i = a[:,:,0] > 0.5 # 先逐个提取每个子张量的目标行 selected_rows = [a[k][i[k]] for k in range(a.size(0))] # 找到所有子张量中的最大行数 max_row_count = max(t.size(0) for t in selected_rows) # 对每个子张量进行padding,补0到最大行数(padding参数格式:(左,右,上,下)) padded_rows = [ torch.nn.functional.pad(tensor, (0, 0, 0, max_row_count - tensor.size(0))) for tensor in selected_rows ] # 堆叠成固定形状的3维张量 result = torch.stack(padded_rows) print(result.shape) # 输出: torch.Size([4, max_row_count, 4]) print(result)
这样得到的就是标准的3维Tensor,所有子张量的行数统一,适合需要固定输入形状的场景。
内容的提问来源于stack exchange,提问作者Eric
相关产品推荐
相关产品推荐

