Stellargraph GraphSAGE示例中model.predict方法报错求助
问题解决:GraphSAGE predict时输入不匹配错误
问题描述
运行Stellargraph官方GraphSAGE节点分类示例时,执行到model.predict(all_mapper)报错:
ValueError: Layer "model_1" expects 3 input(s), but it received 1 input tensors. Inputs received: [<tf.Tensor 'IteratorGetNext:0' shape=(None, None, None) dtype=float32>]
模型summary显示存在3个InputLayer,分别对应(None,1,1433)、(None,10,1433)、(None,50,1433)的输入形状。
原因分析
模型需要3个输入张量,对应GraphSAGE的节点自身特征+一阶邻居特征+二阶邻居特征,但生成器generator.flow(all_nodes)只输出了1个张量,导致输入不匹配。核心原因是生成器配置与模型层数不匹配,或者模型构建时未正确绑定生成器的输入张量。
解决方案
1. 确认生成器num_samples参数与模型层数匹配
GraphSAGE的num_samples参数定义每一层采样的邻居数量,其长度需等于模型的层数(layer_sizes的长度)。例如,若模型是2层(layer_sizes=[32,32]),则num_samples应设为[10,5],此时生成器会自动生成节点自身+一阶10个邻居+二阶5个邻居三个输入张量:
generator = GraphSAGENodeGenerator(G, batch_size=50, num_samples=[10, 5])
2. 模型构建时绑定生成器的输入张量
创建GraphSAGE模型时,必须使用生成器的in_out_tensors()方法获取输入输出张量,确保模型输入与生成器输出对齐:
graphsage = GraphSAGE( layer_sizes=[32, 32], generator=generator, bias=True, dropout=0.5, ) # 获取生成器对应的输入张量和模型输出张量 x_inp, x_out = graphsage.in_out_tensors() # 构建最终分类模型 model = Model(inputs=x_inp, outputs=Dense(7, activation="softmax")(x_out))
3. 验证生成器输出格式
执行以下代码检查生成器的输出是否包含3个张量:
all_mapper = generator.flow(all_nodes) sample_input = next(iter(all_mapper)) print(f"生成器输出张量数量: {len(sample_input)}") # 应输出3
若输出为1,说明生成器配置错误,需重新调整num_samples参数并重新创建生成器。
4. 对齐版本兼容性
确保Stellargraph与TensorFlow的版本与官方文档一致,版本不兼容可能导致生成器与模型的输入输出逻辑异常。
内容的提问来源于stack exchange,提问作者Reza Akraminejad
相关产品推荐
相关产品推荐

