如何在不使用Keras API的TensorFlow中统计不可训练权重的数量
原生TensorFlow 1.x 统计不可训练权重的实现方案
你可以通过全局变量集合减去可训练变量集合的方式得到不可训练权重,完全不需要依赖Keras API,实现逻辑如下:
核心逻辑
TensorFlow 1.x中所有模型变量都会被加入tf.global_variables()返回的全局变量集合,tf.trainable_variables()返回的只是其中标记为可训练的子集,二者的差值就是不可训练权重的集合。
完整实现代码
import numpy as np import tensorflow as tf # 这里替换为你自己的自定义模型构造逻辑 x = np.zeros((1,16,16,3)) x_tf = tf.convert_to_tensor(x, np.float32) z_tf = tf.layers.conv2d(x_tf, filters=32, kernel_size=(3,3)) zz_tf = tf.layers.conv2d(z_tf, filters=32, kernel_size=(3,3)) # 统计可训练参数 trainable_vars = tf.trainable_variables() trainable_count = np.sum([np.prod(v.shape.as_list()) for v in trainable_vars]) # 统计不可训练参数 all_vars = tf.global_variables() non_trainable_vars = [var for var in all_vars if var not in trainable_vars] non_trainable_count = np.sum([np.prod(var.shape.as_list()) for var in non_trainable_vars]) # 结果输出 print('总参数数量: {:,}'.format(trainable_count + non_trainable_count)) print('可训练参数数量: {:,}'.format(trainable_count)) print('不可训练参数数量: {:,}'.format(non_trainable_count))
说明
- 该方法可以覆盖所有非训练权重场景,包括BatchNorm层的滑动均值/方差、ExponentialMovingAverage生成的影子变量等原生TF算子产生的非训练参数
- 调用
shape.as_list()是为了将TensorShape对象转为普通整数列表,避免numpy计算时出现类型报错
内容的提问来源于stack exchange,提问作者Zizi96
相关产品推荐
相关产品推荐

