如何从TensorFlow自定义SWATS优化器中提取lg_err参数?
如何正确从SWATS自定义优化器中提取lg_err参数?
作为刚接触Python和深度学习的开发者,你遇到的问题核心在于混淆了TensorFlow计算图构建逻辑与Python运行时代码的差异,同时文件写入方式也存在错误。下面是详细的问题分析和解决方案:
问题分析
你当前的代码存在两个关键问题:
- 文件写入逻辑错误:每次循环处理参数时都用
'w'模式打开文件,会直接覆盖之前的内容,导致最终CSV仅保留最后一个参数的lg_err值,且丢失所有历史迭代数据。 - 计算图与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
相关产品推荐
相关产品推荐

