如何用PyTorch内置函数移除批量均值计算中的双层for循环
问题描述
现有以下PyTorch代码,目标是根据音高标记构建均值张量:
B = spec_x.size(0) H = spec_x.size(1) T = spec_x.size(2) # Initialize x tensor with zeros z = torch.zeros(B, 256, H).to(pitch.device) # Iterate over each batch element for b in range(B): # Iterate over each pitch index for i in range(256): # Mask spec_x where pitch equals i masked_spec_x = spec_x[b].masked_select(pitch[b] == i) # Compute mean along the time dimension mean_spec_x = torch.mean(masked_spec_x, dim=0) # Assign the mean to the corresponding position in x z[b, i] = mean_spec_x
其中spec_x为形状[B, H, T]的频谱张量,pitch为形状[B, T]的音高标记张量(元素范围0-255),最终要得到形状[B, 256, H]的张量z,使得z[b][i]等于spec_x[b]中所有对应pitch为i的元素的平均值。
当前代码可实现需求,但双层for循环导致运行速度极慢,需用PyTorch内置函数移除循环优化性能。
优化方案
利用PyTorch的scatter_add实现向量化分组求和与计数,避免Python循环开销,具体代码如下:
import torch # 获取张量形状 B, H, T = spec_x.shape # 初始化求和张量与计数张量 sum_spec = torch.zeros(B, 256, H, device=spec_x.device) counts = torch.zeros(B, 256, 1, device=spec_x.device) # 将pitch扩展为[B, 1, T],适配scatter操作维度 pitch_expanded = pitch.unsqueeze(1) # 按pitch索引分组求和:将spec_x转置为[B, T, H]后,在维度1上scatter累加 sum_spec = sum_spec.scatter_add(1, pitch_expanded.expand(-1, -1, H), spec_x.transpose(1, 2)) # 统计每个pitch对应的元素数量 counts = counts.scatter_add(1, pitch_expanded, torch.ones_like(pitch_expanded, device=spec_x.device)) # 计算均值,处理无对应元素的情况(避免除以0,此处默认设为0,可按需调整) z = sum_spec / counts.clamp_min(1) # 若需保留无对应元素时的nan结果,可替换为以下代码: # z = torch.where(counts > 0, sum_spec / counts, torch.tensor(float('nan'), device=spec_x.device))
优化说明
- 向量化操作:完全利用PyTorch的张量运算,避免Python循环的性能损耗,GPU环境下加速效果更明显。
- 分组逻辑:通过
scatter_add将同一音高的频谱元素累加,同时统计每组元素数量,最后通过除法得到均值。 - 边界处理:针对无对应音高元素的情况,提供两种处理方式(设为0或保留nan),可根据实际需求选择。
内容的提问来源于stack exchange,提问作者jmasterx
相关产品推荐
相关产品推荐

