如何使用Python Spektral库DisjointLoader为全连接网络而非GIN喂数据
问题根源
你遇到的维度不匹配是图分类任务的特征层级差异导致的:
- DisjointLoader返回的
x是batch内所有图的节点拼接结果,你的batch大小为7(每张图18个节点),所以x的维度是[7*18=126, 7],对应126个节点的7维属性 - 你直接对节点维度做全连接计算,输出维度为
[126, 类别数],但图分类任务的标签是每张图对应1个,维度为[7, 类别数],二者维度不匹配触发报错
解决方案
你需要把节点级特征聚合为图级特征,再输入全连接层做计算,有两种常用实现方式:
方式1:节点维度映射+全局池化
先用全连接对每个节点特征做维度变换,再用Spektral自带的全局池化层按图索引i聚合同一张图的所有节点特征,得到单图的表征向量后再做分类:
from tensorflow.keras.models import Model from tensorflow.keras.layers import Dense, Dropout from spektral.layers import GlobalAvgPool class FCN0(Model): def __init__(self, channels, outputs): super().__init__() self.dense1 = Dense(channels, activation="relu") self.dropout = Dropout(0.5) self.dense2 = Dense(channels*3, activation="relu") self.pool = GlobalAvgPool() # 分类任务输出层激活函数需要对应调整:二分类用sigmoid,多分类用softmax,回归用linear self.dense3 = Dense(outputs, activation="softmax") def call(self, inputs): x, a, i = inputs x = self.dense1(x) x = self.dropout(x) x = self.dense2(x) # 按图索引i聚合节点特征,得到图级表征 x = self.pool([x, i]) return self.dense3(x)
方式2:单图特征展平后输入全连接
如果要实现完全不利用图结构、纯特征拼接的MLP基线,你可以先把每个图的18个节点特征展平为单个一维向量,再做全连接计算,和GIN的对比更严谨:
import tensorflow as tf from tensorflow.keras.models import Model from tensorflow.keras.layers import Dense, Dropout class FCN0(Model): def __init__(self, channels, outputs): super().__init__() self.dense1 = Dense(channels, activation="relu") self.dropout = Dropout(0.5) self.dense2 = Dense(channels*3, activation="relu") self.dense3 = Dense(outputs, activation="softmax") def call(self, inputs): x, a, i = inputs # 按图索引重组节点特征 batch_size = tf.reduce_max(i) + 1 x = tf.scatter_nd(indices=i[:, None], updates=x, shape=(batch_size, 18, 7)) # 每张图的节点特征展平为18*7=126维向量 x = tf.reshape(x, (batch_size, -1)) x = self.dense1(x) x = self.dropout(x) x = self.dense2(x) return self.dense3(x)
内容的提问来源于stack exchange,提问作者mindstorm84
相关产品推荐
相关产品推荐

