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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 08:33:24