PyTorch torch.cat Function
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)
# 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)
# 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)
# 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;">(</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;">)</span> <span style="color: #a50"># 来自第一层的特征</span><br /> feature2 <span style="color: Gray;">=</span> torch.<span style="color: #05a;">randn</span><span style="color: Olive;">(</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;">)</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;">(</span><span style="color: Olive;">[</span>feature1<span style="color: Gray;">,</span> feature2<span style="color: Olive;">]</span><span style="color: Gray;">,</span> dim<span style="color: Gray;">=</span><span style="color: Maroon;">1</span><span style="color: Olive;">)</span><br /> <br /> <span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"特征1 形状:"</span><span style="color: Gray;">,</span> feature1.<span style="color: #05a;">shape</span><span style="color: Olive;">)</span><br /> <span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"特征2 形状:"</span><span style="color: Gray;">,</span> feature2.<span style="color: #05a;">shape</span><span style="color: Olive;">)</span><br /> <span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"拼接后形状:"</span><span style="color: Gray;">,</span> combined.<span style="color: #05a;">shape</span><span style="color: Olive;">)</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)
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)
Other Extensions