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

如何将Stellargraph对象传入Keras模型?解决数据适配报错

解决Stellargraph图对象传入Keras模型的问题

Stellargraph的StellarGraph对象不能直接通过Keras的Input()层传入,因为Keras的数据适配器无法识别这种自定义类型。以下是正确的处理方式:

1. 基于邻接矩阵的模型(如GCN)

先将图数据转换为Keras可处理的张量格式:

  • 提取节点特征矩阵:
    node_features = sg_graph.node_features()  # 返回numpy数组或张量
    
  • 提取归一化后的邻接矩阵:
    adj_matrix = sg_graph.to_adjacency_matrix(weighted=True, normalized=True)
    
  • 用Input()分别定义输入层:
    from tensorflow.keras.layers import Input
    
    node_input = Input(shape=(node_features.shape[1],))
    adj_input = Input(shape=(node_features.shape[0],))
    
  • 将这两个输入传入对应图神经网络层,再构建后续模型结构。

2. 基于采样的模型(如GraphSAGE)

使用Stellargraph提供的专用生成器处理输入,无需手动调用Input():

from stellargraph.mapper import GraphSAGENodeGenerator
from stellargraph.layer import GraphSAGE

# 初始化生成器
generator = GraphSAGENodeGenerator(sg_graph, batch_size=32, num_samples=[10, 5])
# 生成训练数据迭代器
train_gen = generator.flow(train_nodes, train_labels)

# 获取模型的输入输出张量
graphsage = GraphSAGE(layer_sizes=[64, 32], generator=generator, activation="relu")
x_inp, x_out = graphsage.in_out_tensors()

这里的x_inp就是模型的输入张量,可直接用于构建后续Keras模型。

通用注意事项

  • 禁止直接将StellarGraph对象传给Keras的fit()或Input(),必须转换为张量、数组,或使用Stellargraph生成器。
  • 报错中的NoneType大概率是传入的标签或其他参数为空,需检查数据完整性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 00:49:54