如何将(n,n,d,d)形状张量高效Reshape为(n*d,n*d)分块矩阵
形状为(n,n,d,d)的张量转对应分块矩阵的高效实现
需求说明
给定形状为(n, n, d, d)的张量A,需要构造形状为(n*d, n*d)的分块矩阵B,满足索引映射关系:A[i,j,k,l] = B[i*d + k, j*d + l]。其中A的前两个维度对应分块矩阵的块行、块列索引,后两个维度对应每个d×d大小的块内部的行、列索引。
核心实现思路
注意:不要直接对原始张量调用reshape接口,默认的展平逻辑会打乱块间和块内的元素排布,无法得到正确结果。
正确的变换只需要两步,全程无数据拷贝,效率最高:
- 调整张量轴顺序:将原始轴顺序
(i, j, k, l)重排为(i, k, j, l),也就是把块内行索引对应的k轴移到块行索引i之后,把块内列索引对应的l轴移到块列索引j之后,调整后张量形状为(n, d, n, d)。 - 维度展平:将调整轴顺序后的张量前两个维度展平为
n*d(对应分块矩阵的行维度),后两个维度展平为n*d(对应分块矩阵的列维度),即可得到符合映射要求的目标矩阵B。
常见深度学习/数值计算框架代码示例
PyTorch 实现
使用permute接口调整轴顺序后调用reshape,整个操作是张量视图级操作,不产生数据拷贝:
import torch n, d = 3, 2 A = torch.randn(n, n, d, d) # 轴重排 + 维度展平 B = A.permute(0, 2, 1, 3).reshape(n * d, n * d) # 索引正确性校验 for i in range(n): for j in range(n): for k in range(d): for l in range(d): assert torch.isclose(A[i, j, k, l], B[i*d + k, j*d + l])
NumPy 实现
使用transpose接口调整轴顺序后调用reshape,同样为无拷贝的视图操作:
import numpy as np n, d = 3, 2 A = np.random.randn(n, n, d, d) B = A.transpose(0, 2, 1, 3).reshape(n * d, n * d) # 索引正确性校验 for i in range(n): for j in range(n): for k in range(d): for l in range(d): assert np.isclose(A[i, j, k, l], B[i*d + k, j*d + l])
TensorFlow 实现
使用tf.transpose调整轴顺序后调用tf.reshape展平维度:
import tensorflow as tf n, d = 3, 2 A = tf.random.normal((n, n, d, d)) B = tf.reshape(tf.transpose(A, perm=[0, 2, 1, 3]), (n*d, n*d))
效率说明
上述实现的轴调整操作仅修改张量的步长、维度等元信息,不涉及实际数据的复制搬运,操作耗时和张量本身的元素规模无关,为O(1)复杂度,是该转换需求的最优实现方案,可直接用于大规模张量的转换场景。
内容的提问来源于stack exchange,提问作者Neurobro
相关产品推荐
相关产品推荐

