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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 09:15:41