如何基于PyTorch稀疏张量实现scatter max操作以避免密集化带来的性能问题?
如何基于PyTorch稀疏张量实现scatter max操作以避免密集化带来的性能问题?
嘿,我来帮你搞定这个性能瓶颈!你现在头疼的点就是把大尺寸稀疏张量转密集的速度太慢对吧?其实完全不用走密集化这条路,我们可以直接基于稀疏张量的非零元素实现和原代码一模一样的逻辑,效率能提升一大截。
先帮你拆解下原代码的核心逻辑:你本质上是要对每个像素位置(y,x),取所有行(line)中该位置的最大值——这里要注意,原代码把稀疏张量转密集后,未被稀疏值覆盖的位置会被自动填充为0,所以最终的最大值其实是「该位置所有稀疏值的最大值」和「0」两者中的较大者;如果某个像素位置在所有行里都没有稀疏值,结果就是0。
那我们直接针对稀疏张量的非零元素动手就行,具体步骤和代码如下:
import torch from torch_scatter import scatter_max # 假设你已经有了稀疏张量value_tensor,以及img_size(即原代码中的img_size) # 提取稀疏张量的非零元素坐标和对应值,这俩都是小尺寸张量,处理极快 coords = value_tensor.indices() # 形状为(3, N),每一行对应[行索引, y坐标, x坐标] vals = value_tensor.values() # 把每个非零元素的(y,x)坐标转成扁平索引,和原代码的indices逻辑完全对齐 group_idx = coords[1] * img_size + coords[2] # 按像素扁平索引分组取最大值,dim_size确保覆盖所有像素位置 max_per_group, _ = scatter_max(vals, group_idx, dim=0, dim_size=img_size * img_size) # 还原原代码中密集化后0参与max比较的逻辑:无稀疏值的像素取0,有值的取稀疏值max和0的较大者 max_per_group = torch.max(max_per_group, torch.zeros_like(max_per_group)) # 最后reshape成图像形状,这个张量尺寸很小,完全不会有性能问题 img = max_per_group.reshape(img_size, img_size)
为什么这和原代码结果完全一致?
原代码转密集后,每个像素位置的元素包含「该行的稀疏值(如果有)」和「其他行的0」,scatter max取的是这些元素的最大值,等价于「该位置所有稀疏值的最大值」和「0」的较大者。我们的代码先提取所有稀疏值按像素分组取max,再和0取max,完美复刻了这个逻辑,而且全程没碰那个大尺寸的密集张量。
另外要提一句:最后得到的max_per_group是尺寸为img_size*img_size的密集张量,这个尺寸和原代码中num_lines*img_size*img_size的密集张量比起来小太多了,完全不会有性能瓶颈,真正解决了你原来的密集化慢的问题。
备注:内容来源于stack exchange,提问作者Cedric Martens
相关产品推荐
相关产品推荐

