TensorFlow2中不依赖compat.v1如何通过名称获取指定张量
TF2 原生按张量名称获取张量句柄方案
核心原因说明
你调用tf.compat.v1.get_default_graph().as_graph_def().node返回空列表,本质是两个原因:
- TF2默认开启Eager Execution即时执行模式,不会提前构建全局静态计算图,算子执行完不会在全局默认图留存节点
- 即使你用
tf.function构建静态图,生成的计算图也是和具体函数绑定的独立图,不会挂载到全局默认图上,直接查全局默认图当然拿不到内容
方案1:静态图场景(使用tf.function时,原生API,无需compat模块)
TF2 原生保留了tf.Graph类的get_tensor_by_name()方法,不属于兼容模块,可以直接使用,前提是你拿到张量所在的正确计算图对象,步骤如下:
- 给模型前向传播加
tf.function装饰,传入符合要求的输入规格做追踪,生成绑定具体计算图的ConcreteFunction - 从
ConcreteFunction的.graph属性拿到对应独立计算图 - 直接调用图对象的
get_tensor_by_name()方法,传入完整张量名(带:0这类输出索引后缀)即可拿到句柄
对应示例代码:
import tensorflow as tf class MyModel(tf.keras.Model): def __init__(self): super().__init__() self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu) self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax) def call(self, inputs): # 显式定义目标张量,指定名称 my_happy_tensor = tf.identity(inputs, name="my_happy_tensor") x = self.dense1(my_happy_tensor) return self.dense2(x) model = MyModel() # 定义输入规格,追踪生成具体函数和对应静态图 input_spec = tf.TensorSpec(shape=(None, 10), dtype=tf.float32) concrete_func = tf.function(model).get_concrete_function(input_spec) # 拿到对应计算图,直接调用原生get_tensor_by_name target_graph = concrete_func.graph target_tensor = target_graph.get_tensor_by_name("my_happy_tensor:0")
方案2:纯Eager模式场景
纯Eager模式下没有全局静态图,也不存在全局的「张量名-张量实例」映射表,张量生命周期完全由Python引用计数决定,没有被持有的张量会被直接回收,不存在事后靠名称检索张量的原生API。
这种场景要获取特定张量,最稳妥的方式是在模型定义时提前把需要的张量引用存到实例属性里,前向传播后直接读取即可,示例:
class MyModel(tf.keras.Model): def __init__(self): super().__init__() self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu) self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax) # 提前定义属性存储目标张量引用 self.target_tensors = {} def call(self, inputs): my_happy_tensor = tf.identity(inputs, name="my_happy_tensor") # 把需要后续访问的张量存到属性里 self.target_tensors["my_happy_tensor:0"] = my_happy_tensor x = self.dense1(my_happy_tensor) return self.dense2(x) # Eager模式执行前向 model = MyModel() test_input = tf.random.normal((2, 10)) _ = model(test_input) # 直接读取存好的张量引用 my_happy_tensor = model.target_tensors["my_happy_tensor:0"]
注意事项
tf.Graph.get_tensor_by_name()是TF2核心API,不属于tf.compat.v1兼容模块,可以放心使用- 传入的张量名必须带输出索引后缀,格式为
算子名:输出序号,比如名为my_happy_tensor的算子的第一个输出对应名称是my_happy_tensor:0 - 不要尝试从全局默认图中查找
tf.function内的张量,这类函数的计算图是独立隔离的,必须从对应ConcreteFunction的.graph属性获取
内容的提问来源于stack exchange,提问作者ShlomiF
相关产品推荐
相关产品推荐

