PyTorch torch.flatten Function
Pytorch torch Reference Manual
torch.flattenis a function in PyTorch used to flatten tensors.
Function Definition
torch.flatten(input, start_dim, end_dim, out)
Usage Example
Example
import torch
x = torch.randn(2, 3, 4)
# Flatten to one dimension
y = torch.flatten(x)
print("Original shape:", x.shape)
print("After flattening:", y.shape)
# Flatten from a specified dimension
z = torch.flatten(x, start_dim=1)
print("Flatten from dim=1:", z.shape)
x = torch.randn(2, 3, 4)
# Flatten to one dimension
y = torch.flatten(x)
print("Original shape:", x.shape)
print("After flattening:", y.shape)
# Flatten from a specified dimension
z = torch.flatten(x, start_dim=1)
print("Flatten from dim=1:", z.shape)
The output result is:
原始形状: torch.Size([2, 3, 4]) 展平后: torch.Size([24]) 从 dim=1 展平: torch.Size([2, 12])
Other Extensions