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

如何并行化加速张量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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 08:52:44