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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 21:04:51