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

使用TensorFlow子类API结合Spektral层绘制图神经网络模型图

解决图神经网络模型可视化问题

针对你的GIN模型,由于输入包含多组不同维度的张量(节点特征、邻接矩阵、Batch索引),需要为每个输入创建对应的Input层,同时确保模型的所有层被Keras正确追踪。以下是具体修改方案:

1. 修改GIN0类,添加build_graph方法

import tensorflow as tf
from tensorflow.keras import Model, Input, Dense, Dropout
# 假设GINConv和GlobalAvgPool来自相关GNN库(如tf_geometric)
from tf_geometric.layers import GINConv, GlobalAvgPool

class GIN0(Model):
    def __init__(self, channels, n_layers):
        super().__init__()
        self.conv1 = GINConv(channels, epsilon=0, mlp_hidden=[channels, channels])
        # 用ListWrapper包装卷积层列表,让Keras能识别并追踪这些层
        self.convs = tf.keras.layers.ListWrapper([
            GINConv(channels, epsilon=0, mlp_hidden=[channels, channels])
            for _ in range(1, n_layers)
        ])
        self.pool = GlobalAvgPool()
        self.dense1 = Dense(channels, activation="relu")
        self.dropout = Dropout(0.5)
        self.dense2 = Dense(channels, activation="relu")

    def call(self, inputs):
        x, a, i = inputs
        x = self.conv1([x, a])
        for conv in self.convs:
            x = conv([x, a])
        x = self.pool([x, i])
        x = self.dense1(x)
        x = self.dropout(x)
        return self.dense2(x)
    
    def build_graph(self, input_shapes):
        # input_shapes为三元组:(节点特征形状, 邻接矩阵形状, Batch索引形状)
        x_shape, a_shape, i_shape = input_shapes
        # 为每个输入创建Input层
        x_in = Input(shape=x_shape)
        a_in = Input(shape=a_shape)
        i_in = Input(shape=i_shape)
        # 调用模型的call方法得到输出
        output = self.call([x_in, a_in, i_in])
        # 返回Functional API模型,用于可视化
        return Model(inputs=[x_in, a_in, i_in], outputs=output)

2. 使用示例

固定输入维度场景

如果任务中节点数和特征维度是固定的:

# 定义输入形状:节点数32,节点特征维度128,邻接矩阵32x32,Batch索引对应32个节点
x_shape = (32, 128)
a_shape = (32, 32)
i_shape = (32,)

# 初始化模型
model = GIN0(channels=64, n_layers=3)
# 构建可视化用的模型
graph_model = model.build_graph((x_shape, a_shape, i_shape))

# 打印模型结构
graph_model.summary()

# 绘制模型图(需要安装pydot和graphviz)
from tensorflow.keras.utils import plot_model
plot_model(graph_model, to_file='gin_model.png', show_shapes=True, show_layer_names=True)

动态输入维度场景

如果节点数不固定(支持可变大小的图输入),可以用None表示动态维度:

# 动态形状:节点数可变,节点特征维度128
x_shape = (None, 128)
a_shape = (None, None)  # 邻接矩阵的行列数都可变
i_shape = (None,)       # 节点数可变

graph_model = model.build_graph((x_shape, a_shape, i_shape))
graph_model.summary()

关键说明

  • tf.keras.layers.ListWrapper:用于包装卷积层列表,确保Keras能识别这些子层,否则模型summary中会缺失这些层的信息。
  • 多Input层:GNN的多模态输入需要对应多个Input层,构建Functional API模型后就能正常使用Keras的可视化工具。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 23:35:42