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)
还没有任何评论哟~
