如何无循环实现带梯度保留的索引器张量分批处理与重组?
问题
给定索引器张量(如tensor([0,1,1,0]))和批量输入张量(形状(4,64),即tensor([[foo],[bar],[tee],[aaa]])),需要完成以下操作:
- 按索引器的数值分组输入:索引0对应
[foo, aaa],索引1对应[bar, tee] - 每组输入送入对应
Sequential网络得到输出 - 根据原始索引器将输出重组为原顺序(
[output_foo, output_bar, output_tee, output_aaa]),且必须保留梯度
原循环实现中,通过outputs.data[i].copy_()直接修改张量底层数据,导致梯度流被切断。希望用无循环的张量/矩阵操作实现该逻辑。
原示例代码:
B = 5 # 批量大小 F = 64 # 特征维度 outputs = torch.zeros((B), requires_grad=True) indexer = torch.argmax(masks, dim=1) # (B,) 整数类型 for idx in indexer.unique(): mask = (indexer == idx) # (B,) 布尔类型 curr_inputs = torch.masked_select(inputs, mask.unsqueeze(-1).repeat(1,F)).view(-1,F) # (b, F),仅保留idx对应的条目 # 根据原始索引为输入应用对应的Sequential网络 curr_outputs = apply_to_sequential_by_idx(curr_inputs, idx) # 将输出合并回原形状 (B,) j=0 for i, index in enumerate(indexer): if idx == index: outputs.data[i].copy_(curr_outputs.data[j]) j+=1 return outputs # 此处梯度丢失(无grad_fn)
无循环+保留梯度的实现方案
核心思路
避免直接修改张量的data属性(会切断梯度流),改用原生张量索引操作完成分组处理与结果重组,全程保留计算图的可微分性。
基础实现(仅遍历网络,无批量循环)
假设已预定义对应不同索引的网络字典nets = {0: Sequential(...), 1: Sequential(...)}:
import torch import torch.nn as nn # 模拟定义不同索引对应的处理网络 nets = { 0: nn.Sequential(nn.Linear(64, 1), nn.ReLU()), 1: nn.Sequential(nn.Linear(64, 1), nn.Sigmoid()) } def process_grouped_inputs(inputs, indexer): B, F = inputs.shape # 初始化输出张量,自动跟踪梯度 outputs = torch.zeros((B, 1), device=inputs.device, dtype=inputs.dtype) # 遍历唯一索引(仅遍历网络类型,而非批量元素,不影响梯度流) for idx in torch.unique(indexer): mask = (indexer == idx) # 直接通过索引赋值,PyTorch会自动记录计算图 outputs[mask] = nets[idx](inputs[mask]) return outputs.squeeze() # 若输出为(B,1)可压缩为(B,)
为什么能保留梯度?
- 使用
outputs[mask] = ...的原生张量赋值操作,而非修改data属性,PyTorch会将该操作纳入计算图 - 所有分组、处理、重组逻辑均为可微分操作,梯度可正常反向传播至输入和网络参数
进阶无循环实现(针对多索引场景)
如果索引类型较多,可通过torch.scatter和torch.gather进一步简化,无需遍历索引:
def process_grouped_inputs_scatter(inputs, indexer): B, F = inputs.shape num_indices = len(nets) # 对所有输入应用每个网络,得到形状(num_indices, B, 1)的全量输出 all_outputs = torch.stack([nets[idx](inputs) for idx in range(num_indices)], dim=0) # 利用scatter标记对应索引的输出位置,再gather提取最终结果 index_expand = indexer.unsqueeze(0).unsqueeze(-1) # 形状(1, B, 1) outputs = torch.zeros_like(all_outputs[0]) outputs.scatter_(0, index_expand.repeat(num_indices, 1, 1), all_outputs) outputs = outputs.gather(0, index_expand).squeeze(0) return outputs.squeeze()
内容的提问来源于stack exchange,提问作者Daniel Barahona
相关产品推荐
相关产品推荐

