PyTorch torch.hsplit function


Pytorch torch 参考手册Pytorch torch reference manual

torch.hsplitIt is a function in PyTorch used to split tensors horizontally (along columns).

Function definition

torch.hsplit(input, indices_or_sections)

Usage example

Example

import torch

# Horizontal split of a 1D tensor
x = torch.arange(10)
print("Original 1D tensor:")
print(x)

result = torch.hsplit(x, 2)
print("Split into 2 equal parts:")
for i, t in enumerate(result):
    print(f" Block {i}: {t}")

# Horizontal split of a 2D tensor
y = torch.arange(12).reshape(3, 4)
print("\nOriginal 2D tensor:")
print(y)

result = torch.hsplit(y, 2)
print("Split into 2 parts along columns:")
for i, t in enumerate(result):
    print(f" Block {i}:\n{t}")

# Split by indices
result = torch.hsplit(y, [1, 3])
print("\nSplit by indices [1, 3]:")
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])
平均分为 2 份:
  块 0: tensor([0, 1, 2, 3, 4])
  Chunk 1: tensor([5, 6, 7, 8, 9])

原始二维张量:
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]])

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

Pytorch torch 参考手册Pytorch torch reference manual

Other extensions