PyTorch中如何通过索引张量为目标张量赋值(含排除指定索引场景)
PyTorch张量高效批量赋值方案
给定张量
- 零值张量A,形状为
(batch_size, vocab_size),示例:(16, 10000) - 索引张量B,形状为
(batch_size, seq_len),示例:(16, 20) - 值张量C,形状为
(batch_size, seq_len),示例:(16, 20)
需求
- 把A中对应B索引位置的值替换成C的值,实现类似
A[B] = C的效果 - 同样是替换,但要排除指定索引(比如所有行里排除索引3、5),过滤后没法用等维度张量表示,要实现类似
A[B[valid_indices]] = C[valid_indices]的操作
你当前的低效实现
你用嵌套循环来做,但两层循环耗时太长,代码如下:
for i,row in enumerate(probs): valid_indices = torch.tensor([idx[0] for idx in enumerate(encoder_input_ids[i]) if idx[1] not in [vocab['<pad>'],vocab['<unk>'], vocab['</s>']]]) valid_ids = torch.tensor([idx[0] for idx in enumerate(encoder_input_ids[i]) if idx[1] not in [vocab['<pad>'],vocab['<unk>'], vocab['</s>']]]) # print(valid_ids) # value = probs_c[i][valid_indices] # probs[i][tmp] = value #probs_c[i]
高效解决方案
需求1:直接批量赋值
用PyTorch的高级索引就能搞定,完全不用循环,速度快很多:
import torch # 先初始化示例张量 batch_size = 16 vocab_size = 10000 seq_len = 20 A = torch.zeros(batch_size, vocab_size) B = torch.randint(0, vocab_size, (batch_size, seq_len)) # 生成合法的随机索引 C = torch.rand(batch_size, seq_len) # 核心操作:生成每个batch对应的行索引,和B的列索引配对 batch_indices = torch.arange(batch_size).unsqueeze(1).repeat(1, seq_len) A[batch_indices, B] = C
说明:batch_indices会生成形状和B一样的张量,每个位置对应当前的batch行号,和B里的列索引组合成二维坐标,直接给A赋值,全程向量化运算,比循环快几个数量级。
需求2:过滤指定索引后赋值
先做掩码过滤掉要排除的索引,再提取有效部分批量赋值:
# 定义要排除的索引集合 exclude_indices = {3, 5} # 生成掩码:B中不在排除集合里的位置标记为True mask = ~torch.isin(B, torch.tensor(list(exclude_indices))) # 提取有效索引和对应的值 valid_batch_indices = batch_indices[mask] valid_B = B[mask] valid_C = C[mask] # 执行赋值 A[valid_batch_indices, valid_B] = valid_C
如果是要排除特定token(比如你代码里的<pad>、<unk>、</s>),直接用B和这些token的id做判断就行:
# 获取要排除的token对应的id exclude_token_ids = torch.tensor([vocab['<pad>'], vocab['<unk>'], vocab['</s>']]) # 生成掩码:排除掉这些token的位置 mask = ~torch.isin(B, exclude_token_ids) # 后续操作和上面一样 valid_batch_indices = batch_indices[mask] valid_B = B[mask] valid_C = C[mask] A[valid_batch_indices, valid_B] = valid_C
说明:torch.isin能一次性对整个B张量做判断,生成掩码后直接提取有效部分,最后用高级索引完成赋值,全程没有循环,效率拉满。
内容的提问来源于stack exchange,提问作者jupyter
相关产品推荐
相关产品推荐

