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

如何用PyTorch内置函数替代循环实现分块矩阵拼接?

用PyTorch内置函数替代分块矩阵拼接的嵌套循环

完全可以用PyTorch的维度变换内置函数消除这两层循环,而且效率比循环高得多——毕竟内置操作都是底层优化过的并行实现,没有Python循环的开销。

具体实现思路

原张量y的形状是(B//s2, D2//s1, s1, s2),我们需要把每个分块y[j,i,...]放到大矩阵的[i*s1:(i+1)*s1, j*s2:(j+1)*s2]位置。核心是通过维度重排和形状重塑完成分块的拼接:

  1. 用permute调整维度顺序:把原张量的维度从(J, I, s1, s2)(其中J=B//s2,I=D2//s1)重排为(I, s1, J, s2),这一步对应把分块的行索引i和分块内部的行s1放在一起,列索引j和分块内部的列s2放在一起。
  2. 用view把前两个维度合并为D2(I*s1=D2),后两个维度合并为B(J*s2=B),直接得到完整的大矩阵。

代码示例

import torch

# 自定义参数(确保所有除法结果都是整数)
B = 8
D2 = 6
s1 = 2
s2 = 4

# 生成测试用的4D张量y
y = torch.randn(B // s2, D2 // s1, s1, s2)

# 优化后的实现(无循环)
result_optimized = y.permute(1, 2, 0, 3).contiguous().view(D2, B)

# 原循环实现(用于验证正确性)
result_loop = torch.zeros(D2, B)
for i in range(D2 // s1):
    for j in range(B // s2):
        result_loop[i * s1 : (i+1)*s1, j * s2 : (j+1)*s2] = y[j, i, ...]

# 验证两种方法结果一致
print(torch.allclose(result_loop, result_optimized))  # 输出 True

关键说明

  • permute(1,2,0,3):直接调整张量的维度顺序,把原维度0和1交换,再把维度1和2交换,一步到位得到(I, s1, J, s2)的形状。
  • contiguous():因为permute会改变张量的内存布局,调用这个方法确保内存连续,之后才能安全使用view。
  • view(D2, B):把调整后的4D张量直接重塑为目标的2D矩阵,本质是把所有分块按顺序拼接成完整矩阵。

内容的提问来源于stack exchange,提问作者Goldenalcheese

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 15:02:30