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

如何从TensorFlow自定义SWATS优化器中提取lg_err参数?

如何正确从SWATS自定义优化器中提取lg_err参数?

作为刚接触Python和深度学习的开发者,你遇到的问题核心在于混淆了TensorFlow计算图构建逻辑与Python运行时代码的差异,同时文件写入方式也存在错误。下面是详细的问题分析和解决方案:

问题分析

你当前的代码存在两个关键问题:

  1. 文件写入逻辑错误:每次循环处理参数时都用'w'模式打开文件,会直接覆盖之前的内容,导致最终CSV仅保留最后一个参数的lg_err值,且丢失所有历史迭代数据。
  2. 计算图与Python代码混淆:get_updates方法是用来构建TensorFlow计算图的,里面的Python代码只会在模型初始化时执行一次,而非每次训练迭代都运行。因此你现在的Trial.append(lg_err)和CSV写入逻辑根本不会在每次迭代时触发,自然无法收集到每轮的lg_err数据。

正确实现方案

我们需要通过TensorFlow变量跟踪lg_err,再结合Keras回调函数在训练迭代时提取数值并写入CSV,具体步骤如下:

1. 修改SWATS优化器,添加lg_err跟踪机制

from tensorflow.python.framework import ops
from tensorflow.python.keras import optimizers
from tensorflow.python.keras import backend as K
from tensorflow.python.ops import math_ops
from tensorflow.python.ops import state_ops
from tensorflow.python.ops import control_flow_ops
from tensorflow.python.ops import gen_math_ops
import csv

class SWATS(optimizers.Optimizer):
    def __init__(self,lr=0.001,lr_boost=10.0,beta_1=0.9,beta_2=0.999,epsilon=None,decay=0.,amsgrad=False,**kwargs):
        super(SWATS, self).__init__(**kwargs)
        with K.name_scope(self.__class__.__name__):
            self.iterations = K.variable(0, dtype='int64', name='iterations')
            self.lr = K.variable(lr, name='lr')
            self.beta_1 = K.variable(beta_1, name='beta_1')
            self.beta_2 = K.variable(beta_2, name='beta_2')
            self.decay = K.variable(decay, name='decay')
            # 新增列表存储每个参数的lg_err张量
            self.lg_errs = []
        if epsilon is None:
            epsilon = K.epsilon()
        self.epsilon = epsilon
        self.initial_decay = decay
        self.amsgrad = amsgrad

    def get_updates(self, loss, params):
        def m_switch(pred, tensor_a, tensor_b):
            def f_true():
                return tensor_a
            def f_false():
                return tensor_b
            return control_flow_ops.cond(pred, f_true, f_false, strict=True)
        grads = self.get_gradients(loss, params)
        self.updates = []
        lr = self.lr
        if self.initial_decay > 0:
            lr = lr * ( 1. / (1. + self.decay * math_ops.cast(self.iterations,K.dtype(self.decay))) )
        with ops.control_dependencies([state_ops.assign_add(self.iterations, 1)]):
            t = math_ops.cast(self.iterations, K.floatx())
            lr_bc = gen_math_ops.sqrt(1. - math_ops.pow(self.beta_2, t)) / (1. - math_ops.pow(self.beta_1, t))
        ms = [K.zeros(K.int_shape(p), dtype=K.dtype(p)) for p in params]
        vs = [K.zeros(K.int_shape(p), dtype=K.dtype(p)) for p in params]
        lams = [K.zeros(1, dtype=K.dtype(p)) for p in params]
        conds = [K.variable(False, dtype='bool') for p in params]
        if self.amsgrad:
            vhats = [K.zeros(K.int_shape(p), dtype=K.dtype(p)) for p in params]
        else:
            vhats = [K.zeros(1) for _ in params]
        self.weights = [self.iterations] + ms + vs + vhats + lams + conds

        # 清空之前的lg_err存储,确保每轮迭代都是最新数据
        self.lg_errs.clear()

        for p, g, m, v, vhat, lam, cond in zip(params, grads, ms, vs, vhats, lams, conds):
            beta_g = m_switch(cond, 1.0, 1.0 - self.beta_1)
            m_t = (self.beta_1 * m) + beta_g * g
            v_t = (self.beta_2 * v) + (1. - self.beta_2) * math_ops.square(g)
            if self.amsgrad:
                vhat_t = math_ops.maximum(vhat, v_t)
                p_t_ada = lr_bc * m_t / (gen_math_ops.sqrt(vhat_t) + self.epsilon)
                self.updates.append(state_ops.assign(vhat, vhat_t))
            else:
                p_t_ada = lr_bc * m_t / (gen_math_ops.sqrt(v_t) + self.epsilon)
            gamma_den = math_ops.reduce_sum(p_t_ada * g)
            gamma = math_ops.reduce_sum(gen_math_ops.square(p_t_ada)) / (math_ops.abs(gamma_den) + self.epsilon) * (gen_math_ops.sign(gamma_den) + self.epsilon)
            lam_t = (self.beta_2 * lam) + (1. - self.beta_2) * gamma
            lam_prime = lam / (1. - math_ops.pow(self.beta_2, t))
            lam_t_prime = lam_t / (1. - math_ops.pow(self.beta_2, t))
            lg_err = math_ops.abs( lam_t_prime - gamma )
            
            # 将当前参数的lg_err张量加入跟踪列表
            self.lg_errs.append(lg_err)
            
            cond_update = gen_math_ops.logical_or(gen_math_ops.logical_and(gen_math_ops.logical_and( self.iterations > 1, lg_err < 1e-4 ), lam_t > 0 ), cond )[0]
            lam_update = m_switch(cond_update, lam, lam_t)
            self.updates.append(state_ops.assign(lam, lam_update))
            self.updates.append(state_ops.assign(cond, cond_update))
            p_t_sgd = (1. - self.beta_1) * lam_prime * m_t
            self.updates.append(state_ops.assign(m, m_t))
            self.updates.append(state_ops.assign(v, v_t))
            new_p = m_switch(cond, p - lr * p_t_sgd, p - lr * p_t_ada)
            # Apply constraints.
            if getattr(p, 'constraint', None) is not None:
                new_p = p.constraint(new_p)
            self.updates.append(state_ops.assign(p, new_p))
        return self.updates

    def get_config(self):
        config = {
            'lr': float(K.get_value(self.lr)),
            'beta_1': float(K.get_value(self.beta_1)),
            'beta_2': float(K.get_value(self.beta_2)),
            'decay': float(K.get_value(self.decay)),
            'epsilon': self.epsilon,
            'amsgrad': self.amsgrad
        }
        base_config = super(SWATS, self).get_config()
        return dict(list(base_config.items()) + list(config.items()))
    
    # 新增方法:获取当前所有参数的lg_err实际数值
    def get_current_lg_errs(self):
        return [K.get_value(err) for err in self.lg_errs]

