StellarGraph无监督GraphSage示例运行报错求助
问题排查与解决方案
错误原因分析
这个错误本质是GraphSAGE无监督模型的输入结构与生成器输出不匹配:
- 无监督GraphSAGE训练完成后,用于提取embedding的模型(
embedding_model)需要接收多个输入张量(对应中心节点特征、各层邻居采样的特征),但你的node_gen生成器只输出了单个张量,导致模型输入不匹配。 - 大概率是StellarGraph与TensorFlow版本不兼容,或者生成器配置、模型构建的步骤出现偏差。
具体解决方案
1. 安装兼容版本的StellarGraph和TensorFlow
版本不兼容是这类输入不匹配问题的常见诱因,建议锁定经过验证的版本组合:
pip install stellargraph==1.2.1 tensorflow==2.4.1
安装完成后重启运行环境,确保版本生效。
2. 修正生成器与模型构建的关键步骤
确保整个流程的输入输出结构完全匹配:
- 创建UnsupervisedSampler:
from stellargraph.data import UnsupervisedSampler # 采样参数需与后续GraphSAGE的层数对应 sampler = UnsupervisedSampler( G, nodes=G.nodes(), length=5, number_of_walks=2 ) - 创建GraphSAGENodeGenerator:
from stellargraph.mapper import GraphSAGENodeGenerator # num_samples的长度要和GraphSAGE的layer_sizes长度一致 generator = GraphSAGENodeGenerator( G, batch_size=50, num_samples=[10, 5] # 对应2层GraphSAGE,每层采样10、5个邻居 ) - 构建无监督GraphSAGE模型并提取embedding模型:
from stellargraph.layer import GraphSAGE, link_classification from tensorflow.keras.models import Model graphsage = GraphSAGE( layer_sizes=[128, 128], # 层数与num_samples长度一致 generator=generator, bias=True, dropout=0.5, ) x_inp, x_out = graphsage.in_out_tensors() # 构建训练用的链接分类模型 prediction = link_classification( output_dim=1, output_act="sigmoid", edge_embedding_method="ip" )(x_out) train_model = Model(inputs=x_inp, outputs=prediction) # 构建用于提取embedding的模型(关键:复用x_inp和x_out) embedding_model = Model(inputs=x_inp, outputs=x_out) - 创建节点embedding的生成器:
这里的node_gen = generator.flow(G.nodes())generator.flow(G.nodes())会自动生成模型需要的多输入张量结构(中心节点+各层邻居特征)。
3. 调整predict调用参数
避免因多线程导致的输入解析问题,先尝试单worker模式:
node_embeddings = embedding_model.predict(node_gen, workers=1, verbose=1)
运行正常后再逐步调高workers数值。
验证方法
运行以下代码检查生成器输出结构是否符合要求:
sample = next(iter(node_gen)) print(f"生成器输出的张量数量: {len(sample)}") print(f"每个张量的形状: {[t.shape for t in sample]}")
正常情况下,输出的张量数量应该等于len(num_samples) + 1(中心节点 + 每层邻居),比如num_samples=[10,5]时,应该有3个输入张量,与模型的输入层数量对应。
内容的提问来源于stack exchange,提问作者Shradhit Subudhi
相关产品推荐
相关产品推荐

