,


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.atleast_3dis a function in PyTorch used to convert input tensors to at least 3 dimensions. If the input has fewer than 3 dimensions, dimensions are automatically added to make it 3-dimensional.

]]) Shape: torch.Size([3, 1, 1])

torch.atleast_3d(*tensors)

Usage Example

]]) Shape: torch.Size([1, 1, 1])

import torch

# Convert scalar to 3D tensor
x = torch.atleast_3d(5)
print("After scalar conversion:", x, "Shape:", x.shape)
# Output: After scalar conversion: tensor([[

# Convert 1D tensor to 3D
x = torch.tensor([1, 2, 3])
y = torch.atleast_3d(x)
print("1D to 3D:", y, "Shape:", y.shape)
# Output: 1D to 3D: tensor([[

# Convert 2D tensor to 3D
x = torch.tensor([[1, 2], [3, 4]])
y = torch.atleast_3d(x)
print("2D to 3D:", y, "Shape:", y.shape)
# Output: 2D to 3D: tensor([[[1, 2], [3, 4]]]) Shape: torch.Size([1, 2, 2])

# A tensor that is already 3D remains unchanged
x = torch.tensor([[[1, 2], [3, 4]]])
y = torch.atleast_3d(x)
print("3D tensor:", y, "Shape:", y.shape)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions