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

基于骨架的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操作无法找到对应节点。

具体解决步骤

  1. 检查边集邻接数据
    查看bones边集的邻接矩阵(#index.1对应接收节点),确认是否存在值为25的ID。如果是数据生成时采用了1-based节点编号(比如节点从1到25),需要转换为0-based(减1)。

  2. 添加数据校验逻辑
    在输入模型前,添加代码验证节点索引合法性:

    # 假设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}"
    
  3. 临时修复(不推荐,仅用于验证)
    如果暂时无法修改数据源,可以在模型开头添加索引裁剪层,强制将接收节点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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 23:42:10