PyTorch torch.concatenate Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.concatenateis a function in PyTorch used to concatenate multiple tensors along a specified dimension. It is the same astorch.catis the same function, used to concatenate multiple tensors along a specified dimension into a larger tensor.

Function Definition

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

Parameters:

  • tensors(Sequence of Tensor): The sequence of tensors to be concatenated. 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

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.concatenate([a, b, c])

print(result)

The output is:

tensor([1, 2, 3, 4, 5, 6])

Example

import torch

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

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

print("Shape of a:", a.shape)
print("Shape of b:", b.shape)
print("Shape of c:", c.shape)

The output is:

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

Note:torch.concatenateYestorch.catan alias, and the two have exactly the same functionality. In actual code, the more commonly used istorch.cat。


Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions