PyTorch torch.split_with_sizes Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.split_with_sizesIt is a function in PyTorch used to split a tensor according to specified sizes.

Function Definition

torch.split_with_sizes(input, split_sizes, dim=0)

Usage Example

Example

import torch

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

# Split by sizes [2, 3, 5]
result = torch.split_with_sizes(x, [2, 3, 5])
print("Split results:")
for i, t in enumerate(result):
    print(f" Chunk {i}: {t}")

# Split a 2D tensor by rows
y = torch.arange(12).reshape(4, 3)
print("nOriginal 2D tensor:")
print(y)

result = torch.split_with_sizes(y, [1, 2, 1], dim=0)
print("Split by rows [1, 2, 1]:")
for i, t in enumerate(result):
    print(f" Chunk {i}:n{t}")

The output is:

原始张量:
tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
分割结果:
  块 0: tensor([0, 1])
  Chunk 1: tensor([2, 3, 4])
  Chunk 2: tensor([5, 6, 7, 8, 9])

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

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions