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

如何将(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接口,默认的展平逻辑会打乱块间和块内的元素排布,无法得到正确结果。
正确的变换只需要两步,全程无数据拷贝,效率最高:

  1. 调整张量轴顺序:将原始轴顺序(i, j, k, l)重排为(i, k, j, l),也就是把块内行索引对应的k轴移到块行索引i之后,把块内列索引对应的l轴移到块列索引j之后,调整后张量形状为(n, d, n, d)。
  2. 维度展平:将调整轴顺序后的张量前两个维度展平为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 16:45:39