PyTorch torch.chunk Function
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)
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)
Other Extensions