torch.chunk in PyTorch Explained | Generated by AI

Home 2026.07

Question: What does torch.chunk do in PyTorch?

Answer: Splits a tensor into a specific number of chunks along a given dimension. If the dimension size doesn’t divide evenly, all chunks are equal size except the last one, which is smaller.

import torch

x = torch.arange(10)
chunks = torch.chunk(x, 3)  # split into 3 pieces along dim=0
for c in chunks:
    print(c, c.shape)

# tensor([0, 1, 2, 3]) torch.Size([4])
# tensor([4, 5, 6, 7]) torch.Size([4])
# tensor([8, 9])       torch.Size([2])   <- last chunk smaller, not equal

Key signature:

torch.chunk(input, chunks, dim=0)

Comparison you’ll actually hit in transformer code — splitting QKV projections:

qkv = torch.randn(2, 10, 3 * 768)  # (batch, seq, 3*d_model)
q, k, v = qkv.chunk(3, dim=-1)     # each (2, 10, 768)

This is the standard pattern in attention implementations (e.g. GPT-2 c_attn output) — one big linear projection, then chunk(3, dim=-1) to split into Q, K, V instead of three separate linear layers.

Gotcha: if dim_size % chunks != 0, chunk sizes are uneven (ceil-based), whereas torch.split(x, size) lets you control exact size per split and torch.tensor_split guarantees exactly n chunks even when it doesn’t divide evenly (redistributing the remainder across the first few chunks instead of dumping it all in the last one).

References:


Generated by AI. Curating and sharing still takes effort. If you find it useful, feel free to donate. WeChat: @lzwjavaWeChat QR · X: @lzwjava · Say hi 👋

Back Donate