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内部图)不在同一个图里,触发错误。
修复步骤
- 把特征转换逻辑移到
input_fn内部:让所有张量操作都在Estimator创建的图上下文里执行,避免跨图引用。 - 移除手动创建的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

