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

PyTorch中如何通过索引张量为目标张量赋值(含排除指定索引场景)

PyTorch张量高效批量赋值方案

给定张量

  • 零值张量A,形状为(batch_size, vocab_size),示例:(16, 10000)
  • 索引张量B,形状为(batch_size, seq_len),示例:(16, 20)
  • 值张量C,形状为(batch_size, seq_len),示例:(16, 20)

需求

  1. 把A中对应B索引位置的值替换成C的值,实现类似A[B] = C的效果
  2. 同样是替换,但要排除指定索引(比如所有行里排除索引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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 11:22:02