PyTorch torch.tensor_split Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.tensor_splitIt is a function in PyTorch used to split tensors by index or number of segments.

Function Definition

torch.tensor_split(input, indices_or_sections, dim=0)

Usage Example

Example

import torch

x = torch.arange(12)
print("Original tensor:")
print(x)

# Split by index
result = torch.tensor_split(x, [2, 5, 8])
print("Split by indices [2, 5, 8]:")
for i, t in enumerate(result):
    print(f" Block {i}: {t}")

# Split by number of segments (average into 3 parts)
y = torch.arange(9)
result = torch.tensor_split(y, 3)
print("nAverage split into 3 parts:")
for i, t in enumerate(result):
    print(f" Block {i}: {t}")

# Split the 2D tensor by columns
z = torch.arange(12).reshape(3, 4)
print("n2D tensor:")
print(z)

result = torch.tensor_split(z, 2, dim=1)
print("Split into 2 parts by column:")
for i, t in enumerate(result):
    print(f" Block {i}:n{t}")

The output result is:

原始张量:
tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11])
按索引 [2, 5, 8] 分割:
  块 0: tensor([0, 1])
  Chunk 1: tensor([2, 3, 4])
  Chunk 2: tensor([5, 6, 7])
  Chunk 3: tensor([ 8,  9, 10, 11])

平均分为 3 份:
  块 0: tensor([0, 1, 2])
  Chunk 1: tensor([3, 4, 5])
  Chunk 2: tensor([6, 7, 8])

二维张量:
tensor([[ 0,  1,  2,  3],
        [ 4,  5,  6,  7],
        [ 8,  9, 10, 11]])
按列分为 2 份:
  块 0:
tensor([[0, 1],
        [4, 5],
        [8, 9]])
  Chunk 1:
tensor([[ 2,  3],
        [ 6,  7],
        [10, 11]])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions