PyTorch torch.cat Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.catis a function in PyTorch used to concatenate multiple tensors along a specified dimension. It concatenates multiple tensors along the specified dimension into a larger tensor.

This is a very common operation in deep learning, for example in scenarios such as concatenating feature maps and merging data batches.

Function Definition

torch.cat(tensors, dim=0, out=None)

Parameters:

  • tensors(Sequence of Tensor): The sequence of tensors to concatenate. All tensors must have the same shape in all dimensions except the concatenation dimension.
  • dim(int, optional): The dimension along which to concatenate. Defaults to 0.
  • out(Tensor, optional): The output tensor.

Return Value:

  • torch.Tensor: Returns the concatenated tensor.

Usage Examples

Example 1: Concatenate along the first dimension

Example

import torch

# Create two tensors
a = torch.randn(2, 3)
b = torch.randn(2, 3)

# Concatenate along the first dimension (rows)
c = torch.cat([a, b], dim=0)

print("a's shape:", a.shape)
print("b's shape:", b.shape)
print("c's shape:", c.shape)
print(c)

The output result is:

a 的形状: torch.Size([2, 3])
b 的形状: torch.Size([2, 3])
c 的形状: torch.Size([4, 3])
tensor([[ 0.2532,  0.3643,  0.5341],
        [ 0.9578,  0.9086, -0.2847],
        [-0.7108, -0.0142,  0.7168],
        [-0.1542, -0.9841, -1.4945]])

Example 2: Concatenate along the second dimension

Example

import torch

# Create two tensors
a = torch.randn(2, 3)
b = torch.randn(2, 4)

# Concatenate along the second dimension (columns)
c = torch.cat([a, b], dim=1)

print("a's shape:", a.shape)
print("b's shape:", b.shape)
print("c's shape:", c.shape)

The output result is:

a 的形状: torch.Size([2, 3])
b 的形状: torch.Size([2, 4])
c 的形状: torch.Size([2, 7])

Example 3: Concatenate multiple tensors

Example

import torch

# Create multiple tensors
a = torch.tensor([1, 2])
b = torch.tensor([3, 4])
c = torch.tensor([5, 6])

# Concatenate multiple tensors
result = torch.cat([a, b, c])

print(result)

The output result is:

tensor([1, 2, 3, 4, 5, 6])
</p>

<h3>示例 4: 在神经网络中拼接特征</h3>

<div class="example">
<h2 class="example">实例</h2>
<div class="example_code">
<span style="color: Green;font-weight:bold;">import</span> torch<br />
<br />
<span style="color: #a50"># 模拟来自不同层的特征图</span><br />
feature1 <span style="color: Gray;">=</span> torch.<span style="color: #05a;">randn</span><span style="color: Olive;">&#40;</span><span style="color: Maroon;">1</span><span style="color: Gray;">,</span> <span style="color: Maroon;">64</span><span style="color: Gray;">,</span> <span style="color: Maroon;">32</span><span style="color: Gray;">,</span> <span style="color: Maroon;">32</span><span style="color: Olive;">&#41;</span> &nbsp;<span style="color: #a50"># 来自第一层的特征</span><br />
feature2 <span style="color: Gray;">=</span> torch.<span style="color: #05a;">randn</span><span style="color: Olive;">&#40;</span><span style="color: Maroon;">1</span><span style="color: Gray;">,</span> <span style="color: Maroon;">128</span><span style="color: Gray;">,</span> <span style="color: Maroon;">32</span><span style="color: Gray;">,</span> <span style="color: Maroon;">32</span><span style="color: Olive;">&#41;</span> <span style="color: #a50"># 来自第二层的特征</span><br />
<br />
<span style="color: #a50"># 在通道维度(dim=1)拼接特征</span><br />
combined <span style="color: Gray;">=</span> torch.<span style="color: #05a;">cat</span><span style="color: Olive;">&#40;</span><span style="color: Olive;">&#91;</span>feature1<span style="color: Gray;">,</span> feature2<span style="color: Olive;">&#93;</span><span style="color: Gray;">,</span> dim<span style="color: Gray;">=</span><span style="color: Maroon;">1</span><span style="color: Olive;">&#41;</span><br />
<br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;特征1 形状:&quot;</span><span style="color: Gray;">,</span> feature1.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;特征2 形状:&quot;</span><span style="color: Gray;">,</span> feature2.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;拼接后形状:&quot;</span><span style="color: Gray;">,</span> combined.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
</div>
</div>

<p>输出结果为:</p>
<pre>
特征1 形状: torch.Size([1, 64, 32, 32])
特征2 形状: torch.Size([1, 128, 32, 32])
拼接后形状: torch.Size([1, 192, 32, 32])

In neural networks,torch.catit is commonly used in structures such as Feature Pyramid Networks (FPN) to fuse features from different layers.


Difference between torch.cat and torch.stack

  • torch.cat: Concatenates along existing dimensions; the tensor shapes add up along the concatenation dimension.
  • torch.stack: Stacks along a new dimension, adding a new dimension.

Example

import torch

a = torch.randn(2, 3)
b = torch.randn(2, 3)

# Difference between cat and stack
cat_result = torch.cat([a, b], dim=0)
stack_result = torch.stack([a, b], dim=0)

print("cat result shape:", cat_result.shape)  # (4, 3)
print("stack result shape:", stack_result.shape)  # (2, 2, 3)

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions