使用gcn_conv.GCNConv执行GraphUpdate返回NaN张量问题排查
GCNConv执行GraphUpdate时出现NaN的排查分析
问题背景
使用gcn_conv.GCNConv执行GraphUpdate时,返回张量中部分行全为NaN,但输入张量无NaN值(范围-1到+1的float32),换用SimpleConv无异常,另一组同结构数据也正常。相关代码如下:
g = graphs[0] hidden_size = 32 context_features = g.context.get_features_dict() label = context_features.pop('label') new_graph = g.replace_features(context=context_features) def set_initial_node_state(node_set, node_set_name): features = [ tf.keras.layers.Dense(hidden_size, activation="relu")(node_set['x_dim']), tf.keras.layers.Dense(hidden_size, activation="relu")(node_set['y_dim']), tf.keras.layers.Dense(hidden_size, activation="relu")(node_set['z_dim']) ] return tf.keras.layers.Concatenate()(features) new_graph = tfgnn.keras.layers.MapFeatures( node_sets_fn=set_initial_node_state)(new_graph) result_gcn = tfgnn.keras.layers.GraphUpdate( node_sets = { 'body': tfgnn.keras.layers.NodeSetUpdate({ 'bones': gcn_conv.GCNConv( units = hidden_size)}, tfgnn.keras.layers.NextStateFromConcat( tf.keras.layers.Dense(hidden_size)))})(new_graph)
图结构的TensorSpec如下:
'GraphTensorSpec({'context': ContextSpec({'features': {'label': TensorSpec(shape=(1, 10), dtype=tf.float32, name=None)}, 'sizes': TensorSpec(shape=(1,), dtype=tf.int32, name=None)}, TensorShape([]), tf.int32, None), 'node_sets': {'body': NodeSetSpec({'features': {'x_dim': TensorSpec(shape=(None, 1), dtype=tf.float32, name=None), 'y_dim': TensorSpec(shape=(None, 1), dtype=tf.float32, name=None), 'z_dim': TensorSpec(shape=(None, 1), dtype=tf.float32, name=None)}, 'sizes': TensorSpec(shape=(1,), dtype=tf.int32, name=None)}, TensorShape([]), tf.int32, None)}, 'edge_sets': {'bones': EdgeSetSpec({'features': {}, 'sizes': TensorSpec(shape=(1,), dtype=tf.int32, name=None), 'adjacency': AdjacencySpec({'#index.0': TensorSpec(shape=(None,), dtype=tf.int32, name=None), '#index.1': TensorSpec(shape=(None,), dtype=tf.int32, name=None)}, TensorShape([]), tf.int32, {'#index.0': 'body', '#index.1': 'body'})}, TensorShape([]), tf.int32, None), 'temporal': EdgeSetSpec({'features': {}, 'sizes': TensorSpec(shape=(1,), dtype=tf.int32, name=None), 'adjacency': AdjacencySpec({'#index.0': TensorSpec(shape=(None,), dtype=tf.int32, name=None), '#index.1': TensorSpec(shape=(None,), dtype=tf.int32, name=None)}, TensorShape([]), tf.int32, {'#index.0': 'body', '#index.1': 'body'})}, TensorShape([]), tf.int32, None)}}, TensorShape([]), tf.int32, None)'
可能导致NaN的特定操作分析
1. GCN归一化时的分母为0
GCN核心计算依赖邻接矩阵的归一化,通常公式为:
$$\text{output} = \hat{D}^{-1/2} \hat{A} \hat{D}^{-1/2} X W$$
其中$\hat{A}$是添加自环后的邻接矩阵,$\hat{D}$是$\hat{A}$的度矩阵。如果某个节点的度(含自环)为0,$\hat{D}^{-1/2}$对应位置的元素会是无穷大,计算后直接产生NaN。
而SimpleConv通常采用简单求和、平均等聚合方式,不涉及这种带分母的归一化操作,因此不会触发该问题,你遇到的"部分行全NaN"正好对应这些度为0的孤立节点。
2. GCNConv未处理0度节点的边界情况
如果你的gcn_conv.GCNConv实现没有在归一化步骤中添加小的epsilon(如$1e-8$)来避免除以0,或者没有自动给节点添加自环,那么当图中存在孤立节点时,必然会出现NaN。
3. 特征计算中的极端值传播(可能性较低)
虽然输入特征范围是-1到1,但经过三层Dense+ReLU后,部分节点的特征可能因权重初始化不当被放大到极端值,后续GCN的矩阵乘法可能触发数值溢出/下溢产生NaN。但这种情况通常会导致全局NaN,而非部分行,因此优先级低于前两点。
验证与解决方法
- 检查孤立节点:统计
body节点集中每个节点在bones边集中的入度+出度,确认是否存在度为0的节点。 - 强制添加自环:在输入图中给每个
body节点添加自环,确保所有节点的度至少为1。 - 修改归一化逻辑:在
GCNConv的归一化步骤中,给度矩阵的对角线元素加上一个极小值(如$1e-8$),避免除以0。 - 检查GCNConv实现:确认该实现是否默认添加自环,若没有则手动开启相关参数(如
add_self_loops=True)。
内容的提问来源于stack exchange,提问作者Marianna
相关产品推荐
相关产品推荐

