from torch_geometric.loader import NeighborLoader from torch_geometric.utils import to_networkx # Create batches with neighbor sampling train_loader = NeighborLoader( data, num_neighbors=[5, 10], batch_size=16, input_nodes=data.train_mask, ) # Print each subgraph for i, subgraph in enumerate(train_loader): print(f'Subgraph {i}: {subgraph}') # Plot each subgraph fig = plt.figure(figsize=(16,16)) for idx, (subdata, pos) in enumerate(zip(train_loader, [221, 222, 223, 224])): G = to_networkx(subdata, to_undirected=True) ax = fig.add_subplot(pos) ax.set_title(f'Subgraph {idx}') plt.axis('off') nx.draw_networkx(G, pos=nx.spring_layout(G, seed=0), with_labels=True, node_size=200, node_color=subdata.y, cmap="cool", font_size=10 ) plt.show()