Advertisement

Pytorch: flatten() function used to compress tensor dimensionality.

阅读量:

前言

在处理数据时经常会遇到维度转换的需求。例如,在原始维度为BNWH的情况下(BNWH代表Batch, Height, Width, Channels),我们需要将其转换为两个输出:一个尺寸为512×(N×W×H),另一个尺寸为(N×W×H)×512,并分别对其进行矩阵乘法运算。
具体而言,
(NH W)\times 512 \times 512\times (NH W)=(NH W)\times (NH W)
然后再与原始输出结果BH W进行矩阵乘法运算。
经过上述操作后,
最终得到的结果仍保持相同的形状BH W
这种操作常用于非局部卷积块中以增强模型的表示能力。

代码

复制代码
    import torch
    import numpy as np
    
    input = torch.randn(2,3,4,4)
    # 将从第二个维度开始进行压缩
    # 可以根据自己需要选择从哪里开始压缩
    out = input.flatten(start_dim=1,end_dim=3) 
    out.shape

得到:

复制代码
    torch.Size([2, 48])  # 3*4*4=48

全部评论 (0)

还没有任何评论哟~