PyTorch中的permute, transpose, view, reshape和flatten函数详解(已解决)

1.permute

permute函数用于重新排列张量的维度。它接受一个元组作为参数,表示新的维度顺序。例如,如果我们有一个形状为(2, 3)的二维张量,我们可以使用permute函数将其维度重新排列为(3, 2),如下所示:

复制代码
>>> import torch
>>> x = torch.randn(2,3)
>>> x
tensor([[-0.5945,  0.7441,  0.5515],
        [-1.3831,  0.4533, -0.6908]])
>>> y = x.permute(1,0)
>>> y
tensor([[-0.5945, -1.3831],
        [ 0.7441,  0.4533],
        [ 0.5515, -0.6908]])
>>>

首先创建了一个形状为(2, 3)的二维张量x。然后,我们使用permute函数将其维度重新排列为(3, 2),并将结果存储在变量y中。

2.transpose

transpose函数用于交换张量的两个维度。它接受两个整数作为参数,表示要交换的维度的索引。例如,如果我们有一个形状为(2, 3)的二维张量,我们可以使用transpose函数交换第0维和第1维,如下所示:

复制代码
>>> import torch
>>> x = torch.randn(2,3)
>>> x
tensor([[-0.5945,  0.7441,  0.5515],
        [-1.3831,  0.4533, -0.6908]])

>>> y=x.transpose(0,1)
>>> y
tensor([[-0.5945, -1.3831],
        [ 0.7441,  0.4533],
        [ 0.5515, -0.6908]])
>>>

在上面的例子中,创建了一个形状为(2, 3)的二维张量x。然后,我们使用transpose函数将第0维和第1维交换,并将结果存储在变量y中。

需要注意的是,transpose函数与permute函数不同,它只交换两个特定的维度,而permute函数可以重新排列所有维度。

3.view / reshape

view和reshape函数用于将张量重塑为不同的形状。它们接受一个或两个整数元组作为参数,表示新的形状。例如,如果我们有一个形状为(2, 3)的二维张量,我们可以使用view或reshape函数将其重塑为形状为(6,)的一维张量,如下所示:

复制代码
>>> import torch
>>> x = torch.randn(2,3)
>>> x
tensor([[-0.5945,  0.7441,  0.5515],
        [-1.3831,  0.4533, -0.6908]])

>>> y = x.view(-1)
>>> y
tensor([-0.5945,  0.7441,  0.5515, -1.3831,  0.4533, -0.6908])
# 或者
>>> y=x.reshape(-1)
>>> y
tensor([-0.5945,  0.7441,  0.5515, -1.3831,  0.4533, -0.6908])
>>>

在上面的例子中,创建了一个形状为(2, 3)的二维张量x。然后,我们使用view或reshape函数将x重塑为形状为(6,)的一维张量,并将结果存储在变量y中。

需要注意的是,view和reshape函数实际上不会改变张量中的数据,只是改变了数据的布局方式。因此,新的形状必须与原始形状兼容,否则会抛出错误。具体来说,新的形状的元素总数必须与原始形状的元素总数相同。

4. flatten

flatten函数用于将多维张量展平为一维张量。它接受一个整数作为参数,表示展平后的一维张量的最大长度。例如,如果我们有一个形状为(2, 3)的二维张量,我们可以使用flatten函数将其展平为一维张量,如下所示:

复制代码
>>> import torch
>>> x = torch.randn(2,3)

>>> y=x.flatten(1)
>>> y
tensor([[-0.5945,  0.7441,  0.5515],
        [-1.3831,  0.4533, -0.6908]])
>>> y=x.flatten(0)
>>> y
tensor([-0.5945,  0.7441,  0.5515, -1.3831,  0.4533, -0.6908])
>>> y=x.flatten(2)
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
IndexError: Dimension out of range (expected to be in range of [-2, 1], but got 2)

在上面的例子中,创建了一个形状为(2, 3)的二维张量x。然后,我们使用flatten函数将x展平为一维张量,并将结果存储在变量y中。需要注意的是,flatten函数的参数指定了展平后的一维张量的最大长度。在本例中,我们将最大长度设置为1,因此展平后的张量将具有形状(6,)。如果展平后的长度超过了指定的最大长度,将会抛出错误。

总结:在PyTorch中,permute、transpose、view、reshape和flatten函数都是用于改变张量形状和维度的工具。它们具有不同的用途和特点,可以根据具体需求选择合适的函数来操作张量。

相关推荐
Python大数据分析@1 分钟前
使用大模型MCP采集数据,爬虫已经无门槛
python·网络爬虫
Linguwen7 分钟前
外贸GEO01|GEO是什么?生成式引擎优化,AI时代的新流量密码
人工智能
物质波波波13 分钟前
WS-RPE:面向边缘物理AI实时特征值计算的硬件工作窃取调度器与冗余PE激活架构
人工智能·fpga开发·架构·系统架构·硬件架构
quantdash_cc14 分钟前
告别自建 Requests/BS4 网页爬虫:基于 QuantDash 搭建零维保的高性能量化行情流水线
开发语言·爬虫·python·pandas·量化·quantdash
起个名字好难啊这也被占用了17 分钟前
LangGraphjs可中断可恢复的AI工作流
人工智能·ai编程
@Mr_LiuYang20 分钟前
状态栏动态上下文信息追加到Agent --《深入理解 AI Agent :设计原理与工程实践》实验2-8
人工智能·大模型·动态上下文·状态栏信息追加
65岁退休Coder30 分钟前
LangChain v1.3.4 笔记 - 07 补充:链式调用 LCEL
后端·python·langchain
卷无止境41 分钟前
FastAPI 部署在 Nginx 后面到底该怎么配
后端·python
碳基猿1 小时前
新媒体运营的终局:从“内容创作”走向“运营系统竞争”
人工智能·新媒体运营·产品运营·新媒体矩阵·多账号管理·矩阵分发·矩阵运营方法论
半亩码田1 小时前
C#转Python第3.1篇:Python 的 class 没有访问修饰符?面向对象的另一条路
开发语言·python·c#