tf.get_operations返回值解析及相关代码逻辑疑问咨询
嘿,我来一步步帮你把这些问题掰明白~
通俗解读tf.get_operations()
官方文档有点绕,翻译成大白话就是:
- 这个函数相当于给当前TensorFlow计算图拍了一张全量操作快照,返回的是图里所有已经定义好的
tf.Operation对象列表。 - 你可以直接修改列表里单个操作的属性(比如调整某个操作的参数),但要是你对这个返回的列表本身做增删(比如插个新操作进去、删掉列表里的某一项),完全不会影响原图里的操作——因为它只是个拷贝出来的“清单副本”而已。
- 放心在多线程环境里用它,不会出现并发问题。
关于打印出的操作来源
你猜的没错!这些tf.Operation对象确实是从当前程序加载/构建的TensorFlow计算图里读取出来的。至于是不是来自冻结图,得看你当前的图本身:
- 冻结图的特点是把所有可训练变量都转成了
Const类型的操作(比如你看到的BatchNorm/gamma就是个Const),但普通的计算图里也可能存在原本就定义的常量操作。 - 简单说:如果你的图是从冻结的
.pb文件加载的,那这些操作肯定来自冻结图;如果是动态构建的图(比如自己写的模型),那就是来自你构建的计算图。
拆解集合推导式的工作原理
你贴的这段代码all_tensor_names = {output.name for op in ops for output in op.outputs}是Python的集合推导式,说白了就是用简洁的写法遍历两层循环,收集唯一的张量名称,拆解成普通循环的话就是:
all_tensor_names = set() for op in ops: # 遍历当前操作的所有输出张量 for output in op.outputs: # 把张量的名称添加到集合里(集合自动去重) all_tensor_names.add(output.name)
具体逻辑是:
- 先遍历
tf.get_operations()返回的每一个操作op - 每个操作执行后都会输出一个或多个张量(比如Conv2D会输出特征图张量,Const会输出常量张量),所以再遍历这个操作的所有输出
output - 把每个输出张量的
name属性提取出来,放到集合中——集合的特性是自动去重,所以最终得到的是当前图里所有操作输出张量的唯一名称集合
内容的提问来源于stack exchange,提问作者Simon Kiely
相关产品推荐
相关产品推荐

