如何从TensorFlow变量列表提取去重的BatchNorm截断名称
解决TensorFlow变量BatchNorm路径提取与去重问题
我来帮你实现这个需求,下面是一步步的解决方案和代码示例:
实现步骤
- 提取变量名称:从每个TensorFlow变量中获取完整名称,并去除末尾的
:0后缀(TensorFlow变量默认的索引标识)。 - 筛选含BatchNorm的名称:只保留包含
BatchNorm关键词的变量名。 - 截断到BatchNorm部分:对每个符合条件的名称,截取到
BatchNorm字符串的结尾,得到对应的BatchNorm模块路径。 - 去重并生成结果列表:用集合自动去重,再转换为列表(可选排序保证顺序)。
代码示例
import tensorflow as tf # 你的变量列表(模拟示例) vars_list = [ tf.Variable('resnet_v1_101/conv1/weights:0', shape=(7,7,3,64)), tf.Variable('resnet_v1_101/conv1/BatchNorm/beta:0', shape=(64,)), tf.Variable('resnet_v1_101/conv1/BatchNorm/gamma:0', shape=(64,)), tf.Variable('resnet_v1_101/conv1/BatchNorm/moving_mean:0', shape=(64,)), tf.Variable('resnet_v1_101/conv1/BatchNorm/moving_variance:0', shape=(64,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/shortcut/weights:0', shape=(1,1,64,256)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/shortcut/BatchNorm/beta:0', shape=(256,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/shortcut/BatchNorm/gamma:0', shape=(256,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/shortcut/BatchNorm/moving_mean:0', shape=(256,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/shortcut/BatchNorm/moving_variance:0', shape=(256,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/conv1/weights:0', shape=(1,1,64,64)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/conv1/BatchNorm/beta:0', shape=(64,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/conv1/BatchNorm/gamma:0', shape=(64,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/conv1/BatchNorm/moving_mean:0', shape=(64,)), tf.Variable('resnet_v1_101/block1/unit_1/bottleneck_v1/conv1/BatchNorm/moving_variance:0', shape=(64,)) ] # 核心处理逻辑 batch_norm_paths = set() for var in vars_list: # 去掉变量名末尾的":0"后缀 clean_name = var.name.split(':')[0] if 'BatchNorm' in clean_name: # 找到BatchNorm的结束位置并截断 bn_pos = clean_name.index('BatchNorm') truncated_path = clean_name[:bn_pos + len('BatchNorm')] batch_norm_paths.add(truncated_path) # 转换为有序列表(可选,按字典序排序) result = sorted(batch_norm_paths) print(result)
输出结果
运行后会得到你预期的列表:
['resnet_v1_101/block1/unit_1/bottleneck_v1/conv1/BatchNorm', 'resnet_v1_101/block1/unit_1/bottleneck_v1/shortcut/BatchNorm', 'resnet_v1_101/conv1/BatchNorm']
关键细节说明
- 去除
:0后缀:TensorFlow变量的name属性默认会带上:0标识变量的索引,用split(':')[0]可以快速去除。 - 集合去重:利用Python集合的特性自动过滤重复的BatchNorm路径,避免手动判断重复。
- 精准截断:通过
index找到BatchNorm的起始位置,加上字符串长度确保截断后的路径正好到BatchNorm结尾,不会多也不会少。
内容的提问来源于stack exchange,提问作者Jame
相关产品推荐
相关产品推荐

