torch.cat joins a sequence of tensors along a dimension that already exists. The shapes have to match everywhere except that one dimension, and the result never gains an axis.
torch.cat((a, b), dim=0) # more rows
torch.cat((a, b), dim=1) # more columns
That last part is the entire difference between torch.cat and torch.stack, which does add an axis.
The runs below are on PyTorch 2.14.0 on Python 3.12.5, and the screenshots are the Command Prompt they came out of.
What dim does to the shape
Two tensors of shape (2, 3) can be joined two ways. Along dim=0 you get (4, 3), along dim=1 you get (2, 6):
import torch
a = torch.tensor([[1, 2, 3],
[4, 5, 6]])
b = torch.tensor([[7, 8, 9],
[10, 11, 12]])
print("a:", tuple(a.shape), " b:", tuple(b.shape))
rows = torch.cat((a, b), dim=0) # stack them on top of each other
print("dim=0 ->", tuple(rows.shape))
print(rows)
cols = torch.cat((a, b), dim=1) # glue them side by side
print("dim=1 ->", tuple(cols.shape))
print(cols)
Output:
a: (2, 3) b: (2, 3)
dim=0 -> (4, 3)
tensor([[ 1, 2, 3],
[ 4, 5, 6],
[ 7, 8, 9],
[10, 11, 12]])
dim=1 -> (2, 6)
tensor([[ 1, 2, 3, 7, 8, 9],
[ 4, 5, 6, 10, 11, 12]])
Only the dim you name changes size. Every other dimension is carried through untouched, which is why the others have to agree in the first place.
import torch
x = torch.zeros(2, 3, 4)
y = torch.ones(2, 3, 4)
print("each input:", tuple(x.shape))
for dim in (0, 1, 2, -1):
print(f" dim={dim:>2} -> {tuple(torch.cat((x, y), dim=dim).shape)}")
# dim=-1 is just the last dimension, counted from the right
print("dim=-1 is the same as dim=2:",
torch.equal(torch.cat((x, y), dim=-1), torch.cat((x, y), dim=2)))
Output:
each input: (2, 3, 4)
dim= 0 -> (4, 3, 4)
dim= 1 -> (2, 6, 4)
dim= 2 -> (2, 3, 8)
dim=-1 -> (2, 3, 8)
dim=-1 is the same as dim=2: True
Negative values count from the right, so dim=-1 is the last dimension. That’s handy when you write a function that shouldn’t care how many dimensions it’s handed.
Sizes of tensors must match except in dimension…
This is the error you’ll actually meet, and it reads backwards until you know the trick:
import torch
a = torch.ones(2, 3)
b = torch.ones(2, 4) # 4 columns, not 3
torch.cat((a, b), dim=0)
What PyTorch raises:
Traceback (most recent call last):
File "C:\pyguides\torch_cat_size_mismatch.py", line 6, in <module>
torch.cat((a, b), dim=0)
RuntimeError: Sizes of tensors must match except in dimension 0. Expected size 3 but got size 4 for tensor number 1 in the list.
Read it as two separate facts. “except in dimension 0” is the dim you passed. “expected size 3 but got size 4” is the mismatch, and it’s in some other dimension.
So the fix is usually one of two things: join along the dimension that actually differs, or reshape first.
import torch
a = torch.ones(2, 3)
b = torch.ones(2, 4)
# the sizes differ in dimension 1, so dimension 1 is the one to join on
print(tuple(torch.cat((a, b), dim=1).shape))
Output:
(2, 7)
tensor number 1 is a zero-based index into your sequence, so it’s the second tensor you passed.
On a list of twenty that number is the quickest way to find the odd one out. It helps to know your dtypes and shapes before they surprise you.
torch.concat and torch.concatenate are the same function
PyTorch carries three names for this operation, because concatenate is what NumPy calls it.
import torch
a = torch.tensor([1, 2])
b = torch.tensor([3, 4])
print("cat: ", torch.cat((a, b)))
print("concat: ", torch.concat((a, b)))
print("concatenate:", torch.concatenate((a, b)))
# three separate function objects that do the same work
print("concat is cat:", torch.concat is torch.cat,
"| same result:", torch.equal(torch.concat((a, b)), torch.cat((a, b))))
# every one of them takes dim= and the NumPy-style axis=
m, n = torch.ones(2, 3), torch.ones(2, 3)
print("cat(axis=1): ", tuple(torch.cat((m, n), axis=1).shape))
print("concatenate(dim=1):", tuple(torch.concatenate((m, n), dim=1).shape))
Output:
cat: tensor([1, 2, 3, 4])
concat: tensor([1, 2, 3, 4])
concatenate: tensor([1, 2, 3, 4])
concat is cat: False | same result: True
cat(axis=1): (2, 6)
concatenate(dim=1): (2, 6)
They aren’t the same object — torch.concat is torch.cat comes back False — but they do identical work and take identical arguments.
All three accept dim= and NumPy’s axis=, even though the reference documents concatenate with axis and the other two with dim.
Use whichever reads better where you are. torch.cat is the name you’ll meet in most PyTorch code and in the error messages.
torch.cat vs torch.stack
The one-line version: cat extends an axis, stack creates one.
import torch
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
joined = torch.cat((a, b)) # extends an axis that already exists
stacked = torch.stack((a, b)) # invents a new axis
print("cat ", tuple(joined.shape), joined)
print("stack", tuple(stacked.shape))
print(stacked)
# stack is cat with an unsqueeze in front of it
by_hand = torch.cat((a.unsqueeze(0), b.unsqueeze(0)), dim=0)
print("stack == cat of unsqueezed tensors:", torch.equal(stacked, by_hand))
Output:
cat (6,) tensor([1, 2, 3, 4, 5, 6])
stack (2, 3)
tensor([[1, 2, 3],
[4, 5, 6]])
stack == cat of unsqueezed tensors: True
cat gives 6, stack gives 2×3.torch.cat | torch.stack | |
|---|---|---|
| Number of dimensions | Unchanged | One more |
| Input shapes | Match except in dim | Must match exactly |
Two (3,) tensors | (6,) | (2, 3) |
| Typical use | Growing a batch | Turning a list into a batch |
If you want a batch dimension out of a list of equally shaped samples, you want torch.stack. If you’re appending more samples to a batch that already has one, you want cat.
Why a generator fails
torch.cat wants a real sequence it can measure and index. Hand it a generator expression and it stops before doing any work:
import torch
parts = [torch.ones(1, 3) for _ in range(4)]
# a generator has no length and cannot be indexed, so torch.cat refuses it
torch.cat(p for p in parts)
The error:
Traceback (most recent call last):
File "C:\pyguides\runs\torchcat\ex_generator.py", line 6, in <module>
torch.cat(p for p in parts)
TypeError: cat(): argument 'tensors' (position 1) must be tuple of Tensors, not generator
Wrap it in list() or tuple() and it’s fine. The same example also shows dtype promotion and what happens to a zero-length tensor:
import torch
parts = (torch.ones(1, 3) for _ in range(4))
print(tuple(torch.cat(list(parts)).shape)) # list() drains the generator first
mixed = [torch.ones(2, 2, dtype=torch.float32), torch.ones(2, 2, dtype=torch.float64)]
print("dtype promotion:", torch.cat(mixed).dtype)
empty = torch.empty(0, 3)
print("empty tensors are skipped:", tuple(torch.cat((empty, torch.ones(2, 3))).shape))
Output:
(4, 3)
dtype promotion: torch.float64
empty tensors are skipped: (2, 3)
Mixed dtypes get promoted rather than rejected, and an empty tensor contributes nothing but still has to agree on the other dimensions.
Never call cat inside a loop
This is the one that quietly costs real training time. Growing a tensor by concatenating inside a loop copies everything that came before, on every iteration:
import torch, timeit
parts = [torch.randn(64, 128) for _ in range(200)]
def grow():
out = torch.empty(0, 128)
for p in parts:
out = torch.cat((out, p), dim=0) # copies everything, every time
return out
def once():
return torch.cat(parts, dim=0) # one allocation, one copy
timings = {}
for name, fn in (("cat inside the loop", grow), ("collect, then one cat", once)):
timings[name] = timeit.timeit(fn, number=20) / 20 * 1000
print(f"{name:<23} {timings[name]:7.2f} ms")
print(f"the loop is {timings['cat inside the loop'] / timings['collect, then one cat']:.0f}x slower")
print("identical result:", torch.equal(grow(), once()))
Output:
cat inside the loop 14.67 ms
collect, then one cat 0.55 ms
the loop is 27x slower
identical result: True
Two hundred tensors is enough to put more than an order of magnitude between them, and the gap widens as the list grows.
Each call copies everything already accumulated, so the work is quadratic in the number of pieces.
Collect the pieces in a plain Python list, then call torch.cat once at the end. It’s the same habit that makes converting tensors to NumPy cheap: do the expensive thing once, not per item.
Other PyTorch guides on this site:
- torch.stack in PyTorch
- Convert a PyTorch tensor to NumPy
- PyTorch nn.Linear
- PyTorch nn.Conv2d
- NumPy data types in Python
Frequently asked questions
What does torch.cat do?
It joins a sequence of tensors along an existing dimension. The output has the same number of dimensions as the inputs, with one of them longer. The full signature is in the torch.cat reference.
What is the difference between torch.cat and torch.stack?
cat extends a dimension that already exists, stack adds a new one. Two tensors of shape (3,) give (6,) with cat and (2, 3) with stack.
Why do I get “Sizes of tensors must match except in dimension 0”?
Dimension 0 is the one you asked to join on. The sizes disagree in a different dimension, so either join along that one instead, or reshape first.
Is torch.concat the same as torch.cat?
It does the same thing, and so does torch.concatenate. They are three separate function objects, so torch.concat is torch.cat is False, but the behaviour and the arguments match. All three take dim= and axis=.
Can I pass a generator to torch.cat?
No. It needs a sequence it can index, so a generator raises a TypeError. Wrap it in list() first.
What dim should I use?
The dimension you want to get longer. dim=0 adds rows, dim=1 adds columns, and dim=-1 means the last dimension whatever the rank.
Does torch.cat copy the data?
Yes. It allocates a new tensor and copies every input into it, which is why calling it inside a loop is slow and calling it once is not.
Bijay Kumar is a 13-time Microsoft MVP with more than 18 years in software development, and the founder of Python Guides and TSinfo Technologies. He started out building .NET and SharePoint solutions at HP, TCS and KPIT before moving into Python, machine learning and AI, and he also builds web apps with TypeScript and React. He writes the tutorials here himself, and every example is run before publishing so you see the real output. More about Bijay · Microsoft MVP profile · LinkedIn