MellowMax算子返回+INF求助:高温度参数下TensorFlow实现优化方案
MellowMax是深度Q学习(Deep Q Learning)里替代Max的softmax算子,已被证实无需使用目标网络,对应论文为《MellowMax: A New Approach to Reinforcement Learning》。
为估计目标Q值,需对下一状态的Q值执行MellowMax操作,公式如下:
其中x为Q值张量,w为温度参数。
我的TensorFlow实现代码如下:
def mellow_max(q_values): q_values = tf.cast(q_values, tf.float64) powers = tf.multiply(q_values, DEEP_MELLOW_TEMPERATURE_VALUE) summation_values = tf.math.exp(powers) summation = tf.math.reduce_sum(summation_values, axis=1) val_for_log = tf.multiply(summation,(1/NUM_ACTIONS)) numerator = tf.math.log(val_for_log) mellow_val = tf.math.divide(numerator, DEEP_MELLOW_TEMPERATURE_VALUE).numpy() return mellow_val
当温度参数w设为1000(论文中Atari Breakout测试的最优值)时,函数第三行计算tf.math.exp(powers)会返回+inf值,希望得到解决建议,比如如何在TensorFlow中计算w趋近于1000时的函数极限来避免该问题。
解决思路
1. 数值稳定的公式改写
直接计算exp(w*x)当w很大时必然会溢出,我们可以利用对数的性质对公式变形,彻底避免大数指数运算:
原公式:
$$\text{MellowMax}(x; w) = \frac{1}{w} \log\left( \frac{1}{n} \sum_{i=1}^n e^{w x_i} \right)$$
提取每个样本Q值中的最大值$x_{max} = \max(x_i)$,改写为:
$$\text{MellowMax}(x; w) = x_{max} + \frac{1}{w} \log\left( \frac{1}{n} \sum_{i=1}^n e^{w (x_i - x_{max})} \right)$$
这样$x_i - x_{max} \leq 0$,$w*(x_i - x_{max})$不会出现正的大数,exp计算结果始终≤1,完全避免溢出问题。
2. 修改后的TensorFlow实现
基于上述变形,调整后的代码如下:
def mellow_max(q_values): q_values = tf.cast(q_values, tf.float64) w = DEEP_MELLOW_TEMPERATURE_VALUE n = NUM_ACTIONS # 提取每个样本的Q值最大值,保留维度方便广播 x_max = tf.reduce_max(q_values, axis=1, keepdims=True) # 计算偏移后的指数项,避免溢出 shifted_powers = w * (q_values - x_max) summation_values = tf.math.exp(shifted_powers) summation = tf.reduce_sum(summation_values, axis=1) avg = summation / n log_term = tf.math.log(avg) # 合并结果并转换为numpy数组返回 mellow_val = (x_max[:, 0] + log_term / w).numpy() return mellow_val
3. 极端大w的极限处理
当w趋近于无穷大时,MellowMax的极限就是普通的Max操作。可以设置一个阈值,当w超过该阈值时直接返回Q值的最大值,跳过指数计算,进一步提升效率:
def mellow_max(q_values): q_values = tf.cast(q_values, tf.float64) w = DEEP_MELLOW_TEMPERATURE_VALUE n = NUM_ACTIONS # 阈值可根据数值精度调整,比如1e3或1e4 if w > 1e3: return tf.reduce_max(q_values, axis=1).numpy() # 数值稳定的计算逻辑 x_max = tf.reduce_max(q_values, axis=1, keepdims=True) shifted_powers = w * (q_values - x_max) summation_values = tf.math.exp(shifted_powers) summation = tf.reduce_sum(summation_values, axis=1) avg = summation / n log_term = tf.math.log(avg) mellow_val = (x_max[:, 0] + log_term / w).numpy() return mellow_val
内容的提问来源于stack exchange,提问作者Bryan Carty

