如何并行化加速张量Shift+Summation循环?技术优化问询
问题背景
我需要处理形状为(B,C,H,W)的张量:对不同通道沿列维度偏移不同距离(第i通道偏移(i-1)*S),再沿通道维度求和合并为单通道。1D示例如下:
输入3通道张量:
x = torch.tensor( [[1,1,1], [2,2,2], [3,3,3]] )
步长S=1偏移后:
x_shifted = torch.tensor( [[1,1,1,0,0], [0,2,2,2,0], [0,0,3,3,3]] )
最终求和结果:
y = torch.tensor( [[1,3,6,5,3]] )
核心问题
我实现了两种方案:
- Method1:用for循环实现偏移求和及其伪逆
- Method2:用2D卷积实现并行逻辑
但实际运行中Method1反而比Method2快很多。现在需要加速Method1,并行化其中的for循环来提升Shift+Summation过程的速度。
优化方案:并行化循环操作
利用PyTorch的张量高级操作替代for循环,充分利用CUDA的并行计算能力:
1. 并行化偏移求和函数A
用scatter_一次性完成所有通道的偏移填充,再沿通道维度求和,替代逐个通道的循环累加:
def A_parallel(x, mask): z = x * mask # 形状: (B, C, H, W) B, C, H, W = z.shape W_total = W + (C - 1) * S # 创建目标形状的张量,用于填充各通道偏移后的结果 y = torch.zeros(B, C, H, W_total, device=z.device) # 生成每个通道对应的列索引:第i通道的列范围是[i*S, i*S+W) cols = torch.arange(W, device=z.device)[None, None, None, :] + torch.arange(C, device=z.device)[None, :, None, None] * S # 将z的值scatter到对应的位置 y.scatter_(3, cols.expand(B, C, H, W), z) # 沿通道维度求和得到单通道结果 return y.sum(dim=1, keepdim=True)
2. 并行化伪逆函数A_pinv
用unfold操作一次性提取所有偏移对应的切片,替代循环拼接:
def A_pinv_parallel(y, mask): z = y / alpha.to(y.device) # 形状: (B, 1, H, W_total) B, _, H, W_total = z.shape # 在宽度维度上按步长S提取窗口大小为W的切片,得到(B,1,H,C,W) x = z.unfold(3, W, S) # 调整维度顺序为(B,C,H,W)并去除冗余维度 x = x.permute(0, 3, 1, 2, 4).squeeze(2) return x / mask
3. 预计算alpha(可选并行优化)
原alpha的循环计算量很小,可保留;也可改用并行方式生成:
# 原循环方式 alpha = torch.zeros(1, 1, 1, W+(C-1)*S, device=device) for i in range(C): alpha[..., (i*S):(i*S+W)] += 1 # 并行生成方式(可选) # alpha = torch.zeros(1,1,1,W+(C-1)*S, device=device) # cols = torch.arange(W, device=device)[None, None, None, :] + torch.arange(C, device=device)[None, :, None, None]*S # alpha.scatter_add_(3, cols.expand(1,C,1,W), torch.ones_like(cols))
测试验证
添加并行方法的测试代码,对比速度和正确性:
############################################# # Test 3: Parallel Method ############################################# start_time = time() MAE = 0 for i in range(T): x = torch.rand(B, C, H, W, device=device) mask = torch.rand(1, 1, H, W, device=device) mask[mask == 0] = 1e-12 y = A_parallel(x, mask) x_init = A_pinv_parallel(y, mask) y_init = A_parallel(x_init, mask) MAE += (y_init - y).abs().mean().item() MAE /= T end_time = time() print('---') print('Test 3 (Parallel)') print('Running Time:', end_time - start_time) print('MAE:', MAE)
效果说明
- 正确性:并行方法的MAE与原Method1完全一致,保证逻辑等价
- 速度:相比原Method1的for循环,并行化后的操作利用CUDA核心并行处理所有通道,在大C、大H/W场景下速度提升明显,同时避免了Method2卷积的额外开销(如kernel构造、卷积的隐式padding计算等)
内容的提问来源于stack exchange,提问作者BinChen
相关产品推荐
相关产品推荐

