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

TensorFlow Estimator预测报错:张量来自不同图问题求助

TensorFlow Estimator预测时张量跨图错误排查

我有一段TensorFlow代码,流程是用TFTransformOutput转换从TF.Examples的RecordIO文件读取的原始特征,再把转换后的特征输入到已预热的TensorFlow EstimatorV2中做预测。为了避免混用tf_v1和tf_v2导致图不匹配错误,我把整个流程放在单个TensorFlow图和会话上下文里,但还是一直遇到张量来自不同图的错误。以下是精简代码和报错信息,求帮忙排查。

精简代码

def main(argv: Sequence[str]) -> None:
...
  tf.compat.v1.disable_v2_behavior()

  def transform_features_fn(
      tf_transform_output,
      hparams,
      examples,
      used_features,
  ):
    """Creates an input function for serving based on ExampleMetadata."""
    feature_spec = tf_transform_output.raw_feature_spec()

    parsed_features = tf.io.parse_example(examples, feature_spec)
    transformed_features = tf_transform_output.transform_raw_features(
        parsed_features, drop_unused_features=True
    )
    return transformed_features

  # Define paths and parameters
  tf_transform_output_path = "/some/path"
  checkpoint_path = "some/path/2"
  testdata_filepattern = "some/path/3"

  # Load the transform output
  tf_transform_output = trainer_util.TFTransformOutput(tf_transform_output_path)

  # Define hyperparameters
  hparams = model_hparams.create_hparams()
  # Function to get records
  def get_records(records_list):
    for f in gfile.Glob(testdata_filepattern):
      with recordio.RecordReader(f) as rr:
        for record in rr:
          ex = example_pb2.Example()
          ex.ParseFromString(record)
          records_list.append(ex)
          return  # Return after one example for simplicity

  # Get records
  records_list = []
  get_records(records_list=records_list)

  # Define the graph context
  graph = tf.Graph()
  with graph.as_default():
    with tf.compat.v1.Session(graph=graph) as sess:
      # Serialize the example
      serialized_examples_tensor = tf.constant(
          [records_list[0].SerializeToString()]
      )  # Changed to batch of one serialized example
      # Transform features within the same graph
      transformed_features = transform_features_fn(
          tf_transform_output,
          hparams,
          serialized_examples_tensor,
          used_features,
      )

      # Define input function
      def input_fn():
        dataset = tf.data.Dataset.from_tensor_slices(transformed_features)
        dataset = dataset.batch(1)
        return dataset

      # Create the estimator within the same graph
      my_model = model.create_estimator(
          tf_transform_output,
          hparams,
          warm_start_from=checkpoint_path,
          model_dir=checkpoint_path,
      )
      predictions = my_model.predict(input_fn=input_fn)
      for prediction in predictions:
        print(prediction["predictions"])

报错信息

E0603 04:04:15.467269 3556985 app.py:668] Top-level exception: Tensor("batch_size:0", shape=(), dtype=int64, device=/device:CPU:0) must be from the same graph as Tensor("TensorSliceDataset:0", shape=(), dtype=variant) (graphs are <tensorflow.python.framework.ops.Graph object at 0x51cc1ab7b9c0> and <tensorflow.python.framework.ops.Graph object at 0x51cc24427bc0>).
.......
raise ValueError(
ValueError: Tensor("batch_size:0", shape=(), dtype=int64, device=/device:CPU:0) must be from the same graph as Tensor("TensorSliceDataset:0", shape=(), dtype=variant) (graphs are <tensorflow.python.framework.ops.Graph object at 0x51cc1ab7b9c0> and <tensorflow.python.framework.ops.Graph object at 0x51cc24427bc0>).


问题根源

Estimator的predict方法会自动创建新的图上下文,你手动指定的graph.as_default()不会被input_fn继承——因为input_fn是在Estimator内部调用的,它会使用Estimator自己创建的图,而不是你外部定义的那个图,导致transformed_features(属于外部图)和Dataset操作(属于Estimator内部图)不在同一个图里,触发错误。

修复步骤

  1. 把特征转换逻辑移到input_fn内部:让所有张量操作都在Estimator创建的图上下文里执行,避免跨图引用。
  2. 移除手动创建的Graph和Session上下文:Estimator会自己管理图和会话,手动指定反而会冲突。

修复后的代码示例:

def main(argv: Sequence[str]) -> None:
    tf.compat.v1.disable_v2_behavior()

    def transform_features_fn(
        tf_transform_output,
        hparams,
        examples,
        used_features,
    ):
        """Creates an input function for serving based on ExampleMetadata."""
        feature_spec = tf_transform_output.raw_feature_spec()

        parsed_features = tf.io.parse_example(examples, feature_spec)
        transformed_features = tf_transform_output.transform_raw_features(
            parsed_features, drop_unused_features=True
        )
        return transformed_features

    # Define paths and parameters
    tf_transform_output_path = "/some/path"
    checkpoint_path = "some/path/2"
    testdata_filepattern = "some/path/3"

    # Load the transform output
    tf_transform_output = trainer_util.TFTransformOutput(tf_transform_output_path)

    # Define hyperparameters
    hparams = model_hparams.create_hparams()
    # Function to get records
    def get_records(records_list):
        for f in gfile.Glob(testdata_filepattern):
            with recordio.RecordReader(f) as rr:
                for record in rr:
                    ex = example_pb2.Example()
                    ex.ParseFromString(record)
                    records_list.append(ex)
                    return  # Return after one example for simplicity

    # Get records
    records_list = []
    get_records(records_list=records_list)

    # 将特征转换逻辑移到input_fn内部,移除手动Graph/Session
    def input_fn():
        # 在Estimator的图上下文内创建序列化张量
        serialized_examples_tensor = tf.constant(
            [records_list[0].SerializeToString()]
        )
        transformed_features = transform_features_fn(
            tf_transform_output,
            hparams,
            serialized_examples_tensor,
            used_features,
        )
        dataset = tf.data.Dataset.from_tensor_slices(transformed_features)
        dataset = dataset.batch(1)
        return dataset

    # 创建Estimator,无需手动指定图
    my_model = model.create_estimator(
        tf_transform_output,
        hparams,
        warm_start_from=checkpoint_path,
        model_dir=checkpoint_path,
    )
    predictions = my_model.predict(input_fn=input_fn)
    for prediction in predictions:
        print(prediction["predictions"])

额外说明

  • Estimator API设计上会自动管理图和会话,手动嵌套Graph/Session上下文会打破这种机制,引发跨图问题。
  • 所有需要在Estimator中使用的张量操作,都应该放在input_fn或模型定义函数内部,确保它们处于同一个图上下文。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 20:52:10