在PyTorch里,按索引列表对张量进行分块切割是一个很常见的需求。比如我们有一个形状为[10, 3]的特征张量,想根据索引列表[0,1,2]、[3,4,5,6]、[7,8,9]把它分成三块,就可以通过内置函数或索引操作来实现。

使用 torch.tensor_split 按分割点切割
torch.tensor_split 可以接收一个表示分割位置的索引列表,把张量在指定维度上切开。注意它接收的是分割点,而不是每块的索引集合。
import torch
x = torch.arange(30).reshape(10, 3)
# 想在索引3和7处切开,得到三块
sections = [3, 7]
chunks = torch.tensor_split(x, sections, dim=0)
for i, c in enumerate(chunks):
print(f'块{i}形状: {c.shape}')
按索引列表提取子张量
如果已经明确知道每一块包含哪些索引,可以直接用长整型张量做高级索引。
import torch x = torch.arange(20).reshape(5, 4) idx1 = torch.tensor([0, 2]) idx2 = torch.tensor([1, 3, 4]) chunk_a = x[idx1] chunk_b = x[idx2] print(chunk_a) print(chunk_b)
使用 mask 掩码进行分块
当索引列表是由条件生成时,用布尔掩码更直观。
import torch
x = torch.randn(6, 2)
mask = x.sum(dim=1) > 0
pos = x[mask]
neg = x[~mask]
print('正样本块:', pos)
print('负样本块:', neg)
方法对比
| 方法 | 输入形式 | 适用场景 |
|---|---|---|
| tensor_split | 分割点列表 | 连续区间切分 |
| 高级索引 | 索引列表 | 离散索引提取 |
| 布尔掩码 | 条件mask | 按规则过滤分块 |
小结
按索引列表对张量分块时,连续分割优先用 torch.tensor_split,离散索引用高级索引,条件筛选用掩码。合理选择能让代码清晰且运行高效。
PyTorchtensor_chunkindex_listtensor_split张量分块修改时间:2026-07-25 13:30:16