TensorFlow1.4中如何访问类内神经网络层并保存训练权重
TensorFlow 1.4 自定义网络权重保存与复用方案
你的网络是基于tf.variable_scope+tf.get_variable搭建的,完全适配TF1.x原生的变量管理逻辑,不需要重构现有类结构,就能实现权重存读,跳过重复训练直接做推理。
实现逻辑
所有通过tf.get_variable创建的网络参数,都会被自动加入TF的全局变量集合,配合内置的tf.train.Saver工具,不需要手动逐层遍历网络层,就能完成权重的序列化保存、加载。
具体操作步骤
训练阶段保存权重
训练循环跑完、权重收敛之后,在当前运行的Session中初始化Saver对象,直接调用保存接口即可。如果不需要保存优化器状态、滑动平均这类非网络本身的参数,可以先按作用域过滤出网络的所有权重,减小checkpoint文件体积。
示例代码:# 过滤出你网络两个作用域下的所有权重参数 model_vars = tf.get_collection( tf.GraphKeys.GLOBAL_VARIABLES, scope="Residual" ) + tf.get_collection( tf.GraphKeys.GLOBAL_VARIABLES, scope="Output_projection" ) # 初始化Saver,传入要保存的变量列表 saver = tf.train.Saver(var_list=model_vars) # 保存权重到指定路径 save_path = saver.save(sess, "./checkpoints/res_model.ckpt")推理阶段加载权重直接输出
推理时不需要跑训练流程,只要和训练阶段完全一致地搭建一遍前向计算图(变量名、作用域名、维度参数和训练时保持一致),再通过Saver加载之前存的checkpoint,就能直接计算输出节点out的结果。
示例代码:# 1. 按训练时完全相同的参数配置搭建计算图 model = MODEL() # 配置和训练时一致的hdim、ddim、num_layers等参数 input_x = tf.placeholder(tf.float64, shape=(None, input_dim)) # 调用网络前向传播逻辑,若觉得私有方法调用麻烦,可把类里的__NN重命名为NN去掉私有标识 out = model._MODEL__NN(input_x) # 2. 初始化Saver,变量列表和保存时保持一致 saver = tf.train.Saver() with tf.Session() as sess: # 注意:不要执行tf.global_variables_initializer(),否则会覆盖加载的权重 saver.restore(sess, "./checkpoints/res_model.ckpt") # 3. 传入输入数据直接推理,无需训练 infer_result = sess.run(out, feed_dict={input_x: your_test_data})
辅助功能实现
如果需要实现类似Kerasmodel.summary()的参数查看效果,不需要给类逐层加属性存权重,直接通过变量名从计算图里取对应张量即可:
# 例如获取第0层残差层的权重 w_layer0 = tf.get_default_graph().get_tensor_by_name("Residual/Layer_0/weights:0") b_layer0 = tf.get_default_graph().get_tensor_by_name("Residual/Layer_0/biases:0") # 在Session中run对应张量,就能拿到numpy格式的权重值
遍历所有层的作用域打印参数名、shape,就能得到完整的网络参数统计信息。
注意事项
- 推理阶段搭建计算图时,
hdim/ddim/残差层数量、变量作用域名、参数名必须和训练时完全一致,否则会报变量匹配错误。你当前代码里命名都是固定写死的,只要配置参数和训练时对齐就不会出问题。 - 加载权重前不要执行
tf.global_variables_initializer(),否则会用随机初始化值覆盖训练好的权重。
内容的提问来源于stack exchange,提问作者Carlos
相关产品推荐
相关产品推荐

