You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在保持子矩阵相对位置的前提下合并高维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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.30 16:39:21