如何用PyTorch内置函数替代循环实现分块矩阵拼接?
用PyTorch内置函数替代分块矩阵拼接的嵌套循环
完全可以用PyTorch的维度变换内置函数消除这两层循环,而且效率比循环高得多——毕竟内置操作都是底层优化过的并行实现,没有Python循环的开销。
具体实现思路
原张量y的形状是(B//s2, D2//s1, s1, s2),我们需要把每个分块y[j,i,...]放到大矩阵的[i*s1:(i+1)*s1, j*s2:(j+1)*s2]位置。核心是通过维度重排和形状重塑完成分块的拼接:
- 用
permute调整维度顺序:把原张量的维度从(J, I, s1, s2)(其中J=B//s2,I=D2//s1)重排为(I, s1, J, s2),这一步对应把分块的行索引i和分块内部的行s1放在一起,列索引j和分块内部的列s2放在一起。 - 用
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
相关产品推荐
相关产品推荐

