PyTorch torch.chunk Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.chunkIt is a function in PyTorch used to split a tensor into chunks along a specified dimension.

Function Definition

torch.chunk(tensor, chunks, dim)

Usage Example

Example

import torch

x = torch.arange(12).reshape(3, 4)

# Split into 3 chunks
result = torch.chunk(x, 3, dim=0)

print("Number of chunks:", len(result))
for i, t in enumerate(result):
    print(f"Chunk {i}:", t.shape)

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions