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

在StellarGraph中使用PaddedGraphGenerator自定义训练、验证、测试集训练GCN时遇ValueError问题排查

问题分析与解决方案

嘿,这个错误我之前也碰到过,核心原因很明确:你第二次调用model.fit()的时候直接传了StellarGraph对象列表graphs_train,Keras根本没法直接处理这种类型的输入——StellarGraph的GCN模型必须通过你之前定义的PaddedGraphGenerator来生成模型能识别的张量数据,不能直接喂原始的图对象。

具体问题拆解

咱们看你的代码,有两处明显的问题:

  1. 你连续调用了两次model.fit():第一次用train_gen是正确的打开方式,但第二次直接传graphs_train和标签,完全不符合模型的输入要求,这就是报错的根源。
  2. 代码里提到了valid_gen但压根没定义它,后面的model.evaluate(valid_gen)肯定也会报错。

修正后的完整代码

下面是调整后的代码,把这些问题都解决了:

from stellargraph.mapper import PaddedGraphGenerator
from stellargraph.layer import GCNSupervisedGraphClassification
from tensorflow.keras import Model
from tensorflow.keras.layers import Dense
from tensorflow.keras.optimizers import Adam
from tensorflow.keras.losses import binary_crossentropy
from tensorflow.keras.callbacks import EarlyStopping

# 假设你已经准备好这些数据:
# graphs_train: 训练集的StellarGraph对象列表
# graphs_val: 验证集的StellarGraph对象列表
# graphs_test: 测试集的StellarGraph对象列表
# graphs_train_labels: 训练集标签数组
# graphs_val_labels: 验证集标签数组
# graphs_test_labels: 测试集标签数组

# 1. 初始化生成器:PaddedGraphGenerator需要知道所有图的结构来确定统一的padding尺寸,所以把所有图放一起
all_graphs = graphs_train + graphs_val + graphs_test
generator = PaddedGraphGenerator(graphs=all_graphs)

# 2. 给三个数据集分配对应的索引,创建生成器
train_indices = list(range(len(graphs_train)))
val_indices = list(range(len(graphs_train), len(graphs_train)+len(graphs_val)))
test_indices = list(range(len(graphs_train)+len(graphs_val), len(all_graphs)))

train_gen = generator.flow(train_indices, targets=graphs_train_labels, batch_size=35)
val_gen = generator.flow(val_indices, targets=graphs_val_labels, batch_size=35)
test_gen = generator.flow(test_indices, targets=graphs_test_labels, batch_size=35)

# 3. 早停回调:监控验证集损失,耐心等待20个epoch,恢复最优权重
es = EarlyStopping(monitor="val_loss", min_delta=0, patience=20, restore_best_weights=True)

# 4. 定义GCN模型结构
gc_model = GCNSupervisedGraphClassification(
    layer_sizes=[64, 64], 
    activations=["relu", "relu"], 
    generator=generator, 
    dropout=0.5
)
x_inp, x_out = gc_model.in_out_tensors()
predictions = Dense(units=32, activation="relu")(x_out)
predictions = Dense(units=16, activation="relu")(predictions)
predictions = Dense(units=1, activation="sigmoid")(predictions)

model = Model(inputs=x_inp, outputs=predictions)
model.compile(optimizer=Adam(0.001), loss=binary_crossentropy, metrics=["acc"])

# 5. 训练模型:只需要一次fit,用训练生成器,验证数据用val_gen,加上早停回调
history = model.fit(
    train_gen, 
    epochs=100,  # 调大没关系,早停会自动终止
    validation_data=val_gen, 
    verbose=1,
    callbacks=[es]
)

# 6. 在测试集上评估模型性能
test_metrics = model.evaluate(test_gen, verbose=1)
test_acc = test_metrics[model.metrics_names.index("acc")]
print(f"Test Accuracy: {test_acc}")

关键修正点说明

  • 删掉了第二次model.fit()调用:直接传StellarGraph列表是完全错误的,训练和评估都必须用generator.flow()生成的数据集。
  • 补上了验证集生成器val_gen:你的代码里提到但没定义,现在完整覆盖了训练、验证、测试三个环节的生成器逻辑。
  • 调整了生成器初始化方式:PaddedGraphGenerator需要所有图的结构信息来确定padding的统一尺寸,所以把所有图都传入初始化,再用索引区分不同数据集。
  • 把早停回调加到了正确的fit()里:之前的早停没发挥作用,现在放到训练过程中,会自动监控验证集损失,避免过拟合。

另外关于图的创建方式,只要你的graphs是合法的StellarGraph对象列表(比如用StellarGraph.from_networkx()构建的),那图的创建本身没问题——这次报错的核心是输入数据的传递方式,不是图的创建。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.27 19:37:47