2. 编写Keras回调函数,迭代时记录数据

import tensorflow.keras as keras

class LGErrLogger(keras.callbacks.Callback):
    def __init__(self, optimizer, filename='lg_err_values.csv'):
        super().__init__()
        self.optimizer = optimizer
        self.filename = filename
        # 初始化CSV文件,写入表头(迭代编号+每个参数的lg_err列)
        with open(self.filename, 'w', newline='') as f:
            writer = csv.writer(f)
            params_count = len(self.optimizer.lg_errs)
            writer.writerow(['iteration'] + [f'param_{i}' for i in range(params_count)])
    
    def on_epoch_end(self, epoch, logs=None):
        # 获取当前迭代所有参数的lg_err数值
        lg_err_values = self.optimizer.get_current_lg_errs()
        # 以追加模式写入CSV,避免覆盖历史数据
        with open(self.filename, 'a', newline='') as f:
            writer = csv.writer(f)
            writer.writerow([epoch] + lg_err_values)

3. 训练时使用回调函数收集数据

# 初始化SWATS优化器
optimizer = SWATS(lr=0.001)
# 编译你的模型
model.compile(optimizer=optimizer, loss='your_loss_function')
# 初始化lg_err日志回调
lg_err_logger = LGErrLogger(optimizer)
# 开始训练,传入回调函数
model.fit(x_train, y_train, epochs=10, callbacks=[lg_err_logger])

关键说明

  • 跟踪lg_err张量:在优化器类中添加self.lg_errs列表,存储每个参数对应的lg_err张量,确保计算图运行时能实时更新这些值。
  • 回调函数写入CSV:利用Keras回调函数的on_epoch_end方法(若需要跟踪每个batch可改用on_batch_end),在每次迭代结束时提取lg_err的实际数值,用'a'模式追加写入CSV,避免覆盖历史数据。
  • 区分计算图与Python代码:get_updates中的代码用于构建计算图,而回调函数中的代码是在每次迭代时执行的Python运行时代码,两者分工明确,确保数据收集逻辑正确。

这样修改后,你就能得到每轮训练中所有参数的lg_err值,方便后续跨迭代绘图分析。

内容的提问来源于stack exchange,提问作者Kelvin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:36:07