PyTorch-----torch.flatten()函数

torch.flatten() 是 PyTorch 中的一个函数,用于将输入张量展平为一维张量。它的语法如下:

python 复制代码
torch.flatten(input, start_dim=0, end_dim=-1)
  • input:要展平的输入张量。
  • start_dim(可选):指定从哪个维度开始展平。默认为 0。
  • end_dim(可选):指定从哪个维度结束展平。默认为 -1,表示最后一个维度。

torch.flatten() 函数会将输入张量的指定维度范围内的所有元素展平到一个一维张量中。展平后的张量保持与原始张量相同的数据顺序。例如,如果输入张量是一个 3x4x5 的三维张量,然后你使用 torch.flatten() 函数将它展平,那么结果将是一个包含 60 个元素的一维张量,其中包含原始张量中所有的元素。

以下是一个示例:

python 复制代码
import torch

# 创建一个3x4x5的张量
input_tensor = torch.randn(3, 4, 5)

# 使用torch.flatten()将其展平为一维张量
output_tensor = torch.flatten(input_tensor)

print(output_tensor.size())  # 输出 torch.Size([60])

在此示例中,input_tensor 是一个形状为 (3, 4, 5) 的三维张量,使用 torch.flatten() 函数将其展平为一个一维张量,并打印出了结果张量的大小。

示例:

python 复制代码
import torch

# 创建一个2×3x5x5的张量
input_tensor = torch.randn(2, 3, 5, 5)
print(f"原张量的尺寸为:{input_tensor.size()}") # torch.Size([2, 3, 5, 5])

# 使用torch.flatten()从第一个维度开始展平,从第二个维度结束展平
output_tensor = torch.flatten(input_tensor, start_dim=1, end_dim=2)
print(f"经过展平后的张量的尺寸为:{output_tensor.size()}")  # torch.Size([2, 15, 5])
相关推荐
Henry-SAP6 分钟前
SAP PP核心引擎计划策略业务解析
人工智能·云原生·sap·erp
DP DPharness24 分钟前
拆开 dsh-knowledge 的检索链路,看 RRF 融合与锚点续读
人工智能·dpharness
海上小飞龙26 分钟前
【KMP算法-下篇】同一道题:Java 库函数 2600 微秒,手写 KMP 16 微秒
java·开发语言·算法
I'm a winner36 分钟前
《AI 赋能嵌入式开发:从 0 到全栈工程师》模块1|第4课时
人工智能
johnsong41 分钟前
当决策成为免费商品:思维成本崩溃背后的治理真空
大数据·人工智能
青少儿编程课堂43 分钟前
后缀自动机解析:本质不同子串与最长重复片段统计
c++·python·算法·bfs·信息学竞赛
朝朝辞暮i44 分钟前
C++ 第 12 课:局部变量、作用域、变量生命周期
开发语言·c++·算法
ZhangJun9544 分钟前
在 32GB 内存电脑上本地搭建 Qwen3.6-35B-A3B 大模型踩坑实录
运维·人工智能·阿里云·ai·软件构建
南京码讯光电技术有限公司1 小时前
How to Design an Antenna System for an Industrial WiFi Module
数据库·人工智能
秋名山码民1 小时前
AI 用过一次,下一次怎样做得更好?——拙见 AI OS 的自适应、自迭代与受控演化
人工智能