PyTorch torch.unbind Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.unbindIt is a PyTorch function used to split a tensor into a tuple along a specified dimension.

Function Definition

torch.unbind(input, dim=0)

Usage Example

Example

import torch

x = torch.arange(12).reshape(3, 4)
print(Original tensor:)
print(x)

result = torch.unbind(x, dim=0)
print(Split along dim=0:)
for i, t in enumerate(result):
    print(fBlock {i}: {t})

result = torch.unbind(x, dim=1)
print(Split along dim=1:)
for i, t in enumerate(result):
    print(fBlock {i}: {t})

The output result is:

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

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions