求无需卷积的PyTorch张量手动维度逆转换快速方法(支持批量)
PyTorch张量维度逆转换实现(无卷积)
需要一种无需卷积的快速手动方法,将维度为(Batch, D, M, M)的PyTorch张量转换为(Batch, D//4, M×2, M×2),同时支持批量处理。已知反向转换((Batch, Depth, Size, Size) → (Batch, 4×Depth, Size//2, Size//2))的实现代码,以下是逆操作的完整解决方案。
已知反向转换代码
# Batch, Depth, size, size -> Batch, 4 x Depth, size//2, size//2 x = x.unfold(2, 2, 2).unfold(3, 2, 2).reshape(batch, depth, size // 2, size // 2, 4) x = x.permute(0, 2, 3, 1, 4).reshape(batch, size // 2, size // 2, depth * 4).permute(0, 3, 1, 2)
输入输出示例
输入张量(维度 (12, 2, 2))
[[[ 0, 1, 4, 5, 16, 17, 20, 21, 32, 33, 36, 37], [ 2, 3, 6, 7, 18, 19, 22, 23, 34, 35, 38, 39]], [[ 8, 9, 12, 13, 24, 25, 28, 29, 40, 41, 44, 45], [10, 11, 14, 15, 26, 27, 30, 31, 42, 43, 46, 47]]]
期望输出张量(维度 (3, 4, 4))
[[[ 0, 1, 2, 3], [ 4, 5, 6, 7], [ 8, 9, 10, 11], [12, 13, 14, 15]], [[16, 17, 18, 19], [20, 21, 22, 23], [24, 25, 26, 27], [28, 29, 30, 31]], [[32, 33, 34, 35], [36, 37, 38, 39], [40, 41, 42, 43], [44, 45, 46, 47]]]
逆转换实现与测试
下面是补全后的测试代码,逆转换代码部分实现了从(Batch, 4×Depth, Size//2, Size//2)转回(Batch, Depth, Size, Size)的逻辑:
import torch batch = 2 input = torch.arange(3*4*4*batch).reshape(batch,3, 4, 4) batch, depth, size, _ = input.shape # 正向转换(原反向操作) x = input.unfold(2, 2, 2).unfold(3, 2, 2).reshape(batch, depth, size // 2, size // 2, 4) x = x.permute(0, 2, 3, 1, 4).reshape(batch, size // 2, size // 2, depth * 4).permute(0, 3, 1, 2) # 逆转换代码 batch, new_depth, new_size, _ = x.shape x = x.permute(0, 2, 3, 1).reshape(batch, new_size, new_size, new_depth//4, 4) x = x.permute(0, 3, 1, 2, 4).reshape(batch, new_depth//4, new_size, new_size, 2, 2) x = x.permute(0,1,2,4,3,5).reshape(batch, new_depth//4, new_size*2, new_size*2) # 验证是否与原输入一致 print((input==x).all()) # 输出应为True
逆转换逻辑说明
- 调整维度顺序并拆分深度维度,还原正向转换中的分组结构,得到
(Batch, new_size, new_size, new_depth//4, 4); - 将深度维度移回第二位,重新组织为
(Batch, new_depth//4, new_size, new_size, 4); - 将最后一维的4个元素拆分为
(2,2)的空间块,通过维度重排和reshape恢复为原始的(Batch, Depth, Size, Size)维度。
内容的提问来源于stack exchange,提问作者MasterEpilif
相关产品推荐
相关产品推荐

