PyTorch中如何依据索引匹配从张量列表提取节点特征?
问题:基于多索引从张量列表提取特征填充目标张量
数据背景
现有约108000个节点的4维特征数据集,特征存储在tmp列表中,包含4个shape为(107940,4)的PyTorch张量:
import torch device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu') tmp = [] for _ in range(4): tmp.append(torch.rand((107940, 4), dtype=torch.float).to(device))
已知图batch中的edge_index(行0为目标节点、行1为源节点,为batch内节点索引)和edge_index_class(每条边对应的类别),通过batch.n_id[batch.edge_index]可获取边对应的原始图节点ID,示例输出如下:
# 原始节点ID(对应batch.edge_index的映射) print(batch.n_id[batch.edge_index]) # tensor([[10231, 3059, 32075, ..., 10087, 10158, 10158], # [ 1624, 1624, 6466, ..., 10087, 10158, 10158]], device='cuda:0') # 每条边对应的类别 print(batch.edge_index_class) # tensor([3., 3., 2., ..., 2., 2., 2.], device='cuda:0')
需求
生成shape为(107940,4)的tmp_filled张量,依据edge_index_class的值从tmp对应位置的张量中提取节点特征填充:
- 若某条边的
edge_index_class为3,则从tmp[3]中提取该边两个节点的特征,填入tmp_filled的对应原始节点索引位置。
错误尝试与问题
原代码逻辑错误,导致tmp_filled大部分为0,且取值与预期不符:
# 错误代码 tmp_filled = torch.zeros((107940, 4), dtype=torch.float, device=device) for k in range(len(batch.edge_index_class)): tmp_filled[batch.n_id[torch.unique(batch.edge_index)]] = tmp[int(batch.edge_index_class[k].item())][batch.n_id[torch.unique(batch.edge_index)]]
验证发现取值不符:
tmp_filled[1624] # tensor([0.3438, 0.5555, 0.6229, 0.7983], device='cuda:0') tmp[3][1624] # tensor([0.6895, 0.3241, 0.1909, 0.1635], device='cuda:0')
错误原因:
- 循环中每次对所有唯一节点批量赋值,后续循环的类别会覆盖前面所有节点的取值,最终仅保留最后一次循环的类别特征。
- 索引逻辑错误:未针对每条边的类别匹配对应节点,而是一次性处理所有节点,导致映射关系混乱。
修正方案
方案1:批量索引填充(高效)
通过张量堆叠与批量索引实现,适合大规模数据:
# 初始化目标张量 tmp_filled = torch.zeros((107940, 4), dtype=torch.float, device=device) # 获取所有边对应的原始节点ID,展平为一维数组 original_nodes = batch.n_id[batch.edge_index].flatten() # 每个节点对应的类别:每条边的两个节点共享同一类别,因此重复类别列表两次 all_classes = batch.edge_index_class.repeat(2).long() # 将tmp列表堆叠为三维张量,shape=(4, 107940, 4) tmp_stacked = torch.stack(tmp) # 批量提取对应类别、对应节点的特征 selected_features = tmp_stacked[all_classes, original_nodes] # 填充到目标张量 tmp_filled[original_nodes] = selected_features
方案2:按类别循环填充(直观)
按类别遍历,针对每个类别提取对应边的节点并填充:
# 初始化目标张量 tmp_filled = torch.zeros((107940, 4), dtype=torch.float, device=device) # 遍历每个类别 for cls in range(4): # 筛选出当前类别的边的掩码 cls_edge_mask = (batch.edge_index_class == cls) # 获取当前类别边对应的batch内节点,映射为原始节点ID并展平 cls_original_nodes = batch.n_id[batch.edge_index[:, cls_edge_mask]].flatten() # 从tmp对应类别张量中提取特征,填充到目标张量 tmp_filled[cls_original_nodes] = tmp[cls][cls_original_nodes]
验证修正结果:
# 此时tmp_filled[1624]应与tmp[3][1624]完全一致 print(tmp_filled[1624] == tmp[3][1624]) # tensor([True, True, True, True], device='cuda:0')
内容的提问来源于stack exchange,提问作者Ssong
相关产品推荐
相关产品推荐

