PyTorch torch.stack Function
Pytorch torch Reference Manual
torch.stackIt is a function in PyTorch used to stack multiple tensors along a new dimension. It creates a new dimension and places all input tensors along this new dimension.
This is commonly used in deep learning for scenarios such as creating batch data and stacking outputs from multiple models.
Function Definition
torch.stack(tensors, dim=0, out=None)
Parameters:
tensors(Sequence of Tensor): The sequence of tensors to stack. All tensors must have the same shape.dim(int, optional): The dimension along which to stack, defaults to 0. The new dimension will be inserted at this position.out(Tensor, optional): The output tensor.
Return Value:
torch.Tensor: Returns the stacked tensor.
Usage Examples
Example 1: Basic Stacking
Example
import torch
# Create two one-dimensional tensors
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
# Stack along a new dimension
c = torch.stack([a, b])
print("Shape of a:", a.shape)
print("Shape of b:", b.shape)
print("Shape of c:", c.shape)
print(c)
# Create two one-dimensional tensors
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
# Stack along a new dimension
c = torch.stack([a, b])
print("Shape of a:", a.shape)
print("Shape of b:", b.shape)
print("Shape of c:", c.shape)
print(c)
The output result is:
a 的形状: torch.Size([3])
b 的形状: torch.Size([3])
c 的形状: torch.Size([2, 3])
tensor([[1, 2, 3],
[4, 5, 6]])
Example 2: Specifying the Stack Dimension
Example
import torch
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
# Stack along dim=1
c = torch.stack([a, b], dim=1)
print("Shape of c:", c.shape)
print(c)
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
# Stack along dim=1
c = torch.stack([a, b], dim=1)
print("Shape of c:", c.shape)
print(c)
The output result is:
c 的形状: torch.Size([3, 2])
tensor([[1, 4],
[2, 5],
[3, 6]])
Example 3: Stacking Multiple Tensors
Example
import torch
# Create multiple tensors
tensors = [torch.tensor([i, i+1, i+2]) for i in range(5)]
# Stack all tensors
result = torch.stack(tensors)
print("Result shape:", result.shape)
print(result)
# Create multiple tensors
tensors = [torch.tensor([i, i+1, i+2]) for i in range(5)]
# Stack all tensors
result = torch.stack(tensors)
print("Result shape:", result.shape)
print(result)
The output result is:
结果形状: torch.Size([5, 3])
tensor([[0, 1, 2],
[1, 2, 3],
[2, 3, 4],
[3, 4, 5],
[4, 5, 6]])
Example 4: Saving Multiple States in a Neural Network
Example
import torch
# Simulate saving loss values for multiple epochs
losses = []
for epoch in range(5):
loss = torch.tensor([epoch * 0.1, (epoch + 1) * 0.1])
losses.append(loss)
# Stack the losses of all epochs
all_losses = torch.stack(losses)
print("Loss shape across epochs:", all_losses.shape)
print(all_losses)
# Simulate saving loss values for multiple epochs
losses = []
for epoch in range(5):
loss = torch.tensor([epoch * 0.1, (epoch + 1) * 0.1])
losses.append(loss)
# Stack the losses of all epochs
all_losses = torch.stack(losses)
print("Loss shape across epochs:", all_losses.shape)
print(all_losses)
The output result is:
各 epoch 损失形状: torch.Size([5, 2])
tensor([[0.0000, 0.1000],
[0.1000, 0.2000],
[0.2000, 0.3000],
[0.3000, 0.4000],
[0.4000, 0.5000]])
Difference Between torch.stack and torch.cat
Example
import torch
a = torch.randn(2, 3)
b = torch.randn(2, 3)
# stack creates a new dimension
stack_result = torch.stack([a, b])
print("stack result shape:", stack_result.shape)
# cat does not create a new dimension
cat_result = torch.cat([a, b], dim=0)
print("cat result shape:", cat_result.shape)
a = torch.randn(2, 3)
b = torch.randn(2, 3)
# stack creates a new dimension
stack_result = torch.stack([a, b])
print("stack result shape:", stack_result.shape)
# cat does not create a new dimension
cat_result = torch.cat([a, b], dim=0)
print("cat result shape:", cat_result.shape)
The output result is:
stack 结果形状: torch.Size([2, 2, 3]) cat 结果形状: torch.Size([4, 3])
torch.stack: Stacks along a new dimension; the input tensors must have exactly the same shape, and a new dimension is added.torch.cat: Concatenates along an existing dimension; the input tensors may differ in the concatenation dimension, and no new dimension is added.
Other Extensions