TensorFlow如何批量获取所有名称以Add开头的张量?
如何简洁收集TensorFlow计算图中所有名称以Add开头的张量
当然可以解决这个麻烦!当计算图变复杂时,手动逐个指定Add:0、Add_1:0这种名称确实低效又容易出错,我们可以通过遍历计算图+字符串匹配的方式自动收集目标张量,两种常用方法如下:
方法1:直接遍历所有张量并筛选
这种方式最简单直接,先获取图中所有张量,再筛选名称以Add开头的项:
import tensorflow as tf # 获取默认计算图 graph = tf.get_default_graph() # 列表推导式一键筛选 add_tensors = [tensor for tensor in graph.get_all_tensors() if tensor.name.startswith("Add")]
执行后,add_tensors就会包含所有名称以Add开头的张量(比如Add:0、Add_1:0、Add_2:0等)。
方法2:先筛选Add操作,再收集其输出张量
如果你想更精准地针对Add操作的输出张量(避免其他巧合名称开头的张量),可以先找到所有名称以Add开头的操作,再收集它们的输出:
import tensorflow as tf graph = tf.get_default_graph() add_tensors = [] # 遍历图中所有操作 for op in graph.get_operations(): # 匹配操作名称以Add开头的项 if op.name.startswith("Add"): # 收集该操作的所有输出张量(Add操作通常只有1个输出) add_tensors.extend(op.outputs)
针对TensorFlow 2.x的注意事项
TF2.x默认是即时执行模式,计算图是动态构建的,如果你用的是TF2.x,需要先通过tf.function装饰函数触发图构建,再获取计算图:
import tensorflow as tf @tf.function def build_graph(): a = tf.add(1, 2, name="Add") b = tf.add(3, 4, name="Add_1") c = tf.add(a, b, name="Add_2") return c # 调用函数触发图构建 build_graph() # 获取构建好的计算图 graph = build_graph.get_concrete_function().graph # 用上面任意一种方法筛选张量 add_tensors = [tensor for tensor in graph.get_all_tensors() if tensor.name.startswith("Add")]
这样不管你的计算图有多复杂,都能自动完成张量收集,再也不用手动逐个写名称啦!
内容的提问来源于stack exchange,提问作者Gilfoyle
相关产品推荐
相关产品推荐

