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

MellowMax算子返回+INF求助:高温度参数下TensorFlow实现优化方案

MellowMax在DQN中的实现问题解决

MellowMax是深度Q学习(Deep Q Learning)里替代Max的softmax算子,已被证实无需使用目标网络,对应论文为《MellowMax: A New Approach to Reinforcement Learning》。

为估计目标Q值,需对下一状态的Q值执行MellowMax操作,公式如下:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 05:37:35