PyTorch torch.flatten Function


Pytorch torch 参考手册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)

The output result is:

原始形状: torch.Size([2, 3, 4])
展平后: torch.Size([24])
从 dim=1 展平: torch.Size([2, 12])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions