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])
相关推荐
数字智核12 小时前
2026昆山工厂采购空压机怎么选?哪家公司能做选型和安装
人工智能
AbrahamCS12 小时前
告别无脑召回与死规则:基于国家标准(GB/T 48000.3)与大模型自主编排的 App 智能运营实战
大数据·人工智能·智能体·ontology
山西正方元12 小时前
西安商家公私域联动落地:品牌私域架构拆解与本地化适配
大数据·人工智能·#西安本地运营
腾视科技-AI12 小时前
腾视科技大模型一体机解决方案:低成本私有化落地,重塑行业智能应用新格局
大数据·人工智能·科技·ai·ai大模型·腾视科技·ai算力盒
Thomas.Sir12 小时前
第26课:TensorFlow|循环神经网络RNN原理【时序数据处理、序列依赖关系讲解】
人工智能·rnn·tensorflow
Wang's Blog13 小时前
Vibe Coding一人即团队系列54:云服务器 Node.js 与 MySQL 9 环境搭建及配置指南
服务器·人工智能·mysql·node.js
测试开发Kevin13 小时前
DeepEval + Eval‑Harness 完整讲解(结合 Playwright UI 自动化例子)
人工智能·ai·langchain
张欣-男13 小时前
5分钟理解线性代数v2
人工智能·线性代数·机器学习
pnoker13 小时前
从工业软件到 AI 智能体:工业 AIoT 技术路线的系统梳理
java·人工智能·物联网·microsoft·开源·工业互联网
全栈技术负责人13 小时前
DeepSeek Harness 业务工具权限插件 dsh-tool-permission设计思路
网络·算法·ai·ai编程