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

PyTorch Geometric中to_dense_batch的逆操作是什么?如何转回mini-batch?

将Dense Batch转回Mini-Batch的实现方案

PyTorch Geometric并没有内置名为from_dense_batch的函数,但我们可以利用dense_batch张量和对应的mask掩码手动实现转换,核心是通过掩码筛选有效节点数据,再恢复为PyG标准的mini-batch格式。

核心实现代码

假设你有以下输入:

  • dense_batch: 形状为 [batch_size, max_nodes, feature_dim] 的稠密节点特征张量
  • mask: 形状为 [batch_size, max_nodes] 的布尔掩码,True 对应有效节点位置
import torch
from torch_geometric.data import Batch, Data

def from_dense_batch(dense_batch, mask):
    # 提取所有有效节点的特征
    valid_node_feats = dense_batch[mask]
    
    # 统计每个样本的有效节点数
    node_counts = mask.sum(dim=1).tolist()
    
    # 生成batch向量:标记每个节点所属的样本索引
    batch_vec = torch.repeat_interleave(torch.arange(len(node_counts)), torch.tensor(node_counts))
    
    # 拆分并构造单个Data对象,再合并为mini-batch
    data_list = []
    idx = 0
    for count in node_counts:
        data = Data(x=valid_node_feats[idx:idx+count])
        data_list.append(data)
        idx += count
    
    mini_batch = Batch.from_data_list(data_list)
    return mini_batch, batch_vec

补充说明

  • 函数返回两个结果:PyG标准的Batch对象,以及用于标记节点归属的batch_vec张量
  • 如果你的dense batch包含邻接矩阵等其他数据,可以扩展函数逻辑:比如从稠密邻接矩阵中提取mask对应的有效子矩阵,再转换为COO格式的边索引

内容的提问来源于stack exchange,提问作者hellohawaii

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 13:20:17