您当前的位置:首页 > IT编程 > python
| C语言 | Java | VB | VC | python | Android | TensorFlow | C++ | oracle | 学术与代码 | cnn卷积神经网络 | gnn | 图像修复 | Keras | 数据集 | Neo4j | 自然语言处理 | 深度学习 | 医学CAD | 医学影像 | 超参数 | pointnet | pytorch | 异常检测 | Transformers | 情感分类 | 知识图谱 |

自学教程:Python torch.flatten()函数案例详解

51自学网 2021-10-30 22:17:07
  python
这篇教程Python torch.flatten()函数案例详解写得很实用,希望能帮到您。

先看函数参数:

torch.flatten(input, start_dim=0, end_dim=-1)

input: 一个 tensor,即要被“推平”的 tensor。

start_dim: “推平”的起始维度。

end_dim: “推平”的结束维度。

首先如果按照 start_dim 和 end_dim 的默认值,那么这个函数会把 input 推平成一个 shape 为 [n][n] 的tensor,其中 nn 即 input 中元素个数。

如果我们要自己设定起始维度和结束维度呢?

我们要先来看一下 tensor 中的 shape 是怎么样的:

t = torch.tensor([[[1, 2, 2, 1],                   [3, 4, 4, 3],                   [1, 2, 3, 4]],                  [[5, 6, 6, 5],                   [7, 8, 8, 7],                   [5, 6, 7, 8]]])print(t, t.shape) 运行结果: tensor([[[1, 2, 2, 1],         [3, 4, 4, 3],         [1, 2, 3, 4]],         [[5, 6, 6, 5],         [7, 8, 8, 7],         [5, 6, 7, 8]]])torch.Size([2, 3, 4])

我们可以看到,最外层的方括号内含两个元素,因此 shape 的第一个值是 2;类似地,第二层方括号里面含三个元素,shape 的第二个值就是 3;最内层方括号里含四个元素,shape 的第二个值就是 4。

示例代码:

x = torch.flatten(t, start_dim=1)print(x, x.shape) y = torch.flatten(t, start_dim=0, end_dim=1)print(y, y.shape)  运行结果: tensor([[1, 2, 2, 1, 3, 4, 4, 3, 1, 2, 3, 4],        [5, 6, 6, 5, 7, 8, 8, 7, 5, 6, 7, 8]]) torch.Size([2, 12]) tensor([[1, 2, 2, 1],        [3, 4, 4, 3],        [1, 2, 3, 4],        [5, 6, 6, 5],        [7, 8, 8, 7],        [5, 6, 7, 8]]) torch.Size([6, 4])

可以看到,当 start_dim = 11 而 end_dim = −1−1 时,它把第 11 个维度到最后一个维度全部推平合并了。而当 start_dim = 00 而 end_dim = 11 时,它把第 00 个维度到第 11 个维度全部推平合并了。pytorch中的 torch.nn.Flatten 类和 torch.Tensor.flatten 方法其实都是基于上面的 torch.flatten 函数实现的。

到此这篇关于Python torch.flatten()函数案例详解的文章就介绍到这了,更多相关Python torch.flatten()函数内容请搜索51zixue.net以前的文章或继续浏览下面的相关文章希望大家以后多多支持51zixue.net!


Python之基础函数案例详解
Python正则表达式中的量词符号与组问题小结
万事OK自学网:51自学网_软件自学网_CAD自学网自学excel、自学PS、自学CAD、自学C语言、自学css3实例,是一个通过网络自主学习工作技能的自学平台,网友喜欢的软件自学网站。