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

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. 循环中每次对所有唯一节点批量赋值,后续循环的类别会覆盖前面所有节点的取值,最终仅保留最后一次循环的类别特征。
  2. 索引逻辑错误:未针对每条边的类别匹配对应节点,而是一次性处理所有节点,导致映射关系混乱。

修正方案

方案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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 09:14:56