You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 21:06:27