如何在保持子矩阵相对位置的前提下合并高维Tensor的子矩阵?
张量帧网格重构与高维适配方案
问题背景
给定shape为[z, d, d]的Tensor x(代表视频帧序列),令pz = √z(假设z为完全平方数),需要将x转换为pz×pz的图像网格,最终得到shape为[1,1,pz*d,pz*d]的张量,且元素相对位置与原帧完全一致。
输入示例(shape
[4,2,2]):x = torch.tensor([[[ 0, 1], [ 2, 3]], [[ 4, 5], [ 6, 7]], [[ 8, 9], [10, 11]], [[12, 13], [14, 15]]])期望输出(shape
[1,1,4,4]):tensor([[[[ 0, 1, 4, 5], [ 2, 3, 6, 7], [ 8, 9, 12, 13], [10, 11, 14, 15]]]])直接使用
x.view(1,1,4,4)会打乱帧内像素结构,不符合需求。
同时需要适配更高维度的Tensor(如[b, c, z, d, d]),且必须避免嵌套循环以保证运算效率。
高效实现方案(无循环)
1. 基础3D Tensor [z,d,d] 转网格
通过维度重排与张量拼接实现,完全替代循环逻辑:
import torch z = 4 d = 2 x = torch.arange(z*d*d).view(z, d, d) pz = int(z**0.5) # 步骤1:重排维度为 [pz, pz, d, d] grid_x = x.view(pz, pz, d, d) # 步骤2:按行拼接子图像 → [pz, d, pz*d] row_merged = torch.cat(grid_x.unbind(dim=1), dim=2) # 步骤3:按列拼接行 → [pz*d, pz*d],再扩展维度到目标shape result = torch.cat(row_merged.unbind(dim=0), dim=0).unsqueeze(0).unsqueeze(0) print("转换结果:") print(result) print(f"结果shape: {result.shape}")
2. 高维Tensor [b,c,z,d,d] 转网格
针对批量(b)和通道(c)维度,只需在操作时保留前两维即可实现适配:
b = 2 c = 3 z = 4 d = 2 x = torch.arange(b*c*z*d*d).view(b, c, z, d, d) pz = int(z**0.5) # 步骤1:重排维度为 [b,c,pz,pz,d,d] grid_x = x.view(b, c, pz, pz, d, d) # 步骤2:按行拼接子图像 → [b,c,pz,d,pz*d] row_merged = torch.cat(grid_x.unbind(dim=3), dim=4) # 步骤3:按列拼接行 → [b,c,pz*d,pz*d] result = torch.cat(row_merged.unbind(dim=2), dim=2) print("高维转换结果shape:", result.shape) # 输出: torch.Size([2, 3, 4, 4])
3. 网格张量转回原张量(反向操作)
3D场景反向转换
# 假设result是[1,1,pz*d,pz*d]的张量 flat_grid = result.squeeze(0).squeeze(0) # 拆分行为pz块,每个块shape [d, pz*d] row_blocks = torch.chunk(flat_grid, pz, dim=0) # 每个行块拆分为pz个子图像,拼接成[pz,pz,d,d] grid_rev = torch.stack([torch.chunk(row, pz, dim=1) for row in row_blocks], dim=0) # 转回原shape [z,d,d] original_x = grid_rev.view(z, d, d) print("反向转换结果:") print(original_x) print(f"原shape恢复: {original_x.shape}")
高维场景反向转换
# 假设result是[b,c,pz*d,pz*d]的张量 # 拆分行为pz块,每个块shape [b,c,d,pz*d] row_blocks = torch.chunk(result, pz, dim=2) # 每个行块拆分为pz个子图像,拼接成[b,c,pz,pz,d,d] grid_rev = torch.stack([torch.chunk(row, pz, dim=3) for row in row_blocks], dim=2) # 转回原shape [b,c,z,d,d] original_x = grid_rev.view(b, c, z, d, d) print("高维反向转换结果shape:", original_x.shape) # 输出: torch.Size([2,3,4,2,2])
内容的提问来源于stack exchange,提问作者Skipper
相关产品推荐
相关产品推荐

