如何基于条件为张量指定行赋值?PyTorch高效无循环实现问询
高效实现基于mask与right张量条件的left张量赋值
给定三个PyTorch张量:
left[N, 2]:待修改的目标张量right[N, 2]:用于判断赋值条件的张量mask[1, N]:布尔掩码,筛选需要处理的行
需求
通过mask过滤后,对left的指定行执行如下赋值规则:
- 若right对应行的两个元素全部属于
c1 = [1,3,5],或者全部属于c2 = [11,13],则left该行赋值为[0, 1] - 否则赋值为
[1, 0]
要求实现无for循环的高效版本,可直接在模型forward方法中运行。
示例张量
初始left张量
[[0.9, 0.8], [0.3, 0.0], [0.6, 0.9], [0.7, 0.0], [0.6, 0.8], [0.6, 0.2], [0.6, 0.2]]
mask布尔掩码
[ True, True, False, True, False, False, True]
right条件张量
[[ 1., 3.], [ 1., 5.], [ 7., 0.], [11., 13.], [17., 19.], [21., 1. ], [ 1., 13.]]
无效的尝试代码
你尝试的代码无法工作,因为它用标量判断逻辑处理批量张量,无法实现逐行的条件判断:
left[mask] = torch.tensor([0, 1]) if (right[mask][0] in c1 and right[mask][1] in c1) or ( right[mask][0] in c2 and right[mask][1] in c2) else torch.tensor([1,0])
高效解决方案
使用PyTorch的向量化操作实现批量条件判断与赋值,完全避免循环:
import torch # 将条件集合转为同dtype的张量,适配right的数值类型与设备 c1 = torch.tensor([1, 3, 5], dtype=right.dtype, device=right.device) c2 = torch.tensor([11, 13], dtype=right.dtype, device=right.device) # 筛选mask对应的right子张量 right_masked = right[mask] # 逐行判断:两个元素是否都在c1中 all_in_c1 = torch.isin(right_masked, c1).all(dim=1) # 逐行判断:两个元素是否都在c2中 all_in_c2 = torch.isin(right_masked, c2).all(dim=1) # 合并条件:满足任一集合的全包含要求 valid_condition = all_in_c1 | all_in_c2 # 生成对应赋值张量:满足条件则为[0,1],否则为[1,0] # 保持与left同dtype、同设备,避免类型/设备不匹配错误 assign_tensor = torch.where( valid_condition.unsqueeze(1), torch.tensor([0, 1], dtype=left.dtype, device=left.device), torch.tensor([1, 0], dtype=left.dtype, device=left.device) ) # 批量赋值到left的mask筛选行 left[mask] = assign_tensor
关键步骤说明
torch.isin:批量判断张量元素是否属于目标集合,返回同形状的布尔张量.all(dim=1):对每行的布尔结果取逻辑与,判断该行两个元素是否全部符合集合要求valid_condition.unsqueeze(1):将一维条件张量转为二维(形状从[M]变为[M,1]),与赋值的[2]维度匹配,确保torch.where能逐行生成赋值结果- 设备与类型对齐:所有临时张量的
dtype和device与原张量保持一致,避免运行时错误
内容的提问来源于stack exchange,提问作者topcat
相关产品推荐
相关产品推荐

