基于骨架的TFGNN训练时model.fit()出现索引错误求助
基于TFGNN的骨架动作识别模型报错排查与解决
问题背景
使用TFGNN库构建基于骨架的图神经网络用于动作识别,改编官方Colab示例代码后运行持续报错。已知body节点集含25个节点,bones边集含24条边,尝试过重构图结构、修改图更新层均未解决问题。
输入GraphSchema
GraphTensorSpec({'context': ContextSpec({'features': {}, '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), 'z_dim': TensorSpec(shape=(None, 1), dtype=tf.float32, name=None), 'y_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)}}, TensorShape([]), tf.int32, None)
模型代码
def _build_model( # To be called with the build_model_graph_tensor_spec from above. graph_tensor_spec, # Dimensions of initial states. node_dim=128, # Dimensions for message passing. message_dim=128, next_state_dim=128, # Dimension for the logits. num_classes=3, # Other hyperparameters. l2_regularization=6e-6, dropout_rate=0.2, use_layer_normalization=True, ): # Model building with Keras's Functional API starts with an input object # (a placeholder for future inputs). This works for composite tensors, too. graph = input_graph = tf.keras.layers.Input(type_spec=graph_tensor_spec) graph = graph.merge_batch_to_components() def set_initial_node_state(node_set, node_set_name): if node_set_name == "body": feature_x_embedding = tf.keras.layers.Dense(node_dim, activation="relu") feature_y_embedding = tf.keras.layers.Dense(node_dim, activation="relu") feature_z_embedding = tf.keras.layers.Dense(node_dim, activation="relu") concatenated_features = tf.keras.layers.Concatenate()( [feature_x_embedding(node_set["x_dim"]), feature_y_embedding(node_set["y_dim"]), feature_z_embedding(node_set["z_dim"])]) return concatenated_features graph = tfgnn.keras.layers.MapFeatures( node_sets_fn=set_initial_node_state, name="init_states")(graph) # Abbreviations for repeated building blocks in the GNN. def dense(units, *, use_layer_normalization=False): """A Dense layer with regularization (L2 and Dropout) and normalization.""" regularizer = tf.keras.regularizers.l2(l2_regularization) result = tf.keras.Sequential([ tf.keras.layers.Dense( units, activation="relu", use_bias=True, kernel_regularizer=regularizer, bias_regularizer=regularizer), tf.keras.layers.Dropout(dropout_rate)]) if use_layer_normalization: result.add(tf.keras.layers.LayerNormalization()) return result for i in range(4): graph = tfgnn.keras.layers.GraphUpdate( node_sets={ "body": tfgnn.keras.layers.NodeSetUpdate( {"bones": tfgnn.keras.layers.SimpleConv( tf.keras.layers.Dense(128, "relu"), "mean", receiver_tag=tfgnn.TARGET)}, tfgnn.keras.layers.NextStateFromConcat(tf.keras.layers.Dense(128))) } )(graph) root_states = tfgnn.keras.layers.ReadoutFirstNode(node_set_name="body")(graph) logits = tf.keras.layers.Dense(num_classes)(root_states) return tf.keras.Model(input_graph, logits)
报错信息
Node: 'while/model_1/graph_update_4/node_set_update_4/simple_conv_4/UnsortedSegmentMean/UnsortedSegmentSum' segment_ids[44] = 25 is out of range [0, 25) [[{{node while/model_1/graph_update_4/node_set_update_4/simple_conv_4/UnsortedSegmentMean/UnsortedSegmentSum}}]] [Op:__inference_train_function_10449]
问题分析与解决方法
错误根源
报错显示segment_ids[44] = 25超出范围[0,25),说明bones边集中某条边的接收节点ID为25,但body节点集只有25个节点,索引范围是0~24(节点索引从0开始计数),导致UnsortedSegmentSum操作无法找到对应节点。
具体解决步骤
检查边集邻接数据
查看bones边集的邻接矩阵(#index.1对应接收节点),确认是否存在值为25的ID。如果是数据生成时采用了1-based节点编号(比如节点从1到25),需要转换为0-based(减1)。添加数据校验逻辑
在输入模型前,添加代码验证节点索引合法性:# 假设graph_tensor是输入的GraphTensor max_receiver_id = tf.reduce_max(graph_tensor.edge_sets['bones'].adjacency[tfgnn.TARGET]) num_body_nodes = tf.reduce_sum(graph_tensor.node_sets['body'].sizes) assert max_receiver_id < num_body_nodes, f"非法接收节点ID: {max_receiver_id}, 节点总数: {num_body_nodes}"临时修复(不推荐,仅用于验证)
如果暂时无法修改数据源,可以在模型开头添加索引裁剪层,强制将接收节点ID限制在合法范围内:def clip_edge_indices(graph): adjacency = graph.edge_sets['bones'].adjacency clipped_target = tf.clip_by_value(adjacency[tfgnn.TARGET], 0, tf.reduce_sum(graph.node_sets['body'].sizes)-1) new_adjacency = adjacency.replace({tfgnn.TARGET: clipped_target}) return graph.replace_edge_sets({'bones': graph.edge_sets['bones'].replace(adjacency=new_adjacency)}) # 在merge_batch_to_components后添加 graph = clip_edge_indices(graph)注意:这只是临时方案,根源问题还是要修正数据生成逻辑。
内容的提问来源于stack exchange,提问作者Marianna
相关产品推荐
相关产品推荐

