torch.cat in PyTorch: Joining Tensors Along a Dimension

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]])
Command Prompt showing torch.cat joining two 2x3 PyTorch tensors along dim 0 to give a 4x3 tensor and along dim 1 to give a 2x6 tensor
Same two tensors, same call, one argument different.

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.
Command Prompt showing the PyTorch RuntimeError that sizes of tensors must match except in dimension 0
The message names the dimension you asked for, not the one that’s wrong.

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
Command Prompt comparing torch.cat and torch.stack on two 1D PyTorch tensors, showing shapes 6 and 2 by 3
Two tensors of 3 elements: cat gives 6, stack gives 2×3.
torch.cattorch.stack
Number of dimensionsUnchangedOne more
Input shapesMatch except in dimMust match exactly
Two (3,) tensors(6,)(2, 3)
Typical useGrowing a batchTurning 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
Command Prompt comparing the time taken by calling torch.cat inside a loop against collecting tensors and calling it once
Two hundred tensors, identical output, measured over 20 runs each.

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:

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.