PyTorch张量同时删除指定行与列的高效实现方案
高效删除PyTorch方阵中指定行和列的方法
我需要一种高效方法,同时删除PyTorch中形状为[l,l]的二维方阵张量的指定行和列。目前已通过以下方式实现目标,但速度较慢(列表索引并非张量视图操作,且索引分为两步执行):
# 待删除的索引 idx = 499 keep = [_ for _ in range(l)] keep.remove(idx) t = t[keep,:][:,keep] # 此时t的维度为[l-1,l-1]
能否推荐更快的实现方式?
测试方法1:基于列表的双重索引
import torch import time l = 8000 iterations = 500 idx = 599 t = torch.rand([l,l]) # 测试1 - 列表双重索引 keep = [_ for _ in range(l)] keep.remove(idx) total = 0 for i in range(iterations): start = time.time() t2 = t[keep,:][:,keep] torch.cuda.synchronize() elapsed = time.time() - start total += elapsed print("Took {:.1f}s for {} iterations, {}s/it".format(total,iterations,total/iterations))
耗时:500次迭代共35.0秒,每次迭代0.07009070539474488秒
测试方法2:基于张量的双重索引
(在某些情况下比上述方法略快,但二者均为非连续内存操作)
# 测试2 - 张量双重索引 keep = torch.tensor(keep) total = 0 for i in range(iterations): start = time.time() t2 = t[keep,:][:,keep] torch.cuda.synchronize() elapsed = time.time() - start total += elapsed print("Took {:.1f}s for {} iterations, {}s/it".format(total,iterations,total/iterations))
耗时:500次迭代共34.6秒,每次迭代0.06911029624938965秒
测试方法3:拼接法
# 测试3 - 基于拼接的实现 total = 0 for i in range(iterations): start = time.time() t2 = torch.cat([torch.cat([t[:idx,:idx],t[idx+1:,:idx]],dim = 0),torch.cat([t[:idx,idx+1:],t[idx+1:,idx+1:]],dim = 0)],dim = 1) torch.cuda.synchronize() elapsed = time.time() - start total += elapsed print("Took {:.1f}s for {} iterations, {}s/it".format(total,iterations,total/iterations))
耗时:500次迭代共31.1秒,每次迭代0.06218040370941162秒
测试方法4:移位并删除末尾行/列
(这是一种快速的传统索引操作,但需要克隆t。该方法与上述拼接法速度最快,但删除多行/列时扩展性不佳)
# 测试4 - 移位删除末尾法 total = 0 for i in range(iterations): t2 = torch.clone(t) start = time.time() t2[idx:-1,:] = t[idx+1:,:] t2[:,idx:-1] = t[:,idx+1:] t2 = t2[:-1,:-1] torch.cuda.synchronize() elapsed = time.time() - start total += elapsed print("Took {:.1f}s for {} iterations, {}s/it".format(total,iterations,total/iterations))
耗时:500次迭代共26.5秒,每次迭代0.052913659572601315秒
若排除克隆张量的时间:
耗时:500次迭代共11.7秒,每次迭代0.023439947128295897秒
测试方法5:展平为一维数组进行单次索引
total = 0 for i in range(iterations): start = time.time() o = torch.ones(t.shape,dtype = int) o[:,idx] = 0 o[idx,:] = 0 o = o.view(-1).nonzero().squeeze(1) t2 = t.view(-1)[o].view(l-1,l-1) elapsed = time.time() - start total += elapsed print("Took {:.1f}s for {} iterations, {}s/it".format(total,iterations,total/iterations))
耗时:500次迭代共87.1秒,每次迭代0.11186568689346313秒
测试方法6:二维布尔掩码索引
(底层实现与上述方法大致相同,最终也需要view操作重塑形状,耗时相近)
total = 0 for i in range(iterations): o = torch.ones(t.shape,dtype = bool) start = time.time() o[:,idx] = 0 o[idx,:] = 0 t2 = t[o].view(l-1,l-1) elapsed = time.time() - start total += elapsed print("Took {:.1f}s for {} iterations, {}s/it".format(total,iterations,total/iterations))
耗时:500次迭代共86.1秒,每次迭代0.17228954696655274秒
测试方法7:复制到空张量(参考@trialNerror的答案)
total = 0 for i in range(iterations): start = time.time() H,W = t.shape # 空初始化,仅分配内存 t2 = torch.empty(H-1, W-1) # 左上角块 t2[:idx, :idx] = t[:idx, :idx] # 右上角块 t2[:idx, idx:] = t[:idx, idx+1:] # 左下角块 t2[idx:, :idx] = t[idx+1:, :idx] # 右下角块 t2[idx:, idx:] = t[idx+1:, idx+1:] elapsed = time.time() - start total += elapsed
耗时:500次迭代共16.0秒,每次迭代0.031956102848052975秒
内容的提问来源于stack exchange,提问作者DerekG
相关产品推荐
相关产品推荐

