<?xml version="1.0" encoding="utf-8" standalone="yes"?><rss version="2.0" xmlns:atom="http://www.w3.org/2005/Atom"><channel><title>Math on Jiang Yi(姜祎)'s Homepage</title><link>https://jiangyigithub.github.io/ai.github.io/categories/math/</link><description>Recent content in Math on Jiang Yi(姜祎)'s Homepage</description><generator>Hugo -- gohugo.io</generator><language>en</language><lastBuildDate>Wed, 08 Apr 2026 21:44:44 +0800</lastBuildDate><atom:link href="https://jiangyigithub.github.io/ai.github.io/categories/math/index.xml" rel="self" type="application/rss+xml"/><item><title>(RL series 3) Policy evaluation</title><link>https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/</link><pubDate>Wed, 18 Mar 2026 17:43:57 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/</guid><description>&lt;p&gt;这一节我们讨论如何评估 (evaluate) 一个 policy. 要评估一个 policy $\pi$, 我们需要求出对应的 value function $V^\pi$.&lt;/p&gt;
&lt;h2 id="value-function"&gt;&lt;a href="#value-function" class="header-anchor"&gt;&lt;/a&gt;Value Function
&lt;/h2&gt;&lt;h3 id="monte-carlo"&gt;&lt;a href="#monte-carlo" class="header-anchor"&gt;&lt;/a&gt;Monte Carlo
&lt;/h3&gt;&lt;p&gt;注意到&lt;/p&gt;
$$
V^\pi(s) = \mathbb{E}^\pi\left[\sum_{t=0}^{T-1}\gamma^tr_t\mid s_0=s\right]
$$&lt;p&gt;我们可以利用 Monte Carlo (MC) 方法来进行估计，即从 $s_0=s$ 出发，独立随机采样 $N$ 条轨迹&lt;/p&gt;
$$
\{s, a_0^{(i)},r_0^{(i)},\dots,a_{T^{(i)}-1}^{(i)},r_{T^{(i)}-1}^{(i)}, s^{(i)}_{T^{(i)}}\}, i=1,\dots,N
$$&lt;p&gt;然后我们使用样本平均来近似期望&lt;/p&gt;
$$
V^\pi(s) = \mathbb{E}^\pi\left[\sum_{t=0}^{T-1}\gamma^tr_t\mid s_0=s\right]\approx \frac{1}{N}\sum_{i=1}^N\sum_{t=0}^{T^{(i)}-1}\gamma^tr_t^{(i)}
$$&lt;p&gt;如果 $\mathcal{S}$ 比较小，我们可以遍历 $s\in\mathcal{S}$ 来估计 $V^\pi$.&lt;/p&gt;
&lt;p&gt;当 $\mathcal{S}$ 非常大或者连续时，我们需要使用神经网络 $V_\phi$ 来近似 $V^\pi$,&lt;/p&gt;
$$
\min_{\phi} \mathcal{L}(\phi) =\mathbb{E}_{s\sim p_0}\left[\frac12\left(V_\phi(s)-V^\pi(s)\right)^2\right]
$$&lt;p&gt;对上面的目标函数求导得到&lt;/p&gt;
$$
\begin{align}
\nabla_\phi \mathcal{L}(\phi) &amp;= \mathbb{E}_{s\sim p_0}\left[\left(V_\phi(s)-V^\pi(s)\right)\nabla_\phi V_\phi(s)\right]\\
&amp;= \mathbb{E}_{s\sim p_0}\left[\left(V_\phi(s)- \mathbb{E}^\pi\left[\sum_{t=0}^{T-1}\gamma^tr_t\mid s_0=s\right]\right)\nabla_\phi V_\phi(s)\right]\\
&amp;=\mathbb{E}_{s\sim p_0}\left[\mathbb{E}^\pi\left[\left(V_\phi(s)- \sum_{t=0}^{T-1}\gamma^tr_t\right)\nabla_\phi V_\phi(s)\mid s_0=s\right]\right]\\
&amp;= \mathbb{E}_{s\sim p_0}^\pi\left[\left(V_\phi(s)- \sum_{t=0}^{T-1}\gamma^tr_t\right)\nabla_\phi V_\phi(s)\mid s_0=s\right]\\
\end{align}
$$&lt;p&gt;这样，我们可以结合 SGD 与 MC 来近似 $V^\pi$&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/value-function-MC.png"
width="653"
height="255"
loading="lazy"
alt="Value function approximation with MC"
class="gallery-image"
data-flex-grow="256"
data-flex-basis="614px"
&gt;&lt;/p&gt;
&lt;h3 id="temporal-difference"&gt;&lt;a href="#temporal-difference" class="header-anchor"&gt;&lt;/a&gt;Temporal Difference
&lt;/h3&gt;&lt;p&gt;Temporal difference (TD) learning 通过相邻两步的转换关系来近似函数。注意到，&lt;/p&gt;
$$
V^\pi(s) =\mathbb{E}^\pi[r_0 + \gamma V^\pi(s_1)\mid s_0=s]
$$&lt;p&gt;与前面的方法类似，我们先采样然后使用 MC 来进行估计&lt;/p&gt;
$$
V^\pi(s) =\mathbb{E}^\pi[r_0 + \gamma V^\pi(s_1)\mid s_0=s]\approx \frac{1}N\sum_{i=1}^N\left(r_0^{(i)} + \gamma V^\pi(s_1^{(i)})\right)
$$&lt;p&gt;同样的，当 $\mathcal{S}$ 比较大或者连续时，我们构造损失函数&lt;/p&gt;
$$
\min_{\phi} \mathcal{L}(\phi) =\mathbb{E}_{s\sim p_0}\left[\frac12\left(V_\phi(s)-V^\pi(s)\right)^2\right]
$$&lt;p&gt;对应的梯度为&lt;/p&gt;
$$
\nabla_\phi \mathcal{L}(\phi) =\mathbb{E}_{s\sim p_0}^\pi\left[\left(V_\phi(s)- r_0-\gamma V^\pi(s_1)\right)\nabla_\phi V_\phi(s)\mid s_0=s\right]
$$&lt;p&gt;但是，这里存在的问题在于，我们使用了 $V^\pi(s_1)$, 而这个值是未知的，因此，一个做法是使用当前的 value function $V_\phi$ 来进行代替，即&lt;/p&gt;
$$
\nabla_\phi \mathcal{L}(\phi) \approx \mathbb{E}_{s\sim p_0}^\pi\left[\left(V_\phi(s)- r_0-\gamma V_\phi(s_1)\right)\nabla_\phi V_\phi(s)\mid s_0=s\right]
$$&lt;p&gt;这样，我们结合 TD 和 GSD 的算法就变成了&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/value-function-1-TD-wrong.png"
width="656"
height="171"
loading="lazy"
alt="Value function approximation with TD (wrong)"
class="gallery-image"
data-flex-grow="383"
data-flex-basis="920px"
&gt;&lt;/p&gt;
&lt;p&gt;接下来，我们需要关注如何计算 $g$, 直接通过 $g=\left(V_\phi(s)- r_0-\gamma V_\phi(s_1)\right)\nabla_\phi V_\phi(s)$ 来进行计算，这在数学上是没有问题的，但是从自动微分的角度不对，比如 Pytorch 对应的实现应该是&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt;1
&lt;/span&gt;&lt;span class="lnt"&gt;2
&lt;/span&gt;&lt;span class="lnt"&gt;3
&lt;/span&gt;&lt;span class="lnt"&gt;4
&lt;/span&gt;&lt;span class="lnt"&gt;5
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-python" data-lang="python"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;pred&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;V_phi&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;s0&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;target&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;r&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="n"&gt;gamma&lt;/span&gt; &lt;span class="o"&gt;*&lt;/span&gt; &lt;span class="n"&gt;V_phi&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;s1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;td_error&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;pred&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="n"&gt;target&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;loss&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;0.5&lt;/span&gt; &lt;span class="o"&gt;*&lt;/span&gt; &lt;span class="n"&gt;td_error&lt;/span&gt; &lt;span class="o"&gt;**&lt;/span&gt; &lt;span class="mi"&gt;2&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;loss&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;backward&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;p&gt;可以看到，我们实际上求的梯度是&lt;/p&gt;
$$
\nabla_\phi\left[\frac12\left(V_\phi(s)-(r+\gamma V_\phi(s'))\right)^2\right]=(V_\phi(s)-(r+\gamma V_\phi(s'))(\nabla_\phi V_\phi(s)-\nabla_\phi V_\phi(s'))
$$&lt;p&gt;这显然与 $g$ 不相等。&lt;/p&gt;
&lt;p&gt;为了解决这个问题，我们需要使用 &lt;a class="link" href="stop-gradient.md" &gt;stop-gradient&lt;/a&gt; 的技巧来避免 $V_\phi(s')$ 参与反向传播，这个时候，我们的目标函数就变成了&lt;/p&gt;
$$
\frac12\left(V_\phi(s)-(r+\gamma\ \mathrm{sg}[V_\phi(s')])\right)^2
$$&lt;p&gt;其中 $\mathrm{sg}[\cdot]$ 是 stop-gradient operator, 满足&lt;/p&gt;
$$
\mathrm{sg}[x] = \begin{cases}
x &amp;\text{forward pass}\\
0 &amp;\text{backward pass}
\end{cases}
$$&lt;p&gt;这样其梯度就是&lt;/p&gt;
$$
\nabla_\phi\left[\frac12\left(V_\phi(s)-(r+\gamma\ \mathrm{sg}[V_\phi(s')])\right)^2\right]=(V_\phi(s)-(r+\gamma V_\phi(s'))\nabla_\phi V_\phi(s)=g
$$&lt;p&gt;对应的 python 代码为&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt;1
&lt;/span&gt;&lt;span class="lnt"&gt;2
&lt;/span&gt;&lt;span class="lnt"&gt;3
&lt;/span&gt;&lt;span class="lnt"&gt;4
&lt;/span&gt;&lt;span class="lnt"&gt;5
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-python" data-lang="python"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;pred&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;V_phi&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;s0&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;target&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;r&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="n"&gt;gamma&lt;/span&gt; &lt;span class="o"&gt;*&lt;/span&gt; &lt;span class="n"&gt;V_phi&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;s1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;td_error&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;pred&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="n"&gt;target&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;detach&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt; &lt;span class="c1"&gt;# sop gradient operator&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;loss&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;0.5&lt;/span&gt; &lt;span class="o"&gt;*&lt;/span&gt; &lt;span class="n"&gt;td_error&lt;/span&gt; &lt;span class="o"&gt;**&lt;/span&gt; &lt;span class="mi"&gt;2&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;loss&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;backward&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;p&gt;最终，我们的算法实现为&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/value-function-1-TD.png"
width="650"
height="171"
loading="lazy"
alt="Value function approximation with TD"
class="gallery-image"
data-flex-grow="380"
data-flex-basis="912px"
&gt;&lt;/p&gt;
&lt;p&gt;上述这种做法其实是 semi-gradient methods, 这类方法在参数更新时，只计算部分梯度的方法。虽然这种方法牺牲了严格的梯度下降性质，但是其能够保证更好的表现。&lt;/p&gt;
&lt;h3 id="k-step-td"&gt;&lt;a href="#k-step-td" class="header-anchor"&gt;&lt;/a&gt;K-step TD
&lt;/h3&gt;&lt;p&gt;我们可以进一步推广 TD 到多步的场景，注意到&lt;/p&gt;
$$
\begin{align}
V^\pi(s) &amp;=\mathbb{E}^\pi[r_0 + \gamma V^\pi(s_1)\mid s_0=s]\\
&amp;=\mathbb{E}^\pi[r_0 + \gamma \mathbb{E}^\pi[r_1 + \gamma V^\pi(s_1)\mid s_1]\mid s_0=s]\\
&amp;= \mathbb{E}^\pi[\mathbb{E}^\pi[r_0 + \gamma r_1 + \gamma^2 V^\pi(s_2) \mid s_1] \mid s_0=s]
\end{align}
$$&lt;p&gt;由重期望定理（Law of Total Expectation）：&lt;/p&gt;
$$
\mathbb{E}^\pi[\mathbb{E}^\pi[X \mid s_1] \mid s_0=s] = \mathbb{E}^\pi[X \mid s_0=s]
$$&lt;p&gt;因此：&lt;/p&gt;
$$
V^\pi(s) = \mathbb{E}^\pi[r_0 + \gamma r_1 + \gamma^2 V^\pi(s_2) \mid s_0=s]
$$&lt;p&gt;重复这个过程 k 次,得到 k 步展开：&lt;/p&gt;
$$
\begin{align} V^\pi(s) &amp;= \mathbb{E}^\pi\left[\sum_{i=0}^{k-1} \gamma^i r_i + \gamma^k V^\pi(s_k) \mid s_0=s\right] \end{align}
$$&lt;p&gt;现在，我们可以基于 k-step transition 构建目标函数&lt;/p&gt;
$$
\min_{\phi} \mathcal{L}(\phi) =\mathbb{E}_{s\sim p_0}\left[\frac12\left(V_\phi(s)-\left(\sum_{i=0}^{k-1} \gamma^i r_i + \gamma^k V^\pi(s_k)\right)\right)^2\right]
$$&lt;p&gt;使用前面分析的方法，我们就可以写出类似的算法&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/value-function-k-TD.png"
width="646"
height="230"
loading="lazy"
alt="value function approximation with k-step TD"
class="gallery-image"
data-flex-grow="280"
data-flex-basis="674px"
&gt;&lt;/p&gt;
&lt;p&gt;一般来说，我们会令 $k=5$.&lt;/p&gt;
&lt;h3 id="mc-vs-td"&gt;&lt;a href="#mc-vs-td" class="header-anchor"&gt;&lt;/a&gt;MC v.s. TD
&lt;/h3&gt;&lt;p&gt;TD 本质上一种 bootstrapping, 当需要使用 $V^\pi$ 时，我们使用 $V_\phi$ 来代替。两者对比如下&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;MC 需要完整的一次采样，但是 TD 可以通过 bootstrap 提高采样效率。&lt;/li&gt;
&lt;li&gt;当 $V_\phi$ 与 $V^\pi$ 差距较大时，boostraping 会导致训练不稳定，而 MC 则相对稳定&lt;/li&gt;
&lt;/ol&gt;
&lt;h2 id="q-function"&gt;&lt;a href="#q-function" class="header-anchor"&gt;&lt;/a&gt;Q-function
&lt;/h2&gt;&lt;p&gt;对于 Q-function, 我们也可以设计类似的算法。&lt;/p&gt;
&lt;p&gt;对于 MC, 我们有&lt;/p&gt;
$$
Q^{\pi}(s,a) = \mathbb{E}^{\pi}\left[\sum_{t=0}^{T-1}\gamma^tr_t\mid s_0=s, a_0=a\right]\approx \frac1N\sum_{i=1}^N\sum_{t=0}^{T^{(i)}-1}\gamma^tr_t^{(i)}
$$&lt;p&gt;当使用模型来近似时，我们的目标函数为&lt;/p&gt;
$$
\min_{\phi} \mathcal{L}(\phi) =\mathbb{E}_{s\sim p_0, a\sim\pi(\cdot\mid s)}\left[\frac12\left(Q_\phi(s,a)-Q^\pi(s,a)\right)^2\right]
$$&lt;p&gt;对应的梯度&lt;/p&gt;
$$
\nabla_\phi \mathcal{L}(\phi)= \mathbb{E}_{s\sim p_0}^\pi\left[\left(Q_\phi(s_0,a_0)-\sum_{t=0}^{T-1}\gamma^tr_t\right)\nabla_\phi Q_\phi(s_0,a_0)\right]
$$&lt;p&gt;下面是对应的算法&lt;/p&gt;
&lt;p&gt;MC&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/q-function-1-TD.png"
width="647"
height="178"
loading="lazy"
alt="Q function approximation with TD"
class="gallery-image"
data-flex-grow="363"
data-flex-basis="872px"
&gt;&lt;/p&gt;
&lt;p&gt;1-step Q-function TD learning&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/q-function-1-TD.png"
width="647"
height="178"
loading="lazy"
alt="Q function approximation with TD"
class="gallery-image"
data-flex-grow="363"
data-flex-basis="872px"
&gt;&lt;/p&gt;
&lt;p&gt;k-step Q-function TD learning&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-3-policy-evaluation/q-function-k-TD.png"
width="651"
height="231"
loading="lazy"
alt="Q function approximation with k-step TD"
class="gallery-image"
data-flex-grow="281"
data-flex-basis="676px"
&gt;&lt;/p&gt;
&lt;h2 id="conclusion"&gt;&lt;a href="#conclusion" class="header-anchor"&gt;&lt;/a&gt;Conclusion
&lt;/h2&gt;&lt;p&gt;在本节中，我们介绍了给定 policy $\pi$, 我们如何求解对应的 value function 和 Q-function.&lt;/p&gt;</description></item><item><title>(RL series 2) Bellman Equation</title><link>https://jiangyigithub.github.io/ai.github.io/p/rl-series-2-bellman-equation/</link><pubDate>Wed, 18 Mar 2026 17:39:45 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/rl-series-2-bellman-equation/</guid><description>&lt;h2 id="bellman-equations"&gt;&lt;a href="#bellman-equations" class="header-anchor"&gt;&lt;/a&gt;Bellman Equations
&lt;/h2&gt;&lt;p&gt;本节中，我们将要定义 value function, Q-function 以及这两个函数与 policy 之间的关系，这是我们介绍不同 RL 算法的基础。这一节需要使用到 &lt;a class="link" href="https://maosong.website/p/fix-point-theorem/" target="_blank" rel="noopener"
&gt;不动点定理&lt;/a&gt;.&lt;/p&gt;
&lt;p&gt;我们定义 Value function 如下：&lt;/p&gt;
$$
V^{\pi}(s) = \mathbb{E}^{\pi}[G_0\mid s_0=s]
$$&lt;p&gt;value function 的具体含义为：&lt;em&gt;agent 从当前状态 $s$ 出发，一直遵循当前策略 $\pi$， 最后获取到的 expected discounted return&lt;/em&gt;.&lt;/p&gt;
&lt;p&gt;state-action value function, 或者 Q-value function 定义如下：&lt;/p&gt;
$$
Q^{\pi}(s,a) = \mathbb{E}^{\pi}[G_0\mid s_0=s, a_0=a]
$$&lt;p&gt;其具体含义为：&lt;em&gt;agent 从当前状态 $s$ 出发，执行 action $a$, 再遵循当前策略 $\pi$, 最后获取到的 expected discounted return&lt;/em&gt;.&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;注意
这里为了方便，我们令 $V^{\pi}(\langle term\rangle)=0$, $Q^{\pi}(\langle term\rangle,a)=0,\forall a\in\mathcal{A}$.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;由全概率公式，我们有&lt;/p&gt;
$$
\mathbb{E}[G_0\mid s_0=s] = \sum_{a\in\mathcal{A}}\mathbb{E}\left[G_0\mid s_t=s, a_0=a\right]\pi(a\mid s_0=s)
$$&lt;p&gt;即&lt;/p&gt;
$$
\boxed{V^{\pi}(s) = \mathbb{E}_{a_0\sim \pi(\cdot\mid s_0)}[Q^{\pi}(s,a)]}
$$&lt;p&gt;一般来说，我们会使用 value function 和 Q-value function 的如下递推性质：&lt;/p&gt;
$$
\boxed{
\begin{aligned}V^\pi(s) &amp;= \mathbb{E}_{a_0\sim \pi(\cdot\mid s),\ (r_0,s_0)\sim p(\cdot,\cdot\mid s,a_0)}[r_0 + \gamma V^\pi(s_1)\mid s_0=s]\\
Q^\pi(s,a) &amp;= \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a),\ a'\sim \pi(\cdot\mid s')}[r + \gamma Q^\pi(s',a')\mid s,a]
\end{aligned}}
$$&lt;p&gt;证明需要利用 Markov property:&lt;/p&gt;
$$
\begin{aligned}
V^{\pi}(s) &amp;= \mathbb{E}^{\pi}\left[r_0 + \gamma\sum_{t=1}^{T-1}\gamma^{t-1}r_t\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi}[r_0+\gamma G_1\mid s_0=s]\\
&amp;= \mathbb{E}_{a_0\sim \pi(\cdot\mid s,\ (r_0,s_0)\sim p(\cdot,\cdot\mid s,a_0))}\left[\mathbb{E}^\pi[r_0 + \gamma G_1\mid s_0,a_0,r_0,s_1]\mid s_0=s\right]\\
&amp;= \mathbb{E}_{a_0\sim \pi(\cdot\mid s,\ (r_0,s_0)\sim p(\cdot,\cdot\mid s,a_0))}\left[r_0 + \gamma \mathbb{E}^\pi[G_1\mid s_1]\mid s_0=s\right]\\
&amp;= \mathbb{E}_{a_0\sim \pi(\cdot\mid s,\ (r_0,s_0)\sim p(\cdot,\cdot\mid s,a_0))}\left[r_0 + \gamma V^\pi[s_1)\mid s_0=s\right]
\end{aligned}
$$&lt;p&gt;对于 $Q^\pi(s,a)$ 的推导同理。&lt;/p&gt;
&lt;p&gt;我们接下来介绍 RL 的核心：Bellman equation Theorem, 其阐述了策略 $\pi$ 与 value function 所满足的充分必要条件。&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Theorem&lt;/strong&gt;&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;令 $\pi$ 为一个策略，假设 $\gamma\in(0,1)$, $|\mathcal{S}|&lt;\infty$ 以及 $|r|\leq R&lt;\infty, \mathrm{a.s.}$. 那么 $\pi$ 对应的 value function $V^\pi:\mathcal{S}^+\to\mathbb{R}$ 存在，且满足 Bellman equation:&lt;/p&gt;
&lt;/blockquote&gt;
$$
\boxed{
V^\pi(s) = \mathbb{E}_{a\sim \pi(\cdot\mid s),\ (r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V(s')\mid s\right]
}
$$&lt;blockquote&gt;
&lt;p&gt;反之，如果存在函数 $V:\mathcal{S}^+\to\mathbb{R}$ 满足 Bellman equation, 则 $V=V^\pi$.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;由 $V^\pi$ 定义，我们有&lt;/p&gt;
$$
\left|\sum_{t=0}^{\infty}\gamma^tr_t\right|\leq \sum_{t=0}^\infty \gamma^t R=\frac{R}{1-\gamma}&lt;\infty
$$&lt;p&gt;因此，$\sum_{t=0}^{\infty}\gamma^tr_t$ 是一个绝对收敛的序列，从而其期望存在且有界。&lt;/p&gt;
&lt;p&gt;对于 $V^\pi(s)=\mathbb{E}^{\pi}[G_0\mid s_0=s]$ , 我们有&lt;/p&gt;
$$
\begin{aligned}
V^\pi(s)&amp;=\mathbb{E}^{\pi}\left[G_0\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi}\left[r_0+\gamma G_1\mid s_0=s\right]\\
&amp;=\mathbb{E}_{a\sim\pi(\cdot\mid s), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[\mathbb{E}^\pi\left[r_0+\gamma G_1\mid s_0,a_0,r_0,s'\right]\right]\\
&amp;= \mathbb{E}_{a\sim\pi(\cdot\mid s), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r_0+\gamma\mathbb{E}^\pi\left[ G_1\mid s_0,a_0,r_0,s'\right]\right]\\
&amp;= \mathbb{E}_{a\sim\pi(\cdot\mid s), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r_0+\gamma\mathbb{E}^\pi\left[ G_1\mid s_1\right]\right]\\
&amp;=\mathbb{E}_{a\sim\pi(\cdot\mid s'), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r_0+V^{\pi}(s')\right]
\end{aligned}
$$&lt;p&gt;其中第 5 个等式为 Markov property.&lt;/p&gt;
&lt;p&gt;最后，我们证明方程解的唯一性，假设存在函数 $V:\mathcal{S}^+\to\mathbb{R}$ 满足 Bellman equation. 我们定义 Bellman 算子 $\mathcal{T}^\pi$ 为&lt;/p&gt;
$$
(\mathcal{T}^\pi V)(s) := \mathbb{E}_{a\sim \pi(\cdot\mid s),\ (r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V(s')\mid s\right]
$$&lt;p&gt;我们证明该算子是一个 contraction mapping, 考虑范数 $\|V\|_\infty = \max_{s\in\mathcal{S}} |V(s)|$, 我们有&lt;/p&gt;
$$
\begin{aligned}
\left|(\mathcal{T}^\pi V_1)(s)-(\mathcal{T}^\pi V_2)(s)\right| &amp;= \left|\mathbb{E}_{a\sim\pi(\cdot\mid s'), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[\gamma V_1^{\pi}(s')-\gamma V_2^{\pi}(s')\right]\right|\\
&amp;\leq \gamma \mathbb{E}_{a\sim\pi(\cdot\mid s'), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left|V_1^{\pi}(s')-V_2^{\pi}(s')\right|\\
&amp;\leq \gamma \max_{s''\in\mathcal{S}}\left|V_1(s'')-V_2(s'')\right|\\
&amp;= \gamma \left\|V_1-V_2\right\|_\infty
\end{aligned}
$$&lt;p&gt;上式对于任意 $s\in\mathcal{S}$ 都成立，因此&lt;/p&gt;
$$
\left\|\mathcal{T}^\pi V_1-\mathcal{T}^\pi V_2\right\|_\infty \leq \gamma \left\|V_1-V_2\right\|_\infty
$$&lt;p&gt;由于 $\gamma &lt; 1$, 从而 $\mathcal{T}^\pi$ 是一个 contraction mapping.&lt;/p&gt;
&lt;p&gt;根据不动点定理，已知 $V^\pi$ 是一个不动点，而不动点唯一，因此我们有 $V=V^\pi$. $\blacksquare$&lt;/p&gt;
&lt;p&gt;同理，我们可以推导出关于 Q-function 的 Bellman equation:&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Theorem&lt;/strong&gt;&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;令 $\pi$ 为一个策略，假设 $\gamma\in(0,1)$, $|\mathcal{S}|&lt;\infty$ $|\mathcal{A}|&lt;\infty$ 以及 $|r|\leq R&lt;\infty, \mathrm{a.s.}$. 那么 $\pi$ 对应的 Q-function $Q^\pi:\mathcal{S}^+\times\mathcal{A}\to\mathbb{R}$ 存在，且满足 Bellman equation:&lt;/p&gt;
&lt;/blockquote&gt;
$$
\boxed{
Q^\pi(s, a) = \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a),\ a'\sim \pi(\cdot\mid s')}\left[r+\gamma Q(s', a')\mid s, a\right]
}
$$&lt;blockquote&gt;
&lt;p&gt;反之，如果存在函数 $Q:\mathcal{S}^+\times\mathcal{A}\to\mathbb{R}$ 满足 Bellman equation, 则 $Q=Q^\pi$.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;证明与 value function 的证明基本类似，我们这里略过。&lt;/p&gt;
&lt;h2 id="bellman-optimal-equation"&gt;&lt;a href="#bellman-optimal-equation" class="header-anchor"&gt;&lt;/a&gt;Bellman Optimal Equation
&lt;/h2&gt;&lt;p&gt;前面介绍了针对一般策略的 Bellman equation, 特别地，对于我们的目标最优策略，我们也可以而推导出对应的 Bellman equation.&lt;/p&gt;
&lt;p&gt;首先我们定义最优策略如下&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Definition&lt;/strong&gt;&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;如果策略 $\pi^*$ 满足&lt;/p&gt;
&lt;/blockquote&gt;
$$
V^{\pi^*}(s)\geq V^{\pi}(s), \forall s\in\mathcal{S}^+, \forall \pi
$$&lt;blockquote&gt;
&lt;p&gt;则我们称策略 $\pi^*$ 是 &lt;strong&gt;optimal policy&lt;/strong&gt;. 对应的 $V^{\pi^*}$ 和 $Q^{\pi^*}$ 分别称之为 &lt;strong&gt;optimal value function&lt;/strong&gt; 以及 &lt;strong&gt;optimal Q-function&lt;/strong&gt;, 我们简记为 $V^*=V^{\pi^*}$, $Q^*=Q^{\pi^*}$.
注意 optimal policy 与状态 $s$ 无关，并且 optimal policy 是不唯一的，但是所有的 optimal policy 对应的 optimal value function 是同一个。
$V^*$ 与 $Q^*$ 之间存在如下关系&lt;/p&gt;
&lt;/blockquote&gt;
$$
V^*(s) = \max_{a\sim\pi^*} Q^*(s, a)
$$&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;我们先证明左边小于右边，再证明右边小于左边。&lt;/p&gt;
$$
V^*(s) =\mathbb{E}_{a_0\sim \pi^*(\cdot\mid s_0)}[Q^*(s,a)]\leq \mathbb{E}_{a_0\sim \pi^*(\cdot\mid s_0)}[\max_{a\sim\pi^*}Q^*(s,a)] = \max_{a\sim\pi^*} Q^*(s, a)
$$&lt;p&gt;其次，令 $a^*\in\arg\max_{a}Q^*(s,a)$, 令策略 $\pi'$ 为确定性策略 $\pi'(a^*\mid s)=1$, 则&lt;/p&gt;
$$
\max_aQ^*(s,a^*) = Q^{\pi'}(s,a^*)=V^{\pi'}(s)\leq V^*(s)
$$&lt;p&gt;这样，我们就证明了上面的等式。$\blacksquare$&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Theorem&lt;/strong&gt;&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;假设 $\gamma\in(0,1)$, $|\mathcal{S}|&lt;\infty$ $|\mathcal{A}|&lt;\infty$ 以及 $|r|\leq R&lt;\infty, \mathrm{a.s.}$. 那么 optimal value function $V^\pi:\mathcal{S}^+\to\mathbb{R}$ 存在，且满足 &lt;strong&gt;Bellman optimality equation&lt;/strong&gt;:&lt;/p&gt;
&lt;/blockquote&gt;
$$
\boxed{
V^*(s) = \max_{a\in\mathcal{A}}\mathbb{E}_{\ (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r+\gamma V^*(s')\mid s, a\right]
}
$$&lt;blockquote&gt;
&lt;p&gt;反之，如果存在函数 $V:\mathcal{S}^+\to\mathbb{R}$ 满足 Bellman optimality equation, 则 $V=V^*$, 最终&lt;/p&gt;
&lt;/blockquote&gt;
$$
\pi^*(s) = \arg\max_{a\in\mathcal{A}}\mathbb{E}_{\ (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r+\gamma V^*(s')\mid s, a\right]=\arg\max_{a\in\mathcal{A}} Q^*(s,a)
$$&lt;blockquote&gt;
&lt;p&gt;是一个 optimal deterministic policy.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;我们首先证明存在性，我们定义&lt;/p&gt;
$$
(\mathcal{T}^* V)(s) := \max_{a\in\mathcal{A}}\mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V(s')\mid s, a\right]
$$&lt;p&gt;我们证明 $\mathcal{T}^*V$ 是一个 contraction mapping.&lt;/p&gt;
$$
\begin{aligned}
\left|(\mathcal{T}^* V_1)(s)-(\mathcal{T}^* V_2)(s)\right| &amp;= \left|\max_{a\in\mathcal{A}}\mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V_1(s')\mid s, a\right]- \max_{a\in\mathcal{A}}\mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V_2(s')\mid s, a\right]\right|\\
&amp;\leq \max_{a\in\mathcal{A}}\left|\mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V_1(s')\mid s, a\right]- \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V_2(s')\mid s, a\right]\right|\\
&amp;= \max_{a\in\mathcal{A}}\left|\mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[\gamma V_1(s')-V_2(s')\mid s, a\right]- \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V_2(s')\mid s, a\right]\right|\\
&amp;\leq \gamma \max_{a\in\mathcal{A}}\left|\mathbb{E}_{a\sim\pi(\cdot\mid s'), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[V_1^{\pi}(s')-V_2^{\pi}(s')\mid s,a\right]\right|\\
&amp;\leq \gamma \max_{a\in\mathcal{A}}\mathbb{E}_{a\sim\pi(\cdot\mid s'), (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[\left|V_1^{\pi}(s')-V_2^{\pi}(s')\right|\mid s,a\right]\\
&amp;\leq \gamma \max_{a\in\mathcal{A}}\max_{s''\in\mathcal{S}}\left|V_1(s'')-V_2(s'')\right|\\
&amp;= \gamma \left\|V_1-V_2\right\|_\infty
\end{aligned}
$$&lt;p&gt;这里第一个不等式我们使用了结论&lt;/p&gt;
$$
\left|\max_s v(s)-\max_s u(s)\right|\leq \max_s|u(s)-v(s)|
$$&lt;p&gt;上式对于任意 $s\in\mathcal{S}$ 都成立，因此&lt;/p&gt;
$$
\left\|\mathcal{T}^* V_1-\mathcal{T}^* V_2\right\|_\infty \leq \gamma \left\|V_1-V_2\right\|_\infty
$$&lt;p&gt;由于 $\gamma &lt; 1$, 从而 $\mathcal{T}^*$ 是一个 contraction mapping.&lt;/p&gt;
&lt;p&gt;根据不动点定理，$\mathcal{T}^*$ 存在一个不动点，我们将其记为 $V^*$ .&lt;/p&gt;
&lt;p&gt;接下来，我们定义 $\pi^*$ 为&lt;/p&gt;
$$
\pi^*(s) = \arg\max_{a\in\mathcal{A}}\mathbb{E}_{\ (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r+\gamma V^*(s')\mid s, a\right]
$$&lt;p&gt;此时，我们有&lt;/p&gt;
$$
\begin{aligned}
V^*(s) &amp;= \max_{a\in\mathcal{A}}\mathbb{E}_{\ (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r+\gamma V^*(s')\mid s, a\right]\\
&amp;= \mathbb{E}_{a\sim\pi^*(s),\ (r,s')\sim p(\cdot,\cdot\mid s,a_0)}\left[r+\gamma V^*(s')\mid s\right]\\
&amp;= \mathbb{E}^{\pi^*}\left[r_0+\gamma V^*(s_1)\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi^*}\left[r_0+\gamma \mathbb{E}^{\pi^*}\left[r_1+\gamma V^*(s_2)\mid s_0=s_1\right]\\\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi^*}\left[r_0+\gamma \mathbb{E}^{\pi^*}\left[r_1+\gamma V^*(s_2)\mid s_0=s_1, r_0\right]\\\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi^*}\left[\mathbb{E}^{\pi^*}\left[r_0+\gamma (r_1+\gamma V^*(s_2))\mid s_0=s_1, r_0\right]\\\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi^*}\left[r_0+\gamma r_1+\gamma^2 V^*(s_2)\mid s_0=s\right]\\
&amp;= \mathbb{E}^{\pi^*}\left[r_0+\gamma r_1+\gamma^2r_2+\cdots\mid s_0=s\right]\\
&amp;= V^{\pi^*}(s)
\end{aligned}
$$&lt;p&gt;即 $V^*=V^{\pi^*}$.&lt;/p&gt;
&lt;p&gt;现在我们证明 $\pi^*$ 是最优策略，令 $\pi$ 为任意一个策略，我们有&lt;/p&gt;
$$
V^\pi \leq V^*= \mathcal{T}^*(V^\pi)\leq ( \mathcal{T}^*)^2(V^\pi) \leq\cdots\leq (\mathcal{T}^*)^k(V^\pi) \to V^*=V^{\pi^*}
$$&lt;p&gt;这里第一个等式是 Bellman equation, 第一个不等式是由于下面的 Lemma,最后的极限使用了不动点定理。因此我们有 $V^\pi\leq V^*$, 从而 $\pi^*$ 是最优策略，且 $V^*$ 是对应的 optimal value function. $\blacksquare$&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Lemma&lt;/strong&gt;&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;令 $\pi$ 为一个策略， $\mathcal{T}^\pi$ 和 $\mathcal{T}^*$ 分别是 Bellman operator 和 Bellman optimality operator. 对任意 $V:\mathcal{S}^+\to\mathbb{R}$, 我们有&lt;/p&gt;
&lt;/blockquote&gt;
$$
\mathcal{T}^\pi(V) \leq \mathcal{T}^*(V)
$$&lt;blockquote&gt;
&lt;p&gt;进一步，对任意 $U:\mathcal{S}^+\to\mathbb{R}$, $V:\mathcal{S}^+\to\mathbb{R}$, 如果 $U\leq V$, 则&lt;/p&gt;
&lt;/blockquote&gt;
$$
\mathcal{T}^*(U) \leq \mathcal{T}^*(V)
$$&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
$$
\begin{aligned}
\mathcal{T}^\pi(V)(s) &amp;= \mathbb{E}_{a\sim \pi(\cdot\mid s),\ (r,s')\sim p(\cdot,\cdot\mid s,a_0))}\left[r+\gamma V(s')\mid s\right]\\
&amp;\leq \mathbb{E}_{a\sim \pi(\cdot\mid s)}\left[\max_{a'} \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a')}\left[r+\gamma V(s')\mid s\right]\right]\\
&amp;= \max_{a'} \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a')}\left[r+\gamma V(s')\mid s\right]\\
&amp;= \mathcal{T}^*(V)
\end{aligned}
$$&lt;p&gt;其次，令 $a^*\in\arg\max_{a}\mathbb{E}[r+\gamma U(s')]$, 那么&lt;/p&gt;
$$
\begin{aligned}
\mathcal{T}^*(U)(s) &amp;= \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a'))}\left[r+\gamma U(s')\mid s\right]\\
&amp;\leq \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a'))}\left[r+\gamma V(s')\mid s\right]\\
&amp;\leq \max_a\mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a'))}\left[r+\gamma V(s')\mid s\right]\\
&amp;= \mathcal{T}^*(V)
\end{aligned}
$$&lt;p&gt;证毕。 $\blacksquare$&lt;/p&gt;
&lt;p&gt;对于 Q-function, 我们也有类似结论&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Theorem&lt;/strong&gt;&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;假设 $\gamma\in(0,1)$, $|\mathcal{S}|&lt;\infty$ $|\mathcal{A}|&lt;\infty$ 以及 $|r|\leq R&lt;\infty, \mathrm{a.s.}$. 那么 optimal Q-function $Q^*:\mathcal{S}^+\times\mathcal{A}\to\mathbb{R}$ 存在，且满足 &lt;strong&gt;Bellman optimality equation&lt;/strong&gt;:&lt;/p&gt;
&lt;/blockquote&gt;
$$
\boxed{
Q^*(s, a) = \mathbb{E}_{(r,s')\sim p(\cdot,\cdot\mid s,a)}\left[r+\gamma \max_{a'\in\mathcal{A}}Q(s', a')\mid s, a\right]
}
$$&lt;blockquote&gt;
&lt;p&gt;反之，如果存在函数 $Q:\mathcal{S}^+\times\mathcal{A}\to\mathbb{R}$ 满足 Bellman optimality equation, 则 $Q=Q^*$.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;证明与 value function 类似，我们这里略过。&lt;/p&gt;</description></item><item><title>(RL series 1) Reinforcement Learning basic definitions</title><link>https://jiangyigithub.github.io/ai.github.io/p/rl-series-1-reinforcement-learning-basic-definitions/</link><pubDate>Wed, 18 Mar 2026 17:35:25 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/rl-series-1-reinforcement-learning-basic-definitions/</guid><description>&lt;h2 id="introduction"&gt;&lt;a href="#introduction" class="header-anchor"&gt;&lt;/a&gt;Introduction
&lt;/h2&gt;&lt;p&gt;RL 的基本思想是通过与环境进行交互，获取奖励来进行学习&lt;/p&gt;
&lt;p&gt;RL 的定义如下&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;强化学习是一个通过构建可以与环境进行交互的 agent 来解决控制和决策任务的学习框架，交互的方式为 agent 执行 action, 然后环境给予奖励&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;RL 的执行过程如下所示&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/rl-series-1-reinforcement-learning-basic-definitions/RL-process.png"
width="1400"
height="787"
loading="lazy"
alt="The RL process (source from sutton)"
class="gallery-image"
data-flex-grow="177"
data-flex-basis="426px"
&gt;&lt;/p&gt;
&lt;h2 id="rl-basics"&gt;&lt;a href="#rl-basics" class="header-anchor"&gt;&lt;/a&gt;RL Basics
&lt;/h2&gt;&lt;p&gt;RL 的数学建模依赖于 Markov decision process, 下面我们先介绍相关概念&lt;/p&gt;
&lt;h3 id="definition"&gt;&lt;a href="#definition" class="header-anchor"&gt;&lt;/a&gt;Definition
&lt;/h3&gt;&lt;blockquote&gt;
&lt;p&gt;一个 Markov decision process (MDP) 包含如下模块：&lt;/p&gt;
&lt;/blockquote&gt;
&lt;blockquote&gt;
&lt;ol&gt;
&lt;li&gt;时间 $t=0,1,\dots,T$.&lt;/li&gt;
&lt;li&gt;状态 $s_t\in\mathcal{S}\cup\{\langle term\rangle\}$, 其中 $\mathcal{S}$ 是状态空间 (state space)&lt;/li&gt;
&lt;li&gt;动作 $a_t\in\mathcal{A}$ , 其中 $\mathcal{A}$ 是动作空间 (action space)&lt;/li&gt;
&lt;li&gt;奖励 $r_t\in\mathbb{R}$, 环境对 agent 当前 action $a_t$ 给予的反馈&lt;/li&gt;
&lt;li&gt;终止时间 $T$, 额外定义 $s_T=\langle term\rangle$ 为终止状态&lt;/li&gt;
&lt;li&gt;初始状态分布 $s_o\sim p_0$.&lt;/li&gt;
&lt;li&gt;转换概率 (transition probability): $r_t,s_{t+1}\sim p(\cdot,\cdot\mid s_t,a_t)$.&lt;/li&gt;
&lt;li&gt;轨迹 $\tau=(s_0,a_0,r_0,s_1,a_1,r_1,\dots,s_{T-1},a_{T-1}, r_{T-1}, s_T)$.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;Markov property&lt;/strong&gt;. $t+1$ 时刻的状态仅与 $t$ 时刻的状态与动作相关，即 $p(s_{t+1}\mid s_t,a_t,\dots,a_0,a_0)=p(s_{t+1}\mid s_t,a_t)$ 以及 $p(r_{t}\mid s_t,a_t,\dots,s_0,a_t)=p(r_t\mid s_t,a_t)$.&lt;/li&gt;
&lt;/ol&gt;
&lt;/blockquote&gt;
&lt;p&gt;我们会对一些表达式进行简化：&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;当 $r_t$ 完全由 $(s_t,a_t)$ 决定时，我们记为 $r_t=r(s_t,a_t)$&lt;/li&gt;
&lt;li&gt;我们假设状态转移函数是平稳的 (stationary), 即 $p_t(r,s'\mid s,a)=p(r,s'\mid s,a)$。&lt;/li&gt;
&lt;li&gt;$a_t$ 通常由一个策略 $\pi$ 决定， 当 $\pi$ 完全由 $s_t$ 决定时，我们记为 $a_t=\pi(s_t)$, 反之，$a_t$ 从 $\pi$ 中采样得到，即 $a_t\sim\pi(\cdot\mid s_t)$.&lt;/li&gt;
&lt;li&gt;一般我们会使用一个神经网络来表示 $\pi$, 即 $\pi=\pi_\theta$, 这里 $\theta$ 就是神经网络的参数。&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;&lt;strong&gt;state and observation&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;通常来说，state 包含了当前进行决策所需要的所有信息（环境信息，历史信息，agent 自身信息），这种情况下 MDP 的 Markov 性质成立。实际上 agent 获取的可能只是一部分信息，即 agent 获取的为 observation, 这种情况下 MDP 特化为 Partially Observable Markov Decision Process (POMDP), 此时我们的 Markov 性质就不再对 observation 成立，我们的问题也变的更加困难。&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;Takeaway
State 表示整个环境的完整表示，而 observation 则是环境的部分表示&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;对于 LLM 来说，我们可以获取历史状态的所有信息，因此我们可以将 RL for LLM 视作一个 MDP 问题。&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;action space&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;action space $\mathcal{A}$ 可以是连续的，比如上下左右四个方向，也可以是连续的，比如方向盘的角度&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;episodic and continuing&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;根据任务的终止时间 $T$，我们可以将任务分为 episodic task 和 continuing task, 当 $T=\infty$ 时， 我们的任务称为 continual task，比如游戏, 否则称之为 episodic task, 比如股票预测.&lt;/p&gt;
&lt;p&gt;除了 POMDP 之外，MDP 也有一些其他的推广形式：&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;non-stationary dynamics, 即不同时间步的状态转移函数不同&lt;/li&gt;
&lt;li&gt;pre-determined terminal time $T$. 终止时间指定，因为 $p_{T-1}(\langle term\rangle\mid s_{T-1}, a_{T-1})=1$ 且 $p_{t+1}(\langle term\rangle\mid s_{t}, a_{t})=0$, $\forall t\in[T-2]$, 所以此时 MDP 为 non-stationary,&lt;/li&gt;
&lt;li&gt;policy $\pi_t$ may be time-dependent, 一般来说仅在 non-stationary dynamics 中有必要&lt;/li&gt;
&lt;li&gt;actions may depend on states, 即 $a_t\in\mathcal{A}(s_t)$.&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;由于 $T$ 也是一个随机变量，因此 $\mathbb{E}[\sum_{t=0}^T\cdot]\neq\sum_{t=0}^T\mathbb{E}[\cdot]$.&lt;/p&gt;
&lt;p&gt;为了简化，我们一般会使用一个等价的，不会停止的 MDP:&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;the MDP never stops, 即 $T=\infty$&lt;/li&gt;
&lt;li&gt;$s_t\in\mathcal{S}\cup\{\langle term\rangle\}$, 这里 $\langle term\rangle$ 是一个 normal state.&lt;/li&gt;
&lt;li&gt;absorbing state: 当 $s_t=\langle term\rangle$ 时， $r_t=0, s_{t+1}=\langle term\rangle$, 即状态不会离开 $\langle term\rangle$.&lt;/li&gt;
&lt;li&gt;$\pi(\cdot\mid\langle term\rangle)$ 不产生任何影响&lt;/li&gt;
&lt;/ol&gt;
&lt;h3 id="rl-objective"&gt;&lt;a href="#rl-objective" class="header-anchor"&gt;&lt;/a&gt;RL Objective
&lt;/h3&gt;&lt;p&gt;RL 的最终目标为最大化累计奖励，称为 expected return, 这个目标基于&lt;strong&gt;reward hypothesis&lt;/strong&gt;, 即&lt;em&gt;所有的目标都可以被描述为最大化 expected return&lt;/em&gt;. 目标函数表达式如下&lt;/p&gt;
$$
\max_{\pi}\quad \mathbb{E}_{s_0\sim p_0}^\pi\left[\sum_{t=0}^{T-1}r_t\right]
$$&lt;p&gt;由于未来存在不确定性，我们会对这种不确定性施加惩罚，即时间越远，其对于当前的 reward 就越低，这其实就是经济学中的现值 (present value, PV), 因此我们实际上优化的目标函数是 expected discounted return:&lt;/p&gt;
$$
\max_{\pi}\quad \mathbb{E}_{s_0\sim p_0}^\pi\left[\sum_{t=0}^{T-1}\gamma^tr_t\right]
$$&lt;p&gt;这里 $\gamma\in(0, 1]$ 是一个超参数，我们定义&lt;/p&gt;
$$
G_t = r_t +\gamma r_{t+1}+\cdots+\gamma^{T-1-t}r_{T-1}=\sum_{t'=t}^{T-1}\gamma^{t'-t}r_{t'}
$$&lt;p&gt;$G_t$ 被称为 &lt;strong&gt;discounted return&lt;/strong&gt;.&lt;/p&gt;</description></item><item><title>Fix Point Theorem</title><link>https://jiangyigithub.github.io/ai.github.io/p/fix-point-theorem/</link><pubDate>Mon, 09 Mar 2026 17:16:02 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/fix-point-theorem/</guid><description>&lt;p&gt;Fix Point Theorem, 即不动点定理，是泛函分析中的基本工具，被广泛应用于非线性函数的分析。&lt;/p&gt;
&lt;p&gt;在介绍不动点定理之前，我们先介绍两个概念&lt;/p&gt;
&lt;p&gt;首先是不动点的概念。&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Definition&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;对于函数 $f:\mathbb{R}^n\to\mathbb{R}^n$, 如果一个点 $x^*\in\mathbb{R}^n$ 满足&lt;/p&gt;
$$
f(x^*)=x^*
$$&lt;p&gt;则我们称 $x^*$ 是函数 $f$ 的不动点。&lt;/p&gt;
&lt;p&gt;接下来是 contraction mapping 的概念&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Definition&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;对于函数 $f:\mathbb{R}^n\to\mathbb{R}^n$, 如果存在 $\gamma\in(0,1)$ 满足&lt;/p&gt;
$$
\|f(x_1)-f(x_2)\| \leq \gamma \|x_1-x_2\|,\forall\ x_1,x_2\in\mathbb{R}^n
$$&lt;p&gt;则我们称 $f$ 是一个 contraction mapping. 这里 $\|\cdot\|$ 是一个 matrix norm.&lt;/p&gt;
&lt;p&gt;接下来，我们介绍不动点定理&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Theorem&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;给定 equation $x=f(x)$, 其中 $f:\mathbb{R}^n\to\mathbb{R}^n$, 如果 $f$ 是一个 contraction mapping, 则 $f$ 具有如下性质&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;Existence: 存在 fixed point $x^*\in\mathbb{R}^n$ 满足 $f(x^*)=x^*$.&lt;/li&gt;
&lt;li&gt;Uniqueness: fixed point $x^*$ 唯一。&lt;/li&gt;
&lt;li&gt;Algorithm: 对任意 $x^0\in\mathbb{R}^n$, 使用迭代算法 $x_{k+1}=f(x_k)$ 产生的序列 $\{x_k\}_{k=0}^{\infty}$ 收敛到 fixed point $x*$, 且收敛速度为指数级。&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;证明需要用到柯西列的概念。&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Definition&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;一个序列 $x_1,x_2,\dots$ 被称为柯西列 (Cauchy sequence) 当且仅当对任意 $\epsilon&gt;0$, 都存在 $N&gt;0$, 使得&lt;/p&gt;
$$
\|x_m-x_n\| &lt;\epsilon,\forall m, n&gt;N
$$&lt;p&gt;柯西列的一个重要性质为柯西列一定是收敛列。&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;证明&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;我们首先证明由 $x_k=f(x_{k=1})$ 产生的序列 $\{x_k\}_{k=1}^{\infty}$ 是收敛的，我们通过证明序列 $\{x_k\}_{k=1}^{\infty}$ 是一个柯西列来证明这一点。&lt;/p&gt;
&lt;p&gt;注意到 $f$ 是一个 contraction mapping, 因此&lt;/p&gt;
$$
\|x_{k+1}-x_k\| = \|f(x_k)-f(x_{k-1})\|\leq \gamma \|x_k-x_{k-1}\|
$$&lt;p&gt;迭代下去，我们就得到&lt;/p&gt;
$$
\|x_{k+1}-x_k\|\leq \gamma \|x_k-x_{k-1}\|\leq\cdots\leq \gamma^k \|x_1-x_{0}\|
$$&lt;p&gt;现在我们证明序列 $\{x_k\}_{k=1}^{\infty}$ 是一个柯西列：&lt;/p&gt;
$$
\begin{aligned}
\|x_m-x_n\| &amp;= \|x_m-x_{m-1}+x_{m-1}-\cdots-x_{n+1}+x_{n+1}-x_n\|\\
&amp;\leq \sum_{i=n}^{m-1}\|x_{i+1}-x_i\|\\
&amp;\leq \sum_{i=n}^{m-1}\gamma^i \|x_{1}-x_0\|\\
&amp;\leq \frac{\gamma^n}{1-\gamma}\|x_{1}-x_0\|.
\end{aligned}
$$&lt;p&gt;从而，序列 $\{x_k\}_{k=1}^{\infty}$ 是一个柯西列, 因此也是一个收敛列。&lt;/p&gt;
&lt;p&gt;接下来，我们证明 $x^*=\lim_{k\to\infty}x_k$ 是 $f(x)$ 的不动点，注意到&lt;/p&gt;
$$
\|f(x_k)-x_k\| = \|x_{k+1}-x_k\|\leq \gamma^k\|x_1-x_0\| \to 0, k\to\infty
$$&lt;p&gt;我们有 $\lim_{k\to\infty}f(x_k)=\lim_{k\to\infty}x_k$ , 由于 contraction mapping 一定是连续的，因此我们就可以得到 $f(x^*)=x^*$.&lt;/p&gt;
&lt;p&gt;然后，我们证明不动点唯一。假设还存在一个另外一个不动点 $x'\neq x^*$ 满足 $f(x')=x'$, 那么&lt;/p&gt;
$$
\|x'-x^*\| = \|f(x')-f(x')\| \leq \gamma \|x'-x^*\|
$$&lt;p&gt;由于 $\gamma\in(0,1)$, 因此上述等式当且仅当 $\|x'-x^*\|=0$, 这与前面假设矛盾，因而不动点是唯一的&lt;/p&gt;
&lt;p&gt;最后，我们证明 $x_{k+1}=f(x_k)$ 这个算法的收敛速度为指数级，注意到&lt;/p&gt;
$$
\|x^*-x_n\| = \lim_{m\to\infty}]\|x_m-x_n\| \leq \frac{\gamma^n}{1-\gamma}\|x_1-x_0\|
$$&lt;p&gt;因为 $\gamma &lt;1$, 因此收敛速度为指数级&lt;/p&gt;</description></item><item><title>Notes on KL divergence</title><link>https://jiangyigithub.github.io/ai.github.io/p/notes-on-kl-divergence/</link><pubDate>Sat, 24 Jan 2026 16:32:14 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/notes-on-kl-divergence/</guid><description>&lt;h2 id="introduction"&gt;&lt;a href="#introduction" class="header-anchor"&gt;&lt;/a&gt;Introduction
&lt;/h2&gt;&lt;p&gt;在本节中，我们先介绍 KL divergence 的基本定义，然后我们介绍 KL divergence 的一般形式，即 f-divergence.&lt;/p&gt;
&lt;h3 id="kl-divergence"&gt;&lt;a href="#kl-divergence" class="header-anchor"&gt;&lt;/a&gt;KL-divergence
&lt;/h3&gt;&lt;p&gt;KL divergence 用于衡量近似概率分布 $Q(x)$ 到真实概率分布 $P(x)$ 的误差，我们可以将其理解为：如果我们用 $Q(x)$ 来替换 $P(x)$, 会有多大的信息损失？&lt;/p&gt;
&lt;p&gt;连续概率分布的 KL divergence 的定义如下&lt;/p&gt;
$$
D_{KL}(P\parallel Q) =\mathbb{E}_{x\sim P}\left[\log \frac{P(x)}{Q(x)}\right]=\int P(x)\log\left(\frac{P(x)}{Q(x)}\right)dx
$$&lt;p&gt;离散概率分布的 KL divergence 定义如下&lt;/p&gt;
$$
D_{KL}(P\parallel Q) = \sum_{x} P(x)\log\left(\frac{P(x)}{Q(x)}\right)
$$&lt;p&gt;KL divergence 有两几个关键性质：&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;非负性：$D_{KL}(P\parallel Q)\geq0$, 且 $D_{KL}(P\parallel Q)=0$ 当且仅当 $P(x)=Q(x)$ 对任意 $x$ 成立&lt;/li&gt;
&lt;li&gt;非对称性： 一般情况下，$D_{KL}(P\parallel Q)\neq D_{KL}(Q\parallel P)$.&lt;/li&gt;
&lt;li&gt;有限性：如果存在 $x$ 使得 $P(x)&gt;0$ 但是 $Q(x)=0$, 则 $D_{\mathrm{KL}}(P\parallel Q)=\infty$.&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;一般我们称 $D_{KL}(P\parallel Q)$ 为 &lt;strong&gt;forward KL&lt;/strong&gt; (相对于 $Q$), 对应的还有 &lt;strong&gt;reverse KL&lt;/strong&gt; $D_{KL}(Q\parallel P)$ (相对于 $Q$).&lt;/p&gt;
&lt;h3 id="f-divergence"&gt;&lt;a href="#f-divergence" class="header-anchor"&gt;&lt;/a&gt;F-divergence
&lt;/h3&gt;&lt;p&gt;KL divergence 是 f-divergence 的一种特殊情况。 f-divergence 是一类衡量不同概率分布 $P$ 和 $Q$ 的函数 $D_f(P\parallel Q)$.&lt;/p&gt;
&lt;p&gt;假设函数 $f:(0,\infty)\to\mathbb{R}$ 是一个凸函数，且 $f(1)=0$. $P$ 和 $Q$ 是两个概率分布，则 f-divergence 定义如下&lt;/p&gt;
$$
D_f(P\parallel Q) = \mathbb{E}_{x\sim Q}\left[ f\left(\frac{P(x)}{Q(x)}\right)\right]=\int Q(x)f\left(\frac{P(x)}{Q(x)}\right)dx
$$&lt;p&gt;我们称 $f$ 为 $D_f$ 的 &lt;strong&gt;generator&lt;/strong&gt;.&lt;/p&gt;
&lt;p&gt;以下是几种常见的 f-divergence:&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Name&lt;/th&gt;
&lt;th&gt;generator&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;forward KL divergence&lt;/td&gt;
&lt;td&gt;$f(x)=x\log x$&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;reverse KL divergence&lt;/td&gt;
&lt;td&gt;$f(x)=-\log x$&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;Total variation&lt;/td&gt;
&lt;td&gt;$f(x)=1/2\vert x-1\vert$&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\chi^2$-divergence&lt;/td&gt;
&lt;td&gt;$f(x)=(x-1)^2$&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;JS-divergence&lt;/td&gt;
&lt;td&gt;$f(x)=x\log\frac{2x}{x+1}+\log\frac{2}{x+1}$&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;p&gt;我们这里推导一下 KL divergence 对应的 generator.&lt;/p&gt;
&lt;p&gt;对于 forward KL, 注意到&lt;/p&gt;
$$
D_f(P \parallel Q) = \int Q(x) \left( \frac{P(x)}{Q(x)} \log \frac{P(x)}{Q(x)} \right) dx = \int P(x) \log \frac{P(x)}{Q(x)} dx = D_{KL}(P \parallel Q)
$$&lt;p&gt;因此 forward KL 对应的 generator 为 $f=x\log x$.&lt;/p&gt;
&lt;p&gt;对于 reverse KL, 注意到&lt;/p&gt;
$$
D_f(P \parallel Q) = \int Q(x) \left( -\log \frac{P(x)}{Q(x)} \right) dx = \int Q(x) \log \frac{Q(x)}{P(x)} dx = D_{KL}(Q \parallel P)
$$&lt;p&gt;因此 forward KL 对应的 generator 为 $f=-\log x$.&lt;/p&gt;
&lt;h4 id="properties-of-f-divergence"&gt;&lt;a href="#properties-of-f-divergence" class="header-anchor"&gt;&lt;/a&gt;Properties of F-divergence
&lt;/h4&gt;&lt;p&gt;f-divergence 性质如下&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;linearity: $D_{a_1f_1+a_2f_2}=a_1D_{f_1}+a_2D_{f_2}$.&lt;/li&gt;
&lt;li&gt;$D_f=D_g$ 当且仅当存在 $c\in\mathbb{R}$ 使得 $f(x)=g(x)+c(x-1)$.&lt;/li&gt;
&lt;li&gt;non-negativity. $D_f(P\parallel Q)\geq0$ 且 $D_f(P\parallel Q)$ 当且仅当 $P=Q$.&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;性质 2 证明如下：&lt;/p&gt;
&lt;p&gt;如果 $f(x)=g(x)+c(x-1)$, 则通过定义，我们可以验证得到 $D_f=D_g$.&lt;/p&gt;
&lt;p&gt;反之，如果 $D_f=D_g$, 令 $h=f-g$, 对任意两个在集合 $\{0, 1\}$ 上的概率分布 $P,Q$, 由于 $D_f(P\parallel Q) - D_g(P\parallel Q)=0$, 我们有&lt;/p&gt;
$$
h\left(\frac{P(1)}{Q(1)}\right) = -\frac{Q(0)}{Q(1)}h\left(\frac{P(0)}{Q(0)}\right)
$$&lt;p&gt;我们不妨假设 $P(0)=aQ(0)$, $P(1)=bQ(1)$, 结合 $P(0)+P(1)=1$ 和 $Q(0)+Q(1)=1$ 我们有&lt;/p&gt;
$$
Q(0) = \frac{1-a}{b-a}, Q(1) = \frac{b-1}{b-a}
$$&lt;p&gt;从而&lt;/p&gt;
$$
\frac{h(b)}{b-1}=\frac{h(a)}{a-1}
$$&lt;p&gt;由于我们可以任意选定 $P$ 和 $Q$, 因此 $h$ 是一个线性函数，形式为 $h(x)=c(x-1)$. $\blacksquare$&lt;/p&gt;
&lt;h2 id="approximation"&gt;&lt;a href="#approximation" class="header-anchor"&gt;&lt;/a&gt;Approximation
&lt;/h2&gt;&lt;p&gt;本节中，我们将介绍针对 KL divergence 的三种近似形式。&lt;/p&gt;
&lt;p&gt;在实际计算 KL divergence 时，由于：&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;完整计算 KL divergence 需要的算力或内存过高&lt;/li&gt;
&lt;li&gt;没有闭式解&lt;/li&gt;
&lt;li&gt;我们可以仅保存 log-probability, 而不是整个概率分布&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;因此，我们假设我们只能计算输入 $x$ 对应的概率 $P(x)$ 和 $Q(x)$. 一般来说，我们会通过 Monte Carlo estimate 来进行近似。即我们先对 $P$ 进行采样得到 $x_1,\dots,x_N\sim P$, 然后我们构建估计量。&lt;/p&gt;
&lt;p&gt;一个高的估计量应该是无偏 (unbiased) 并且方差低 (low variance) 的。John Schulman 给出了三种 estimator. 我们分别针对 forward KL 和 reverse KL 进行介绍。这里我们定义&lt;/p&gt;
$$
r = \frac{P(x)}{Q(x)}
$$&lt;h3 id="forward-kl-estimation"&gt;&lt;a href="#forward-kl-estimation" class="header-anchor"&gt;&lt;/a&gt;Forward KL Estimation
&lt;/h3&gt;&lt;p&gt;对于 forward KL $D_{KL}(P\parallel Q)$, 其对应的 generator 为 $f(x)=x\log x$, 注意到 $\mathbb{E}_{x\sim Q}[r]=1$, 且 $f$ 是一个凸函数，因此我们有 $f(r)-f'(1)(r-1)\geq0$, 从而我们可以得到一个新的估计为 $\boxed{k=r\log r - (r-1)}$.&lt;/p&gt;
&lt;h3 id="reverse-kl-estimation"&gt;&lt;a href="#reverse-kl-estimation" class="header-anchor"&gt;&lt;/a&gt;Reverse KL Estimation
&lt;/h3&gt;&lt;p&gt;对于 reverse KL $D_{KL}(Q\parallel P)$, 其对应的 generator 为 $f(x)=-\log x$, 由概率性质，$\boxed{k_1=-\log r}$ 是 $D_{KL}(Q\parallel P)$ 的一个无偏估计。但是 $k_1$ 的问题在于 当 $r$ 非常小时，$k_1$ 会变得非常大。也就是说，$k_1$ 的 variance 比较高。&lt;/p&gt;
&lt;p&gt;John Schulman 基于 f-divergence 泰勒展开给出了一个新的估计 $k_2$, 其定义为&lt;/p&gt;
$$
\boxed{k_2 = \frac12(\log r)^2}
$$&lt;p&gt;其期望为&lt;/p&gt;
$$
\mathbb{E}_Q[k_2] = \mathbb{E}_Q\left[\frac12(\log r)^2\right]
$$&lt;p&gt;这是一个 f-divergence, 对应的 generator 为 $f_{k_2}(x)=1/2(\log x)^2$, 而 $D_{KL}(Q\parallel P)$ 对应的 generator 为 $f_{k_1}(x)=-\log x$.&lt;/p&gt;
&lt;p&gt;当 $P$ 和 $Q$ 比较靠近时，我们记 $\theta=r-1$， 对 $D_{f}(P\parallel Q)$ 在 $x=1$ 处进行展开得到&lt;/p&gt;
$$
\begin{aligned}
D_f(P\parallel Q) &amp;= \mathbb{E}_{x\sim Q}\left[ f(r)\right]\\
&amp;= \mathbb{E}_{x\sim Q}\left[ f(1) + f'(1)\theta + \frac{f''(1)}{2}f(1+\lambda)\theta^2+O(\theta^3)\right]\\
&amp;= \frac{f''(1)}{2}F\theta^2+O(\theta^3)
\end{aligned}
$$&lt;p&gt;这里我们应用了 $f(1)=0$, $\mathbb{E}[\theta]=0$, $F=\mathbb{E}[f(1+\lambda\theta)$ 是 Fisher information matrix.&lt;/p&gt;
&lt;p&gt;我们分别带入 $f_{k_1}(x)$ 和 $f_{k_2}(x)$ 得到 $f_{k_1}''(1)=f_{k_2}''(1)=1$, 即 $k_1$ 和 $k_2$ 在 $P$ 和 $Q$ 比较靠近时二阶近似是相同的。因此，**$k_2$ 表面上是一个二阶近似，在分布接近时有效，但本质上是在优化 &lt;em&gt;另一个 f-divergence&lt;/em&gt;*&lt;/p&gt;
&lt;p&gt;John Schulman 还构造了第三种估计。回顾前面 f-divergence 的性质 2，即当 $f(x)=g(x)+c(x-1)$ 时，我们有 $D_f=D_g$, 因此我们可以选取合适的 $c$ 来降低估计的 variance. 注意到 $k_1$ 的主要问题在于存在负数的可能性，因此我们就构建一个对应的估计量来解决这个问题。注意到 $\log x \leq x -1$, 因此我们可以令 $c=1$, 此时就得到了新的估计&lt;/p&gt;
$$
\boxed{k_3 =(r-1)- \log r }
$$&lt;p&gt;$k_3$ 继承了 $k_1$ 的无偏性，并且 $k_3$ 通过 f-divergence 等价类消除了负值，兼顾无偏与低方差，解决了 $k_1$ variance 过大的问题&lt;/p&gt;
&lt;h3 id="experiments-on-approximation"&gt;&lt;a href="#experiments-on-approximation" class="header-anchor"&gt;&lt;/a&gt;Experiments on Approximation
&lt;/h3&gt;&lt;p&gt;对于分布 $P=\mathcal{N}(0,1)$ 以及 $Q=\mathcal{N}(0.1, 1)$, 真实的 KV divergence 为 0.005, 三个 estimator 的误差如下表所示&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Method&lt;/th&gt;
&lt;th&gt;Bias&lt;/th&gt;
&lt;th&gt;Std Dev&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;$k_1$&lt;/td&gt;
&lt;td&gt;0.0001&lt;/td&gt;
&lt;td&gt;20.0005&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$k_2$&lt;/td&gt;
&lt;td&gt;0.0025&lt;/td&gt;
&lt;td&gt;1.4175&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$k_3$&lt;/td&gt;
&lt;td&gt;0.0000&lt;/td&gt;
&lt;td&gt;1.4163&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;p&gt;当 $P=\mathcal{N}(1,1)$, $Q=\mathcal{N}(0.1, 1)$ 时， 真实的 KV divergence 为 0.405, 三个 estimator 的误差如下表所示&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Method&lt;/th&gt;
&lt;th&gt;Bias&lt;/th&gt;
&lt;th&gt;Std Dev&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;$k_1$&lt;/td&gt;
&lt;td&gt;-0.0000&lt;/td&gt;
&lt;td&gt;2.2223&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$k_2$&lt;/td&gt;
&lt;td&gt;0.2025&lt;/td&gt;
&lt;td&gt;1.6762&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$k_3$&lt;/td&gt;
&lt;td&gt;0.0000&lt;/td&gt;
&lt;td&gt;1.6342&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;p&gt;可以看到 $k_1$ 的 variance 非常大，$k_2$ 是一个有偏估计，$k_3$ 既满足了无偏又满足了 low variance.&lt;/p&gt;
&lt;h3 id="summary"&gt;&lt;a href="#summary" class="header-anchor"&gt;&lt;/a&gt;Summary
&lt;/h3&gt;&lt;p&gt;我们接下来总结 reverse KL $D_{KL}(Q\parallel P)$ 的近似 $k_1$, $k_2$ 和 $k_3$ 的性质如下 ($r=P(x)/Q(x)$)&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;estimation&lt;/th&gt;
&lt;th&gt;definition&lt;/th&gt;
&lt;th&gt;motivation&lt;/th&gt;
&lt;th&gt;bias&lt;/th&gt;
&lt;th&gt;variance&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;$k_1$&lt;/td&gt;
&lt;td&gt;$-\log r$&lt;/td&gt;
&lt;td&gt;naive estimation&lt;/td&gt;
&lt;td&gt;unbiased&lt;/td&gt;
&lt;td&gt;high&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$k_2$&lt;/td&gt;
&lt;td&gt;$\frac12(\log r)^2$&lt;/td&gt;
&lt;td&gt;f-divergence, taylor expansion&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;low&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$k_3$&lt;/td&gt;
&lt;td&gt;$(r-1)- \log r$&lt;/td&gt;
&lt;td&gt;f-divergence, non-negativity&lt;/td&gt;
&lt;td&gt;unbiased&lt;/td&gt;
&lt;td&gt;low&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;h2 id="applications-to-ml"&gt;&lt;a href="#applications-to-ml" class="header-anchor"&gt;&lt;/a&gt;Applications to ML
&lt;/h2&gt;&lt;blockquote&gt;
&lt;p&gt;Remark
本节内容主要参考了 &lt;a class="link" href="https://dibyaghosh.com/blog/probability/kldivergence/" target="_blank" rel="noopener"
&gt;KL Divergence for Machine Learning&lt;/a&gt;&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;我们假设真实目标分布和近似的目标分布分别记为 $p_{data}(x)$ 和 $p_\theta(x)$. 由于 KL divergence 的非对称性，因此我们需要考虑两种目标函数：&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;forward KL: $\arg\min_\theta D_{KL}(p_{data}\parallel p_\theta)$&lt;/li&gt;
&lt;li&gt;reverse KL: $\arg\min_\theta D_{KL}(p_\theta \parallel p_{data})$&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;我们将会看到，这两种不同的目标函数导致的结果也不尽相同&lt;/p&gt;
&lt;h3 id="forward-kl"&gt;&lt;a href="#forward-kl" class="header-anchor"&gt;&lt;/a&gt;Forward KL
&lt;/h3&gt;&lt;p&gt;对目标函数进行简化得到&lt;/p&gt;
$$
\arg\min_\theta D_{KL}(p_{data}\parallel p_\theta) = \arg\max_\theta \mathbb{E}_{x\sim p_{data}}\left[\log p_\theta(x)\right]
$$&lt;p&gt;实际在计算时，我们会使用 Monte Carlo 的方式对真实分布进行采样然后进行估计。&lt;/p&gt;
&lt;p&gt;Forward KL 其代表的含义为，我们从分布 $p_{data}$ 中进行采样，然后求 $p_\theta$ 的最大似然估计。最终的结果满足：&lt;strong&gt;当 $p_{data}(x)$ 概率很高时，$p_\theta(x)$ 的概率也需要很高&lt;/strong&gt;. 这是一种 &lt;strong&gt;mean-seeking&lt;/strong&gt; behavior, 因为 $p_\theta$ 必须覆盖 $p_{data}$ 的所有 modes.&lt;/p&gt;
&lt;p&gt;一般来说，supervised learning 对应的就是 forward KL. 我们可以证明 forward KL divergence 和 MLE 是等价的。也就是说，最大似然估计得到的分布就是 KL divergence 最小的近似分布。我们将 $p_{data}(x)$ 和 $p_\theta(x)$ 对应的 KL divergence 进行展开得到&lt;/p&gt;
$$
\begin{aligned}
\theta_{KL}^* &amp;= \arg\min_{\theta}D_{KL}(p_{data}(x)\parallel p_\theta(x))\\
&amp;= \arg\min_{\theta} \int p_{data}(x)\frac{p_{data}(x)}{p_\theta(x)} dx\\
&amp;= \arg\min_{\theta}\int p_{data}(x)\log p_{data}(x) dx - \int p_{data}(x)\log p_\theta(x)dx \\
&amp;= \arg\min_{\theta} - \int p_{data}(x)\log p_\theta(x)dx \\
&amp;= \arg\max_{\theta} \int p_{data}(x)\log p_\theta(x)dx
\end{aligned}
$$&lt;p&gt;实际上，真实的数据分布 $p_{data}(x)$ 是未知的，我们只有从 $p_{data}(x)$ 采样得到的一批数据 $X=\{x_1,\dots,x_n\}\sim p_{data}(x)$. 基于大数定律，我们有&lt;/p&gt;
$$
\frac{1}{n}\sum_{i=1}^n\log p(\theta_i\mid \theta)=\mathbb{E}_{x\sim p_{data}}[\log p_\theta(x)] = \int p_{data}(x)\log p_\theta(x)dx, n\to \infty
$$&lt;p&gt;这样，最大似然估计就与最小化 KL divergence 构建起了联系：&lt;/p&gt;
$$
\begin{aligned}
\theta_{MLE}^*&amp;=\arg\max_{\theta} \sum_{i=1}^n \log p(x_i\mid \theta)\\
&amp;= \arg\max_{\theta} \int p_{data}(x)\log p_\theta(x)dx\\
&amp;= \theta_{KL}^*, n\to\infty.
\end{aligned}
$$&lt;p&gt;也就是说，当采样样本足够多的时候，最大似然估计和最小 KL divergence 是等价的。监督学习中，我们先从真实分布 $p_{data}(x,y)$ 中收集一个数据集 $\mathcal{D}=\{(x_i,y_i)\}$, 然后我们会基于模型 $f_\theta:\mathcal{X}\to\mathcal{Y}$ 和损失函数 $\mathcal{L}:\mathcal{Y}\times\mathcal{Y}\to\mathbb{R}$ 来优化模型参数 $\theta$:&lt;/p&gt;
$$
\arg\min_\theta \mathbb{E}_{(x_i,y_i)\sim\mathcal{D}}[\mathcal{L}(f_\theta(x_i), y_i)]
$$&lt;p&gt;对于使用 cross-entropy loss 的分类问题以及 MSE loss 的回归问题，其目标函数实际上都是最小化 KL divergence.&lt;/p&gt;
&lt;h3 id="reverse-kl"&gt;&lt;a href="#reverse-kl" class="header-anchor"&gt;&lt;/a&gt;Reverse KL
&lt;/h3&gt;&lt;p&gt;对目标函数进行简化，得到&lt;/p&gt;
$$
\arg\min_\theta D_{KL}(Q_\theta\parallel p_{data}) = \arg\max_\theta \mathbb{E}_{x\sim Q_\theta}\left[\log p_{data}(x)\right] - \mathbb{E}_{x\sim Q_\theta}\left[\log Q_\theta(x)\right]
$$&lt;p&gt;实际在计算时，我们需要知道真实概率分布在采样点上的概率值 $p_{data}(x)$.&lt;/p&gt;
&lt;p&gt;Reverse KL 代表的含义为，我们从分布 $p_\theta(x)$ 中进行采样，然后最大化采样点在 $p_{data}(x)$ 中的概率分布。entropy item 鼓励 $p_\theta$ 尽可能均匀分布（覆盖广），从而最终结果满足：&lt;strong&gt;当 $p_\theta(x)$ 概率很高时，$p_{data}(x)$ 的概率也需要很高&lt;/strong&gt;。注意到与 forward KL 不同，Reverse KL 中包含 entropy 项，其避免了 $p_\theta$ 收缩到 $p_{data}$ 的某一个 非常窄的 mode 上，最终结果是 $p_\theta$ 会找到 $p_{data}$ 的一个 &lt;strong&gt;high probability&lt;/strong&gt; 以及 &lt;strong&gt;wide support&lt;/strong&gt; 的 mode, 然后进行覆盖。&lt;/p&gt;
&lt;p&gt;一般来说，reinforcement learning 对应的就是 reverse KL, 这是因为我们希望 policy model 不要离 reference model 太远，并不一定要 cover 所有的 mode.&lt;/p&gt;
&lt;h3 id="experiments-on-forward-and-reverse-kl"&gt;&lt;a href="#experiments-on-forward-and-reverse-kl" class="header-anchor"&gt;&lt;/a&gt;Experiments on forward and Reverse KL
&lt;/h3&gt;&lt;p&gt;我们通过概率分布来可视化 forward RL 与 reverse RL 的区别，验证 forward KL 与 reverse KL 不同的模式。&lt;/p&gt;
&lt;p&gt;我们假设 $p_{data}=w_1\mathcal{N}(\mu_1, \sigma_1^2)+w_2\mathcal{N}(\mu_2, \sigma_2^2)$, 然后我们用一个 normal distribution $p_\theta=\mathcal{N}(\mu, \sigma^2)$ 来近似 $p_{data}$, 这里 $\theta=(\mu, \sigma^2)$. 对于 forward KL, 我们可以从理论上得出最优解，对应的 $\mu=w_1\mu_1+w_2\mu_2$, 而 reverse KL 则只能通过优化的方式进行求解，并且解与初始化条件相关，下面是相关的实验结果&lt;/p&gt;
&lt;p&gt;首先我们令 $w_1=w_2=0.5$, $\mu_1=\mu_2=4.0$, $\sigma_1=\sigma_2=1$, reverse KL 的初始化条件为 $\theta_0=(2,1)$, 对应的结果为&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/notes-on-kl-divergence/KL-divergence-reverse-kl-vis1.png"
width="1010"
height="549"
loading="lazy"
alt="visualization of forward KL v.s. reverse KL (1)"
class="gallery-image"
data-flex-grow="183"
data-flex-basis="441px"
&gt;&lt;/p&gt;
&lt;p&gt;接下来我们改变 reverse KL 的初始化条件为 $\theta_0=(-2,1)$, 对应的结果为&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/notes-on-kl-divergence/KL-divergence-reverse-kl-vis2.png"
width="1010"
height="549"
loading="lazy"
alt="visualization of forward KL v.s. reverse KL (2)"
class="gallery-image"
data-flex-grow="183"
data-flex-basis="441px"
&gt;&lt;/p&gt;
&lt;p&gt;可以看到，与前面分析一致，使用 forward KL 时，最终得到的 $p_\theta$ 会倾向于拟合分布的中心 (mean seeking), 即 $\mu(p_\theta)=\mu(p_{data})$, 而使用 reverse KL 时，最终得到的 $P$ 会倾向于拟合分布的 mode (mode seeking).&lt;/p&gt;
&lt;h2 id="applications-to-rl"&gt;&lt;a href="#applications-to-rl" class="header-anchor"&gt;&lt;/a&gt;Applications to RL
&lt;/h2&gt;&lt;blockquote&gt;
&lt;p&gt;Remark
本节内容主要参考了 &lt;a class="link" href="https://xihuai18.github.io/reinforcement-learning/2025/12/01/kl-estimators-en.html" target="_blank" rel="noopener"
&gt;Understanding KL Divergence Estimators in RL: From Value Approximation to Gradient Estimation&lt;/a&gt;&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;在本节中，我们将基于 RL 来推导 KL 的相关性质。为了统一，这里我们使用 RL 中常见的 notation 来进行计算&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;notation&lt;/th&gt;
&lt;th&gt;description&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;$\pi_\theta$&lt;/td&gt;
&lt;td&gt;policy model with parameter $\theta$&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\pi_{ref}$&lt;/td&gt;
&lt;td&gt;reference model&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\pi_{old}$&lt;/td&gt;
&lt;td&gt;behavior model to sample from&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$s_\theta(x)=\nabla_\theta \log \pi_\theta(x)$&lt;/td&gt;
&lt;td&gt;score function&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\rho(x)=\pi_\theta(x)/\pi_{old}(x)$&lt;/td&gt;
&lt;td&gt;importance weight&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\mathrm{sg}(\cdot)$&lt;/td&gt;
&lt;td&gt;stop gradient operation&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;p&gt;首先 score function 有一个期望为 0 的性质：&lt;/p&gt;
$$
\mathbb{E}_{x\sim\pi_\theta}[s_\theta(x)]=\int_x \pi_\theta(x)\nabla_\theta \log \pi_\theta(x)dx = \int_x\nabla_\theta \pi_\theta(x)dx= \nabla_\theta\int_x \pi_\theta(x)dx =\nabla_\theta1 = 0
$$&lt;p&gt;接下来，我们分别推导 forward KL 和 reverse KL 的梯度。对于 forward KL, 我们有&lt;/p&gt;
$$
\nabla_\theta D_{KL}(\pi_{ref}\parallel \pi_\theta) = -\int \pi_{ref}\nabla_\theta \log \pi_\theta dx=-\mathbb{E}_{\pi_{ref}}[s_\theta] = \boxed{-\mathbb{E}_{\pi_\theta}\left[\frac{\pi_{ref}}{\pi_\theta}s_\theta\right]}
$$&lt;p&gt;对于 reverse KL,我们有&lt;/p&gt;
$$
\begin{aligned}
\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})&amp; = \int\left[\nabla_\theta \pi_\theta\cdot\log\frac{\pi_\theta}{\pi_{ref}} + \pi_\theta \nabla_\theta\log \frac{\pi_\theta}{\pi_{ref}}\right]dx\\
&amp;= \int \pi_\theta s_\theta\log \frac{\pi_\theta}{\pi_{ref}}dx + \int \pi_\theta s_\theta dx\\
&amp;= \mathbb{E}_{\pi_\theta}\left[s_\theta\log \frac{\pi_\theta}{\pi_{ref}}\right]+\mathbb{E}_{\pi_\theta}[s_\theta]\\
&amp;= \boxed{\mathbb{E}_{\pi_\theta}\left[s_\theta\log \frac{\pi_\theta}{\pi_{ref}}\right]}
\end{aligned}
$$&lt;p&gt;这里我们使用了 $\nabla_\theta\pi_\theta=\pi_\theta s_\theta$ , $\nabla_\theta\log\pi_\theta=s_\theta$ 以及 前面推导的 $\mathbb{E}_{\pi_\theta}[s_\theta]=0$ 的结论.&lt;/p&gt;
&lt;p&gt;RL 的目标函数如下&lt;/p&gt;
$$
\mathcal{J}(\theta) = \mathbb{E}_{\tau\sim \pi_\theta}\left[\sum_{t=0}^T\gamma^tr(s_t,a_t)\right] - \beta D_{KL}(\pi_\theta\parallel \pi_{ref})
$$&lt;h3 id="ki-as-loss"&gt;&lt;a href="#ki-as-loss" class="header-anchor"&gt;&lt;/a&gt;Ki as Loss
&lt;/h3&gt;&lt;p&gt;由于 KL divergcne 不能直接计算（或者计算难度较大），因此，基于前面对 KL divergence estimation 的分析，我们可以使用如下代理损失函数来优化我们的模型：&lt;/p&gt;
$$
\mathcal{J}_1(\theta) = \mathbb{E}_{\tau\sim \pi_\theta}\left[\sum_{t=0}^T\gamma^tr(s_t,a_t)\right] - \beta k_i(\pi_\theta, \pi_{ref})
$$&lt;p&gt;这里 $i\in\{1,2,3\}$ 代表了我们使用的估计。从直觉上来说，这样做是没问题的，但是我们将从数学分析上说明，$k_1,k_3$ 作为损失函数都存在问题。其核心问题在于&lt;/p&gt;
$$
\mathbb{E}[\widehat{D_{KL}}]=D_{KL} \nRightarrow \mathbb{E}[\nabla_\theta \widehat{D_{KL}}] =\nabla_\theta D_{KL}
$$&lt;p&gt;也就是说，&lt;strong&gt;KL divergence estimation 的无偏性不能推导出 KL divergence estimation gradient 的无偏性，这是因为我们在求期望时，对应的概率分布可能也与参数相关&lt;/strong&gt;。实际上，我们有&lt;/p&gt;
$$
\begin{aligned}
\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref}) &amp;= \nabla_\theta \mathbb{E}_{x\sim\pi_\theta}[\widehat{D_{KL}}(\pi_\theta\parallel \pi_{ref})]\\
&amp;= \mathbb{E}_{x\sim\pi_\theta}[\nabla_\theta \widehat{D_{KL}}(\pi_\theta\parallel \pi_{ref})] + \mathbb{E}_{x\sim\pi_\theta}[\widehat{D_{KL}}(\pi_\theta\parallel \pi_{ref})\nabla_\theta \pi_\theta(x)]\\
&amp;\neq \mathbb{E}_{x\sim\pi_\theta}[\nabla_\theta \widehat{D_{KL}}(\pi_\theta\parallel \pi_{ref})]
\end{aligned}
$$&lt;p&gt;因此 $\nabla_\theta \widehat{D_{KL}}$ 是 $\nabla_\theta D_{KL}$ 的一个有偏估计。&lt;/p&gt;
&lt;p&gt;我们分别来分析一下 $k_1,k_2,k_3$ 梯度，&lt;/p&gt;
$$
\begin{aligned}
\nabla_\theta k_1 &amp;= \nabla_\theta\left[-\log \frac{\pi_{ref}}{\pi_\theta}\right] = s_\theta\\
\nabla_\theta k_2 &amp;= \nabla_\theta\left[\frac12\left(\log \frac{\pi_{ref}}{\pi_\theta}\right)^2\right] = -\log \frac{\pi_{ref}}{\pi_\theta}s_\theta\\
\nabla_\theta k_3 &amp;= \nabla_\theta\left[\frac{\pi_{ref}}{\pi_\theta}-1- \log \frac{\pi_{ref}}{\pi_\theta}\right] = \left(1 - \frac{\pi_{ref}}{\pi_\theta}\right)s_\theta
\end{aligned}
$$&lt;p&gt;此时对应的梯度的期望为&lt;/p&gt;
$$
\begin{aligned}
\mathbb{E}_{\pi_{\theta}}[\nabla_\theta k_1] &amp;= \mathbb{E}_{\pi_{\theta}}[s_\theta]=0\\
\mathbb{E}_{\pi_{\theta}}[\nabla_\theta k_2] &amp;= \mathbb{E}_{\pi_{\theta}}\left[-\log \frac{\pi_{ref}}{\pi_\theta}s_\theta\right]=\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})\\
\mathbb{E}_{\pi_{\theta}}[\nabla_\theta k_3] &amp;= \mathbb{E}_{\pi_{\theta}}\left[\left(1 - \frac{\pi_{ref}}{\pi_\theta}\right)s_\theta\right]=\nabla_\theta D_{KL}(\pi_{ref}\parallel \pi_\theta)\\
\end{aligned}
$$&lt;p&gt;也就是说，$k_1$ 估计的梯度的期望为 0，对整体训练没有任何帮助，$k_3$ 估计的梯度的期望等价于优化 forward KL, **只有 $k_2$ 估计的梯度的期望等价于优化 reverse KL.&lt;/p&gt;
&lt;hr&gt;
&lt;p&gt;在实际代码实现的时候，KL divergence 有两种不同的实现形式：&lt;/p&gt;
&lt;p&gt;第一种是根据定义将 KL divergence 作为损失函数的一部分，此时我们的 KL divergence 参与反向传播，对应的实现方式如下&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt;1
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-fallback" data-lang="fallback"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;loss = -advantage * log_prob + beta * kl
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;p&gt;第二种是只调整 reward, 而不参与反向传播（通过 $\mathrm{sg}(\cdot)$ 实现），对应的实现方式如下所示&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt;1
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-fallback" data-lang="fallback"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;shaped_reward = reward - beta * kl.detach()
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;p&gt;这两者对于模型的训练影响很大，下面我们分别来进行介绍&lt;/p&gt;
&lt;h3 id="kl-as-loss"&gt;&lt;a href="#kl-as-loss" class="header-anchor"&gt;&lt;/a&gt;KL as Loss
&lt;/h3&gt;&lt;p&gt;为了统一 on-policy 和 off-policy 两种形式，我们使用一个统一的表达形式，即&lt;/p&gt;
$$
L=\rho k_i
$$&lt;p&gt;此时对应的 RL 目标函数为&lt;/p&gt;
$$
\mathcal{J}_2(\theta) = \mathbb{E}_{\tau\sim \pi_\theta}\left[\sum_{t=0}^T\gamma^tr(s_t,a_t)\right] - \beta\rho k_i(\pi_\theta, \pi_{ref})
$$&lt;p&gt;这里&lt;/p&gt;
$$
\rho = \frac{\pi_\theta}{\mathrm{sg}(\pi_{old})}
$$&lt;p&gt;是 importance weight,&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;当算法为 on-policy 时，$\pi_\theta=\pi_{old}$, $\rho\equiv1$.&lt;/li&gt;
&lt;li&gt;当算法为 off-policy 时，$\rho=\pi_\theta/\pi_{old}$, $\nabla_\theta \rho=\rho s_\theta$.&lt;/li&gt;
&lt;/ol&gt;
&lt;p&gt;通过这种方式，我们使得参数分布本身不会对梯度计算产生影响，从而使得对期望进行求导和对导数求期望相等，即&lt;/p&gt;
$$
\nabla_\theta\mathbb{E}_{\pi_{old}}[k] = \int \pi_{old}(x)\nabla_\theta kdx= \mathbb{E}_{\pi_{old}}[\nabla_\theta k]
$$&lt;p&gt;接下来我们来计算对应估计的梯度的期望，即 $\mathbb{E}[\nabla_\theta(\rho k_i)]$, 首先我们计算对应的梯度&lt;/p&gt;
$$
\begin{aligned}
\nabla_\theta (\rho k_1) &amp;= \rho s_\theta k_1+r\rho_\theta=\rho s_\theta(k_1+1)\\
\nabla_\theta (\rho k_2) &amp;= \rho s_\theta k_2+\rho\left(-\log \frac{\pi_{ref}}{\pi_\theta}s_\theta\right)=\rho s_\theta(k_1+k_2)\\
\nabla_\theta (\rho k_3) &amp;= \rho s_\theta k_3+\rho\left(1 - \frac{\pi_{ref}}{\pi_\theta}\right)s_\theta=\rho s_\theta\left(k_3+1-\frac{\pi_{ref}}{\pi_\theta}\right)=\rho s_\theta k_1
\end{aligned}
$$&lt;p&gt;注意到 $\mathbb{E}_{\pi_{old}} [\rho k_i]=\mathbb{E}_{\pi_{\theta}}[k_i]$ 以及 $\mathbb{E}_{\pi_{\theta}}[s_\theta]=0$, 我们对上述梯度求期望得到&lt;/p&gt;
$$
\begin{aligned}
\mathbb{E}_{\pi_{old}}[\nabla_\theta (\rho k_1)] &amp;= \mathbb{E}_{\pi_{old}}[\rho s_\theta(k_1+1)]=\mathbb{E}_{\pi_{\theta}}[s_\theta k_1]=\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})\\
\mathbb{E}_{\pi_{old}}[\nabla_\theta (\rho k_2)] &amp;= \mathbb{E}_{\pi_{old}}[\rho s_\theta(k_1+k_2)]=\nabla_\theta \mathbb{E}_{\pi_\theta}[k_2]\\
\mathbb{E}_{\pi_{old}}[\nabla_\theta (\rho k_3)] &amp;= \mathbb{E}_{\pi_{old}}[\rho s_\theta k_1]=\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})
\end{aligned}
$$&lt;p&gt;这里在计算 $\mathbb{E}_{\pi_{old}}[\nabla_\theta (\rho k_2)]$ 时，我们使用了 Leibniz 乘法法则：&lt;/p&gt;
$$
\mathbb{E}_{\pi_{old}}[\rho s_\theta(k_1+k_2)]= \mathbb{E}_{\pi_{\theta}}[s_\theta k_2]+\mathbb{E}_{\pi_{\theta}}[\nabla_\theta k_2]=\nabla_\theta\mathbb{E}_{\pi_{\theta}}[k_2]
$$&lt;p&gt;可以看到，$\rho k_1$ 和 $\rho k_3$ 都满足梯度与期望的可交换性，而 $\rho k_2$ 不满足，为了解决这个问题，我们可以使用 stop gradient, 即 $\mathrm{sg}(\rho)l_2$, 此时，我们有&lt;/p&gt;
$$
\nabla_\theta(\mathrm{sg}(\rho) k_2) = \mathrm{sg}(\rho)\nabla_\theta k_2 = \rho s_\theta k_1
$$&lt;p&gt;对其求期望有&lt;/p&gt;
$$
\mathbb{E}_{\pi_{old}}[\nabla_\theta(\mathrm{sg}(\rho) k_2)] = \mathbb{E}_{\pi_{old}}[\rho s_\theta k_1] = \mathbb{E}_{\pi_{\theta}}[s_\theta k_1]=\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})
$$&lt;p&gt;我们将如上结果总结为下表&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Loss&lt;/th&gt;
&lt;th&gt;gradient&lt;/th&gt;
&lt;th&gt;expected gradient&lt;/th&gt;
&lt;th&gt;objective&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;$\rho k_1$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta (k_1+1)$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\rho k_2$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta (k_1+k_2)$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta\mathbb{E}_{\pi_{\theta}}[k_2]$&lt;/td&gt;
&lt;td&gt;f-divergence&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\mathrm{sg}(\rho) k_2$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta k_1$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;$\rho k_3$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta k_1$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;p&gt;接下来，我们就可以分析在 on-policy 和 off-policy 场景下分析不同 estimator 的性质了。&lt;/p&gt;
&lt;p&gt;如果说，我们显式加入 $\rho$, 则根据上表我们可以使用上表的 $\rho k_1$, $\mathrm{sg}(\rho) k_2$ 以及 $\rho k_3$ 都可以作为损失函数的代替。&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;注
实际上 on-policy 场景下使用 $k_2$ 也有用的原因在于 $\nabla_\theta k_2=s_\theta k_1$, 也就是 $k_2$ 和 $\rho k_3$ 的梯度相同，其本质上是一个等效梯度。但是其收敛得到的 policy 与 target optimal policy 不同&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;接下来，我们来分析一下 $\rho k_1, \mathrm{sg}(\rho)k_2, \rho k_3$ 这三种估计的梯度的 variance, 为了避免混淆，【2】使用了 &amp;ldquo;projection variance in any direction&amp;rdquo; 的概念，即任意取一个向量 $u$, 然后计算 $\rho k_1$ 和后两者之间对应的 variance 的差（由于 $\mathrm{sg}(\rho)k_2$ 的梯度与 $\rho k_3$ 相同，因此这里我们仅计算 $\rho k_3$），得到:&lt;/p&gt;
$$
\begin{aligned}
\mathrm{var}[\nabla_\theta (\rho k_1)^Tu] - \mathrm{var}[\nabla_\theta (\rho k_3)^Tu] &amp;= (\mathbb{E}_{\pi_{old}}[(\nabla_\theta (\rho k_1)^Tu)^2] -\mathbb{E}_{\pi_{old}}^2[\nabla_\theta (\rho k_1)^Tu] ) - (\mathbb{E}_{\pi_{old}}[(\nabla_\theta (\rho k_3)^Tu)^2] -\mathbb{E}_{\pi_{old}}^2[\nabla_\theta (\rho k_3)^Tu] ) \\
&amp;= \mathbb{E}_{\pi_{old}}[(\nabla_\theta (\rho k_1)^Tu)^2] - \mathbb{E}_{\pi_{old}}[(\nabla_\theta (\rho k_3)^Tu)^2]\\
&amp;= \mathbb{E}_{\pi_{old}}[\rho(x)^2(s(\theta)(x)^Tu)^2(2k_1(x)+1)]
\end{aligned}
$$&lt;p&gt;当 $\pi_\theta$ 和 $\pi_{ref}$ 比较接近时，我们有&lt;/p&gt;
$$
\frac{\pi_{ref}(x)}{\pi_\theta(x)} = 1+\epsilon(x), \text{ where } |\epsilon(x)| &lt;&lt; 1
$$&lt;p&gt;此时&lt;/p&gt;
$$
2k_1(x) + 1 = 1-2\log(1+\epsilon(x))\approx 1-2\epsilon(x) \geq 0
$$&lt;p&gt;从而我们有&lt;/p&gt;
$$
\boxed{\mathrm{var}[\nabla_\theta (\rho k_1)]\geq \mathrm{var}[\nabla_\theta (\rho k_3)]=\mathrm{var}[\nabla_\theta (\mathrm{sg}(\rho)k_2)]}
$$&lt;p&gt;即当 $\pi_\theta$ 和 $\pi_{ref}$ 比较接近时，$\rho k_3$ 的 variance 比 $\rho k_1$ 更小，这是由于 $\rho s_\theta (k_1+1)$ 额外包含了一个 期望为零的项，这导致了其 variance 比较高。在 &lt;a class="link" href="https://maosong.website/p/notes-on-deepseek-v3.2/" target="_blank" rel="noopener"
&gt;DeepSeek-V3.2&lt;/a&gt; 中，作者就使用了 $\rho k_3$ 来降低梯度的 variance, 提高训练的稳定性。&lt;/p&gt;
&lt;p&gt;【3】将相关的估计总结为了下表的形式&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Type&lt;/th&gt;
&lt;th&gt;Loss&lt;/th&gt;
&lt;th&gt;Gradient&lt;/th&gt;
&lt;th&gt;Expected gradient&lt;/th&gt;
&lt;th&gt;Objective&lt;/th&gt;
&lt;th&gt;Biased&lt;/th&gt;
&lt;th&gt;Variance&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho k_1$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta (k_1+1)$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;unbiased&lt;/td&gt;
&lt;td&gt;high&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho k_2$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta (k_1+k_2)$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta\mathbb{E}_{\pi_{\theta}}[k_2]$&lt;/td&gt;
&lt;td&gt;f-divergence&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\mathrm{sg}(\rho) k_2$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta k_1$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;unbiased&lt;/td&gt;
&lt;td&gt;low&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho k_3$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta k_1$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;unbiased&lt;/td&gt;
&lt;td&gt;low&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;p&gt;【3】还强调了一点就是我们的损失函数必须显式包含 $\rho$, 在 on-policy 场景下，虽然 $\rho\equiv1$, 但是在反向传播时我们通过 $\nabla_\theta \rho=s_\theta$ 保留了采样信息从而避免了梯度估计期望的错配问题。&lt;/p&gt;
&lt;p&gt;对于 $\rho k_1$ variance 比较高的特点，我们还可以采用 variance reduction 的方法来降低不同估计的 variance. 【TODO】&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;analytic gradient&lt;/strong&gt;
当 action space 有限时，我们还可以使用解析梯度【TODO】&lt;/p&gt;
&lt;h3 id="as-a-reward-reshaping-item"&gt;&lt;a href="#as-a-reward-reshaping-item" class="header-anchor"&gt;&lt;/a&gt;As a Reward Reshaping Item
&lt;/h3&gt;&lt;p&gt;接下来我们来探究一下第二种形式，即 KL divergence 只影响最终的 reward, 而不参与反向传播。对应的代理目标函数形式为&lt;/p&gt;
$$
\mathcal{J}_3(\theta) = \mathbb{E}_{\tau\sim \pi_\theta}\left[R\right] - \beta\ \mathrm{sg}(k_i(\pi_\theta, \pi_{ref}))
$$&lt;p&gt;这里 $R=\sum_{t=0}^T\gamma^tr(s_t,a_t)$ 为 accumulative reward&lt;/p&gt;
&lt;p&gt;首先，基于前面分析，我们可以得到原始目标函数的梯度为&lt;/p&gt;
$$
\begin{aligned}
\nabla_\theta \mathcal{J}(\theta) &amp;= \nabla_\theta\mathbb{E}_{\pi_\theta}\left[R\right] - \beta \nabla_\theta D_{KL}(\pi_\theta, \pi_{ref})\\
&amp;= \mathbb{E}_{\pi_\theta}\left[s_\theta R\right]-\beta \mathbb{E}_{\pi_\theta}\left[s_\theta\log \frac{\pi_\theta}{\pi_{ref}}\right]\\
&amp;= \mathbb{E}_{\pi_\theta}\left[s_\theta(R-\beta k_1) \right]
\end{aligned}
$$&lt;p&gt;代理目标函数的梯度为&lt;/p&gt;
$$
\nabla_\theta \mathcal{J}_3(\theta) = \mathbb{E}_{\pi_\theta}\left[s_\theta(R-\beta k_i) \right]
$$&lt;p&gt;显然，当我们使用 $k_1$ 时，我们有 $\nabla_\theta \mathcal{J}(\theta)=\nabla_\theta \mathcal{J}_3(\theta)$.&lt;/p&gt;
&lt;p&gt;当我们使用 $k_2$ 时，带入 $k_2$ 表达式易知 $\nabla_\theta \mathcal{J}_3(\theta)\neq \nabla_\theta \mathcal{J}(\theta)$,&lt;/p&gt;
&lt;p&gt;当我们使用 $k_3$ 时，&lt;/p&gt;
$$
\begin{aligned}
\mathbb{E}_{\pi_\theta}\left[s_\theta k_i \right] &amp;= \mathbb{E}_{\pi_\theta}\left[s_\theta \left(\frac{\pi_{ref}}{\pi_\theta}-1- \log \frac{\pi_{ref}}{\pi_\theta} \right)\right]\\
&amp;= \mathbb{E}_{\pi_\theta}\left[s_\theta \frac{\pi_{ref}}{\pi_\theta}\right] - \mathbb{E}_{\pi_\theta}\left[s_\theta \right] - \mathbb{E}_{\pi_\theta}\left[s_\theta \log \frac{\pi_{ref}}{\pi_\theta} \right]\\
&amp;=s_\theta k_1 -\nabla_\theta D_{KL}(\pi_{ref}\parallel \pi_\theta)
\end{aligned}
$$&lt;p&gt;此时，$\nabla_\theta \mathcal{J}_3(\theta)\neq \nabla_\theta \mathcal{J}(\theta)$. 因此，&lt;strong&gt;在 on-policy 场景下，只有 $k_1$ 对应的梯度是无偏的&lt;/strong&gt;&lt;/p&gt;
&lt;p&gt;在 off-policy 场景下，由于 Off-policy 只影响 $R$ 的计算，因此原始目标函数和代理目标函数的梯度仍然保持不变，on-policy 场景的结论也适用。&lt;/p&gt;
&lt;p&gt;总之，&lt;strong&gt;当我们将 KL divergence 作为 reward reshaping item 时，只有 $k_1$ 产生的梯度是无偏的。&lt;/strong&gt;&lt;/p&gt;
&lt;h3 id="comparison-of-two-paradigms"&gt;&lt;a href="#comparison-of-two-paradigms" class="header-anchor"&gt;&lt;/a&gt;Comparison of Two Paradigms
&lt;/h3&gt;&lt;p&gt;接下来我们来比较一下 KL divergence 作为 loss 和 reward shaping item 的异同之处。首先，两者对于梯度的贡献分别为&lt;/p&gt;
$$
\begin{align}
&amp;\rho s_\theta k_1\tag{loss}\\
&amp; \mathbb{E}_{\pi_{old}}[\rho s_\theta k_1]\tag{reward shaping}
\end{align}
$$&lt;p&gt;即两者在期望上时一致的。但是两者也存在不一致的地方，即 KL divergence 作为 loss 时不会影响 $R$, 而作为 reward shaping item 时会影响。因此这就导致两者的优化方向不一致。&lt;/p&gt;
&lt;h3 id="experiments"&gt;&lt;a href="#experiments" class="header-anchor"&gt;&lt;/a&gt;Experiments
&lt;/h3&gt;&lt;p&gt;首先，我们来验证前面的结论，我们构造一个包含 $100$ 个 arms 的 multi-arm bandits, 然后令&lt;/p&gt;
$$
\pi_{ref}=\epsilon_1, \pi= \epsilon_1+\epsilon_2
$$&lt;p&gt;其中 $\epsilon_1,\epsilon_2\sim\mathcal{N}(0,1)$, 我们实验 100 次然后取平均值，然后分别计算 estimator 与真实 KL divergence 之间的 MSE 和 estimator gradient 与真实 kl divergence gradient 的 RMSE, 结果如下图所示&lt;/p&gt;
&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/notes-on-kl-divergence/KL-divergence-estimator-gradient-bias.png"
width="1389"
height="489"
loading="lazy"
alt="bias of KL divergence estimators and their gradients"
class="gallery-image"
data-flex-grow="284"
data-flex-basis="681px"
&gt;&lt;/p&gt;
&lt;p&gt;可以看到，这验证了我们之前分析的结论，即 $k_1$ 和 $k_3$ 是无偏估计，而在计算梯度时，只有 $k_2$ 梯度的期望与真实 KL divergence 的梯度相同。&lt;/p&gt;
&lt;h3 id="overview"&gt;&lt;a href="#overview" class="header-anchor"&gt;&lt;/a&gt;Overview
&lt;/h3&gt;&lt;p&gt;我们在本节总结前面的分析，如下表所示&lt;/p&gt;
&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Type&lt;/th&gt;
&lt;th&gt;Loss&lt;/th&gt;
&lt;th&gt;Gradient&lt;/th&gt;
&lt;th&gt;Expected gradient&lt;/th&gt;
&lt;th&gt;Objective&lt;/th&gt;
&lt;th&gt;Biased&lt;/th&gt;
&lt;th&gt;Variance&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;on-policy&lt;/td&gt;
&lt;td&gt;$k_1$&lt;/td&gt;
&lt;td&gt;$s_\theta$&lt;/td&gt;
&lt;td&gt;$0$&lt;/td&gt;
&lt;td&gt;constants&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on-policy&lt;/td&gt;
&lt;td&gt;$k_2$&lt;/td&gt;
&lt;td&gt;$-\log r s_\theta$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;unbiased&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on-policy&lt;/td&gt;
&lt;td&gt;$k_3$&lt;/td&gt;
&lt;td&gt;$(1-r)s_\theta$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_{ref}\parallel \pi_\theta)$&lt;/td&gt;
&lt;td&gt;forward KL&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho k_1$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta (k_1+1)$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;unbiased&lt;/td&gt;
&lt;td&gt;high&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho k_2$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta (k_1+k_2)$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta\mathbb{E}_{\pi_{\theta}}[k_2]$&lt;/td&gt;
&lt;td&gt;f-divergence&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\mathrm{sg}(\rho) k_2$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta k_1$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;unbiased&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;low&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho k_3$&lt;/td&gt;
&lt;td&gt;$\rho s_\theta k_1$&lt;/td&gt;
&lt;td&gt;$\nabla_\theta D_{KL}(\pi_\theta\parallel \pi_{ref})$&lt;/td&gt;
&lt;td&gt;reverse KL&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;unbiased&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;low&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho\mathrm{sg}(k_1)$&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;unbiased&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho \mathrm{sg}(k_2)$&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;on/off-policy&lt;/td&gt;
&lt;td&gt;$\rho \mathrm{sg}(k_3)$&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;td&gt;biased&lt;/td&gt;
&lt;td&gt;-&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;
&lt;h2 id="conclusion"&gt;&lt;a href="#conclusion" class="header-anchor"&gt;&lt;/a&gt;Conclusion
&lt;/h2&gt;&lt;p&gt;在本文中，我们详细介绍了 KL-divergence 的基本性质，相关估计方法以及在机器学习特别是 RL 领域中的应用。最终结论为：&lt;/p&gt;
&lt;ol&gt;
&lt;li&gt;如果希望稳定可控，则将 KL divergence 作为 loss item; 如果希望更灵活，与奖励信号结合的话，则将其作为 reward shaping item.&lt;/li&gt;
&lt;li&gt;使用 KL divergence 作为 loss item 时，on-policy 场景下使用 $k_2$ 近似 KL divergence 效果最好；off-policy 场景下，使用 $\mathrm{sg}(\rho)k_2, \rho k_3$ 效果最好&lt;/li&gt;
&lt;li&gt;使用 KL divergence 作为 reward shaping item 时，$k_1$ 的效果最好&lt;/li&gt;
&lt;/ol&gt;
&lt;h2 id="references"&gt;&lt;a href="#references" class="header-anchor"&gt;&lt;/a&gt;References
&lt;/h2&gt;&lt;ul&gt;
&lt;li&gt;&lt;a class="link" href="http://joschu.net/blog/kl-approx.html" target="_blank" rel="noopener"
&gt;Approximating KL Divergence&lt;/a&gt;&lt;/li&gt;
&lt;li&gt;&lt;a class="link" href="https://dibyaghosh.com/blog/probability/kldivergence/" target="_blank" rel="noopener"
&gt;KL Divergence for Machine Learning&lt;/a&gt;&lt;/li&gt;
&lt;li&gt;&lt;a class="link" href="https://xihuai18.github.io/reinforcement-learning/2025/12/01/kl-estimators-en.html" target="_blank" rel="noopener"
&gt;Understanding KL Divergence Estimators in RL: From Value Approximation to Gradient Estimation&lt;/a&gt;&lt;/li&gt;
&lt;li&gt;&lt;a class="link" href="http://arxiv.org/abs/2506.09477" target="_blank" rel="noopener"
&gt;On a few pitfalls in KL divergence gradient estimation for RL&lt;/a&gt;&lt;/li&gt;
&lt;/ul&gt;</description></item><item><title>Notes on Softmax</title><link>https://jiangyigithub.github.io/ai.github.io/p/notes-on-softmax/</link><pubDate>Sat, 27 Dec 2025 16:39:53 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/notes-on-softmax/</guid><description>&lt;h2 id="introduction"&gt;&lt;a href="#introduction" class="header-anchor"&gt;&lt;/a&gt;Introduction
&lt;/h2&gt;&lt;p&gt;softmax 函数用于将 $K$ 个实数转换为一个 $K$ 维概率分布。其具体做法是先对所有元素指数化，即求 $e^x$, 然后每个元素除以所有指数的和。即&lt;/p&gt;
$$
\begin{aligned}
\mathrm{softmax}:\mathbb{R}^K&amp;\to (0,1)^K\\
\mathrm{softmax}(\mathbf{z}) &amp;=\left[\frac{e^{z_1}}{\sum_{j=1}^Ke^{z_j}},\dots,\frac{e^{z_K}}{\sum_{j=1}^Ke^{z_j}}\right]
\end{aligned}
$$&lt;h2 id="analysis"&gt;&lt;a href="#analysis" class="header-anchor"&gt;&lt;/a&gt;Analysis
&lt;/h2&gt;&lt;h3 id="properties"&gt;&lt;a href="#properties" class="header-anchor"&gt;&lt;/a&gt;Properties
&lt;/h3&gt;&lt;p&gt;softmax 的第一个性质是 shift invariance, 即&lt;/p&gt;
$$
\mathrm{softmax}(\mathbf{z}+c) = \mathrm{softmax}(\mathbf{z})
$$&lt;p&gt;证明比较容易：&lt;/p&gt;
$$
\mathrm{softmax}(\mathbf{z}+c)_i = \frac{e^{z_i+c}}{\sum_{j=1}^Ke^{z_j+c}} = \frac{e^ce^{z_i}}{e^c\sum_{j=1}^Ke^{z_j}} = \frac{e^{z_i}}{\sum_{j=1}^Ke^{z_j}}=\mathrm{softmax}(\mathbf{z})_i,\ i=1,\dots,K
$$&lt;h3 id="gradient"&gt;&lt;a href="#gradient" class="header-anchor"&gt;&lt;/a&gt;Gradient
&lt;/h3&gt;&lt;p&gt;向量输入下 Softmax 函数的 Jacobian 矩阵推导&lt;/p&gt;
&lt;p&gt;设输入为向量 $\mathbf{z} = [z_1, z_2, \dots, z_d]^\top \in \mathbb{R}^d$，Softmax 函数的输出为向量 $\mathbf{a} = [a_1, a_2, \dots, a_d]^\top \in \mathbb{R}^d$，其中每个元素定义为：&lt;/p&gt;
$$
a_j = \text{softmax}(\mathbf{z})_j = \frac{e^{z_j}}{\sum_{k=1}^d e^{z_k}}
$$&lt;p&gt;记分母（归一化因子）为 $S = \sum_{k=1}^d e^{z_k}$，则 $a_j = e^{z_j}/S$.&lt;/p&gt;
&lt;p&gt;我们分两种情况计算 $\frac{\partial a_j}{\partial z_k}$：&lt;/p&gt;
&lt;p&gt;当 $j = k$ 时， 此时求 $a_j$ 对自身输入 $z_j$ 的偏导数：&lt;/p&gt;
$$
\frac{\partial a_j}{\partial z_j} = \frac{\partial}{\partial z_j} \left( \frac{e^{z_j}}{S} \right) = \frac{e^{z_j}S-e^{z_j}e^{z_j}}{S^2}=\frac{e^{z_j}}{S}\left(1-\frac{e^{z_j}}{S}\right)=a_j(1-a_j)
$$&lt;p&gt;当 $j \neq k$ 时， 此时求 $a_j$ 对输入 $z_k$ 的偏导数有：&lt;/p&gt;
$$
\frac{\partial a_j}{\partial z_k} = \frac{\partial}{\partial z_k} \left( \frac{e^{z_j}}{S} \right) = \frac{0\cdot S-e^{z_j}e^{z_k}}{S^2}=-\frac{e^{z_j}e^{z_k}}{S}=-a_ja_j
$$&lt;p&gt;综合以上两种情况，Jacobian 矩阵 $\mathbf{J}$ 可表示为：&lt;/p&gt;
$$
\mathbf{J} = \text{diag}(\mathbf{a}) - \mathbf{a} \mathbf{a}^\top
$$&lt;h2 id="interpretation"&gt;&lt;a href="#interpretation" class="header-anchor"&gt;&lt;/a&gt;Interpretation
&lt;/h2&gt;&lt;h3 id="soft-argmax"&gt;&lt;a href="#soft-argmax" class="header-anchor"&gt;&lt;/a&gt;Soft Argmax
&lt;/h3&gt;&lt;p&gt;softmax 是 argmax 的 smooth approximation, 所以实际上 softmax 指的是 “soft argmax&amp;quot;. 为了证明这一点，我们首先定义如下函数&lt;/p&gt;
$$
\mathrm{softmax}(\mathbf{z};\tau) =\mathrm{softmax}(\mathbf{z}/\tau)=\left[\frac{e^{z_1/\tau}}{\sum_{j=1}^Ke^{z_j/\tau}},\dots,\frac{e^{z_K/\tau}}{\sum_{j=1}^Ke^{z_j/\tau}}\right]
$$&lt;p&gt;易知， $\mathrm{softmax}(\mathbf{z})=\mathrm{softmax}(\mathbf{z};1)$. 并且，$\mathrm{softmax}$ 还是一个光滑函数&lt;/p&gt;
&lt;p&gt;我们定义 smooth approximation 为&lt;/p&gt;
&lt;blockquote&gt;
&lt;p&gt;Definition
如果 $\lim_{\tau\to0^+}\mathrm{softmax}(\mathbf{z};\tau)=\mathbb{1}_{\arg\max(\mathbf{z})}$, 则我们说 $\mathrm{softmax}(\cdot;\tau)$ 是 $\arg\max$ 的光滑近似，特别地，$\mathrm{softmax}(\cdot)$ 是 $\arg\max$ 的光滑近似。
这里 $\arg\max(\mathbf{z})=\arg\max_k z_k$ 是最大值的索引， $\mathbb{1}\in\{0,1\}^K$ 是示性函数 (indicator function), 即 $\mathbb{1}_{\arg\max(\mathbf{z})}[i]=1$ 当且仅当 $z_i=\max_jz_j$.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;我们下面来进行证明。我们不妨假设最大值唯一，其 index 为 $m$, 即 $z_m = \max_i z_i$. 由前面的性质，我们有：&lt;/p&gt;
$$
\mathrm{softmax}(\mathbf{z};\tau) = \mathrm{softmax}(\mathbf{z}-z_m;\tau) =\left[\frac{e^{(z_1-z_m)/\tau}}{\sum_{j=1}^Ke^{(z_j-z_m)/\tau}},\dots,\frac{e^{(z_K-z_m)/\tau}}{\sum_{j=1}^Ke^{(z_j-z_m)/\tau}}\right]
$$&lt;p&gt;此时，我们有&lt;/p&gt;
$$
\lim_{\tau\to0^+}\mathrm{softmax}(\mathbf{z};\tau)_i = \begin{cases}
1, &amp;\text{if }i = m\\
0, &amp;\text{otherwise}
\end{cases}
$$&lt;p&gt;当最大值不唯一的时候，我们记 $\mathcal{I} = \{i\in[K]\mid z_i=\max_j z_j\}$, 与上面方法类似，最终 $\mathrm{softmax}(\cdot;\tau)$ 的结果为&lt;/p&gt;
$$
\lim_{\tau\to0^+}\mathrm{softmax}(\mathbf{z};\tau)_i = \begin{cases}
1/|\mathcal{I}|, &amp;\text{if }i \in \mathcal{I}\\
0, &amp;\text{otherwise}
\end{cases}
$$&lt;p&gt;因此，我们就证明了 softmax 是 argmax 函数的 smooth approximation.&lt;/p&gt;
&lt;h3 id="statistical-mechanics"&gt;&lt;a href="#statistical-mechanics" class="header-anchor"&gt;&lt;/a&gt;Statistical Mechanics
&lt;/h3&gt;&lt;h3 id="temperature"&gt;&lt;a href="#temperature" class="header-anchor"&gt;&lt;/a&gt;Temperature
&lt;/h3&gt;&lt;p&gt;我们前面介绍了 $\mathrm{softmax}(\mathbf{z};\tau)$ 函数，这里的 $\tau$ 实际上被称为温度 (temperature), 它控制了输入的 variance, $T$ 越大，输入的 variance 越低，输出就倾向于均匀分布，而 $T$ 越小，则说明输入的 variance 越高，输出就倾向于 one-hot 分布。&lt;/p&gt;
&lt;p&gt;我们前面已经证明了后者，现在我们来证明一下前者，证明思路也很简单，$T\to+\infty$ 时，$e^{x/T}\to 1$, 因而&lt;/p&gt;
$$
\lim_{\tau\to+\infty}\mathrm{softmax}(\mathbf{z};\tau)_i =\frac1K,\ i=1,\dots,K
$$&lt;p&gt;下面是可视化的代码以及结果&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt; 1
&lt;/span&gt;&lt;span class="lnt"&gt; 2
&lt;/span&gt;&lt;span class="lnt"&gt; 3
&lt;/span&gt;&lt;span class="lnt"&gt; 4
&lt;/span&gt;&lt;span class="lnt"&gt; 5
&lt;/span&gt;&lt;span class="lnt"&gt; 6
&lt;/span&gt;&lt;span class="lnt"&gt; 7
&lt;/span&gt;&lt;span class="lnt"&gt; 8
&lt;/span&gt;&lt;span class="lnt"&gt; 9
&lt;/span&gt;&lt;span class="lnt"&gt;10
&lt;/span&gt;&lt;span class="lnt"&gt;11
&lt;/span&gt;&lt;span class="lnt"&gt;12
&lt;/span&gt;&lt;span class="lnt"&gt;13
&lt;/span&gt;&lt;span class="lnt"&gt;14
&lt;/span&gt;&lt;span class="lnt"&gt;15
&lt;/span&gt;&lt;span class="lnt"&gt;16
&lt;/span&gt;&lt;span class="lnt"&gt;17
&lt;/span&gt;&lt;span class="lnt"&gt;18
&lt;/span&gt;&lt;span class="lnt"&gt;19
&lt;/span&gt;&lt;span class="lnt"&gt;20
&lt;/span&gt;&lt;span class="lnt"&gt;21
&lt;/span&gt;&lt;span class="lnt"&gt;22
&lt;/span&gt;&lt;span class="lnt"&gt;23
&lt;/span&gt;&lt;span class="lnt"&gt;24
&lt;/span&gt;&lt;span class="lnt"&gt;25
&lt;/span&gt;&lt;span class="lnt"&gt;26
&lt;/span&gt;&lt;span class="lnt"&gt;27
&lt;/span&gt;&lt;span class="lnt"&gt;28
&lt;/span&gt;&lt;span class="lnt"&gt;29
&lt;/span&gt;&lt;span class="lnt"&gt;30
&lt;/span&gt;&lt;span class="lnt"&gt;31
&lt;/span&gt;&lt;span class="lnt"&gt;32
&lt;/span&gt;&lt;span class="lnt"&gt;33
&lt;/span&gt;&lt;span class="lnt"&gt;34
&lt;/span&gt;&lt;span class="lnt"&gt;35
&lt;/span&gt;&lt;span class="lnt"&gt;36
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-python" data-lang="python"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="kn"&gt;import&lt;/span&gt; &lt;span class="nn"&gt;numpy&lt;/span&gt; &lt;span class="k"&gt;as&lt;/span&gt; &lt;span class="nn"&gt;np&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="kn"&gt;import&lt;/span&gt; &lt;span class="nn"&gt;matplotlib.pyplot&lt;/span&gt; &lt;span class="k"&gt;as&lt;/span&gt; &lt;span class="nn"&gt;plt&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="kn"&gt;from&lt;/span&gt; &lt;span class="nn"&gt;scipy.interpolate&lt;/span&gt; &lt;span class="kn"&gt;import&lt;/span&gt; &lt;span class="n"&gt;make_interp_spline&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="k"&gt;def&lt;/span&gt; &lt;span class="nf"&gt;softmax&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x&lt;/span&gt;&lt;span class="p"&gt;):&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;e_x&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;np&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;exp&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="n"&gt;np&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;max&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;e_x&lt;/span&gt; &lt;span class="o"&gt;/&lt;/span&gt; &lt;span class="n"&gt;e_x&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;sum&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;num_elements&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;15&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;indices&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;np&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;arange&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;num_elements&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;np&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;linspace&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mf"&gt;3.5&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mf"&gt;3.5&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;num_elements&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;scales&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mf"&gt;0.01&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mf"&gt;0.1&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mf"&gt;1.0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mf"&gt;5.0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mf"&gt;10.0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mf"&gt;100.0&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;figure&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;figsize&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="mi"&gt;10&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mi"&gt;6&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="k"&gt;for&lt;/span&gt; &lt;span class="n"&gt;s&lt;/span&gt; &lt;span class="ow"&gt;in&lt;/span&gt; &lt;span class="n"&gt;scales&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;probs&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;softmax&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;*&lt;/span&gt; &lt;span class="n"&gt;s&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;x_smooth&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;np&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;linspace&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;indices&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;min&lt;/span&gt;&lt;span class="p"&gt;(),&lt;/span&gt; &lt;span class="n"&gt;indices&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;max&lt;/span&gt;&lt;span class="p"&gt;(),&lt;/span&gt; &lt;span class="mi"&gt;300&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;spl&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;make_interp_spline&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;indices&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;probs&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;k&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;3&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;y_smooth&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;np&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;clip&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;spl&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x_smooth&lt;/span&gt;&lt;span class="p"&gt;),&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="kc"&gt;None&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="c1"&gt;# Clip to ensure no negative artifacts&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;plot&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x_smooth&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;y_smooth&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;label&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Scale = &lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="n"&gt;s&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;linewidth&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;2&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;uniform_prob&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;1.0&lt;/span&gt; &lt;span class="o"&gt;/&lt;/span&gt; &lt;span class="nb"&gt;len&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;indices&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;axhline&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;y&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;uniform_prob&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;color&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;black&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;linestyle&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;:&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;alpha&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mf"&gt;0.6&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;label&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Uniform distribution&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;xticks&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;indices&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;xlabel&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Logit Index&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;fontsize&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;12&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;ylabel&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Softmax Probability&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;fontsize&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;12&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;title&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Impact of Variance Scaling on Softmax Distribution&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;fontsize&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;14&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;legend&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;title&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="s2"&gt;&amp;#34;Variance Scale&amp;#34;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;grid&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="kc"&gt;True&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;linestyle&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;--&amp;#39;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;alpha&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mf"&gt;0.5&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tight_layout&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="n"&gt;plt&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;show&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;p&gt;&lt;img src="https://jiangyigithub.github.io/ai.github.io/p/notes-on-softmax/softmax_impact_variance.png"
width="989"
height="590"
loading="lazy"
alt="impact of variance on softmax"
class="gallery-image"
data-flex-grow="167"
data-flex-basis="402px"
&gt;&lt;/p&gt;
&lt;p&gt;可以看到，当 variance 比较小的时候，输出的分布接近于均匀分布，而 variance 越大，输出的分布越接近 One-hot 分布。&lt;/p&gt;
&lt;p&gt;在 attention 的计算过程中，我们也有 softmax 函数，为了在 softmax 过程中避免 variance 的影响，现在会在计算 softmax 之前加入 normalization layer 来提前进行归一化。见 &lt;a class="link" href="https://maosong.website/p/notes-on-qk-norm/" target="_blank" rel="noopener"
&gt;QK-norm&lt;/a&gt;.&lt;/p&gt;
&lt;h2 id="algorithms"&gt;&lt;a href="#algorithms" class="header-anchor"&gt;&lt;/a&gt;Algorithms
&lt;/h2&gt;&lt;h3 id="implementation"&gt;&lt;a href="#implementation" class="header-anchor"&gt;&lt;/a&gt;Implementation
&lt;/h3&gt;&lt;p&gt;由于 $e^x$ 在实际计算时，非常容易溢出，因此在实现的时候，我们往往会考虑其数值稳定性。实际上，现在的 softmax 函数基本由 logsumexp 实现，logsumexp 函数定义如下&lt;/p&gt;
$$
\mathrm{logsumexp}(\mathbf{z}) = \log \left(\sum_{i=1}^K e^{z_i}\right)
$$&lt;p&gt;softmax 函数与 logsumexp 函数的关系如下&lt;/p&gt;
$$
\begin{aligned}
\mathrm{softmax}(\mathbf{z}) &amp;=\exp\log\left(\frac{e^{\mathbf{z}}}{\sum_{j=1}^Ke^{z_j}}\right)\\
&amp;= \exp\left(\mathbf{z} - \log\left(\sum_{i=1}^K e^{z_i}\right)\right)\\
&amp;= \exp(\mathbf{z} - \mathrm{logsumexp}(\mathbf{z}))
\end{aligned}
$$&lt;p&gt;考虑前面提到的 $e^x$ 数值溢出的问题，我们的输入会先经过 shift, 减掉最大值。此时我们有&lt;/p&gt;
$$
\mathrm{softmax}(\mathbf{z}) = \mathrm{softmax}(\mathbf{z}-c) = \exp((\mathbf{z}-c) - \mathrm{logsumexp}(\mathbf{z}-c))
$$&lt;p&gt;这里我们使用了前面推导出来的 shift invariance 性质。对应的代码实现如下：&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt;1
&lt;/span&gt;&lt;span class="lnt"&gt;2
&lt;/span&gt;&lt;span class="lnt"&gt;3
&lt;/span&gt;&lt;span class="lnt"&gt;4
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-python" data-lang="python"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="k"&gt;def&lt;/span&gt; &lt;span class="nf"&gt;softmax&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="n"&gt;torch&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Tensor&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="o"&gt;-&amp;gt;&lt;/span&gt; &lt;span class="n"&gt;torch&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Tensor&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;x&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;x&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="n"&gt;x&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;max&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;keepdim&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;True&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;values&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;log_sum_exp&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;torch&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;log&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;torch&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;sum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;torch&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;exp&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x&lt;/span&gt;&lt;span class="p"&gt;),&lt;/span&gt; &lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;keepdim&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;True&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;torch&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;exp&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;x&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="n"&gt;log_sum_exp&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;h3 id="gumbel-softmax-reparametrization-trick"&gt;&lt;a href="#gumbel-softmax-reparametrization-trick" class="header-anchor"&gt;&lt;/a&gt;Gumbel-softmax Reparametrization Trick
&lt;/h3&gt;&lt;p&gt;TODO&lt;/p&gt;
&lt;h3 id="online-softmax"&gt;&lt;a href="#online-softmax" class="header-anchor"&gt;&lt;/a&gt;Online Softmax
&lt;/h3&gt;&lt;p&gt;注意到我们在计算 softmax 时，需要加载 $\mathbf{z}$ 的全部信息，如果 $\mathbf{z}$ 非常大的话，会产生频繁的内存读写进而影响整体效率。因此 &lt;a class="link" href="https://maosong.website/p/notes-on-flashattention/" target="_blank" rel="noopener"
&gt;flash attention&lt;/a&gt; 中提出了 online softmax 算法来减少内存访问开销。&lt;/p&gt;
&lt;p&gt;其具体做法是假设我们的输入被分为若干个 block, 即 $\mathbf{z}=[\mathbf{z}^1;\dots,\mathbf{z}^n]\in\mathbb{R}^K$, 这里 $\mathbf{z}^i\in\mathbb{R}^{K/n}$ ($K\mod n=0$).&lt;/p&gt;
&lt;p&gt;对于 $\mathbf{z}\in\mathbb{R}^K$, flash attention 定义如下结果&lt;/p&gt;
$$
m(\mathbf{z}) = \max_i z_i,\ f(\mathbf{z}) = [e^{z_1-m(\mathbf{z})},\dots,e^{z_K-m(\mathbf{z})}], \ \ell(\mathbf{z})=\sum_if(z)_i, \ \mathrm{softmax}(\mathbf{z}) = \frac{f(\mathbf{z})}{\ell(\mathbf{z})}
$$&lt;p&gt;对于 $\mathbf{z}=[\mathbf{z}^1;\dots,\mathbf{z}^n]\in\mathbb{R}^K$, 我们现在的计算方式为&lt;/p&gt;
$$
\begin{aligned}
m_i(\mathbf{z}) &amp;= \max([\mathbf{z}^1;\dots;\mathbf{z}^i]) = \max(m_{i-1}(\mathbf{z}),m(\mathbf{z}^i))\\
\ell_i(\mathbf{z}) &amp;= \sum_{j=1}^if(\mathbf{z}^j) = \exp(m_{i-1}(\mathbf{z}) - m_i(\mathbf{z}))\ell(\mathbf{z}^{i-1}) + \exp(\mathbf{z}^i-m_i(\mathbf{z}))
\end{aligned}
$$&lt;p&gt;因此，如果我们额外记录 $m(x)$ 以及 $\ell(x)$ 这两个量，那么我们可以每次仅计算 softmax 的一个 block. 计算完毕之后，$m_i(\mathbf{z})$ 和 $\ell_i(\mathbf{z})$ 就分别代表了 global max 和 global denominator.&lt;/p&gt;
&lt;h2 id="conclusion"&gt;&lt;a href="#conclusion" class="header-anchor"&gt;&lt;/a&gt;Conclusion
&lt;/h2&gt;&lt;p&gt;我们回顾了机器学习中 softmax function 的基本定义与性质&lt;/p&gt;
&lt;h2 id="references"&gt;&lt;a href="#references" class="header-anchor"&gt;&lt;/a&gt;References
&lt;/h2&gt;&lt;ul&gt;
&lt;li&gt;&lt;a class="link" href="https://en.wikipedia.org/wiki/Softmax_function" target="_blank" rel="noopener"
&gt;Softmax function&lt;/a&gt;&lt;/li&gt;
&lt;li&gt;&lt;a class="link" href="https://blog.ando.ai/posts/softmax-to-the-max/" target="_blank" rel="noopener"
&gt;Softmax to the Max&lt;/a&gt;&lt;/li&gt;
&lt;li&gt;&lt;a class="link" href="https://arxiv.org/pdf/1805.02867" target="_blank" rel="noopener"
&gt;Online normalizer calculation for softmax&lt;/a&gt;&lt;/li&gt;
&lt;/ul&gt;</description></item><item><title>compression is intelligence</title><link>https://jiangyigithub.github.io/ai.github.io/p/compression-is-intelligence/</link><pubDate>Thu, 06 Mar 2025 17:57:51 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/compression-is-intelligence/</guid><description>&lt;p&gt;我们知道，基于decoder-only transformer的LLM的训练目标是最小化next-token-prediction loss，即给定sequence $x=(x_1,\dots, x_n)\in D$，我们的目标为求解以下优化问题&lt;/p&gt;
$$
\min_{\theta} -\sum_{x\in D}\log P_{\theta}(x_i|x_1,\dots,x_{i-1})
$$&lt;p&gt;这里 $\theta$ 就是我们的模型参数，$D$ 是我们的训练数据集。&lt;/p&gt;
&lt;p&gt;无数模型通过实际效果告诉我们，这个优化目标可以很好地训练出具有良好泛化能力的大语言模型。但是，我们的问题是，为什么这个优化目标可以训练出智能的模型？ 本文将从压缩即智能的角度来理解这个问题。&lt;/p&gt;
&lt;h1 id="压缩即智能"&gt;&lt;a href="#%e5%8e%8b%e7%bc%a9%e5%8d%b3%e6%99%ba%e8%83%bd" class="header-anchor"&gt;&lt;/a&gt;压缩即智能
&lt;/h1&gt;&lt;h2 id="一个例子"&gt;&lt;a href="#%e4%b8%80%e4%b8%aa%e4%be%8b%e5%ad%90" class="header-anchor"&gt;&lt;/a&gt;一个例子
&lt;/h2&gt;&lt;p&gt;我们首先来看一个简单的例子。给定如下三个0-1字符串：&lt;/p&gt;
&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt;1
&lt;/span&gt;&lt;span class="lnt"&gt;2
&lt;/span&gt;&lt;span class="lnt"&gt;3
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-fallback" data-lang="fallback"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;01010101010101010101
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;01001000100001000001
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;01101000101010100101
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;&lt;p&gt;我们该如何描述这三个字符串的规律？显然，第一个字符串最简单，它是&lt;code&gt;01&lt;/code&gt;字符串重复得到的结果；第二个字符串稍微复杂一些，它在每个&lt;code&gt;1&lt;/code&gt;之前插入重复次数的&lt;code&gt;0&lt;/code&gt;；第三个字符串则最复杂，它是我随手写的一个字符串，基本没有任何规律，因此，我们只能直接存储这个字符串。&lt;/p&gt;
&lt;p&gt;这个例子告诉我们，一个字符串的规律越简单，我们越容易描述它，因此，我们越容易压缩它。实际上，大语言模型做的也是类似的事情。它们的核心思想是，压缩即智能。&lt;/p&gt;
&lt;h2 id=""&gt;&lt;a href="#" class="header-anchor"&gt;&lt;/a&gt;
&lt;/h2&gt;&lt;h1 id="结论"&gt;&lt;a href="#%e7%bb%93%e8%ae%ba" class="header-anchor"&gt;&lt;/a&gt;结论
&lt;/h1&gt;&lt;p&gt;本文中，我们从压缩即智能的角度来理解大语言模型的原理。我们发现，大语言模型的next-token-prediction其实就是压缩。我们通过压缩让大语言模型学习到了语言中的规律，从而让模型具有了智能。&lt;/p&gt;
&lt;h1 id="参考文献"&gt;&lt;a href="#%e5%8f%82%e8%80%83%e6%96%87%e7%8c%ae" class="header-anchor"&gt;&lt;/a&gt;参考文献
&lt;/h1&gt;</description></item><item><title>1137. N-th Tribonacci Number</title><link>https://jiangyigithub.github.io/ai.github.io/p/1137.-n-th-tribonacci-number/</link><pubDate>Wed, 24 Apr 2024 18:53:26 +0800</pubDate><guid>https://jiangyigithub.github.io/ai.github.io/p/1137.-n-th-tribonacci-number/</guid><description>&lt;p&gt;Compute the $n$-th tribonacci number.&lt;/p&gt;
&lt;h1 id="intuition"&gt;&lt;a href="#intuition" class="header-anchor"&gt;&lt;/a&gt;Intuition
&lt;/h1&gt;&lt;p&gt;Same as compute the $n$-th fibonacci number, we use three numbers to remember the state.&lt;/p&gt;
&lt;h1 id="approach"&gt;&lt;a href="#approach" class="header-anchor"&gt;&lt;/a&gt;Approach
&lt;/h1&gt;&lt;p&gt;We use three numbers to represent $n-2$, $n-1$ and $n$-th tribonacci number respectively&lt;/p&gt;
&lt;h1 id="complexity"&gt;&lt;a href="#complexity" class="header-anchor"&gt;&lt;/a&gt;Complexity
&lt;/h1&gt;&lt;ul&gt;
&lt;li&gt;
&lt;p&gt;Time complexity:
&lt;/p&gt;
$$O(n)$$&lt;/li&gt;
&lt;li&gt;
&lt;p&gt;Space complexity:
&lt;/p&gt;
$$O(1)$$&lt;/li&gt;
&lt;/ul&gt;
&lt;h1 id="code"&gt;&lt;a href="#code" class="header-anchor"&gt;&lt;/a&gt;Code
&lt;/h1&gt;&lt;div class="highlight"&gt;&lt;div class="chroma"&gt;
&lt;table class="lntable"&gt;&lt;tr&gt;&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code&gt;&lt;span class="lnt"&gt; 1
&lt;/span&gt;&lt;span class="lnt"&gt; 2
&lt;/span&gt;&lt;span class="lnt"&gt; 3
&lt;/span&gt;&lt;span class="lnt"&gt; 4
&lt;/span&gt;&lt;span class="lnt"&gt; 5
&lt;/span&gt;&lt;span class="lnt"&gt; 6
&lt;/span&gt;&lt;span class="lnt"&gt; 7
&lt;/span&gt;&lt;span class="lnt"&gt; 8
&lt;/span&gt;&lt;span class="lnt"&gt; 9
&lt;/span&gt;&lt;span class="lnt"&gt;10
&lt;/span&gt;&lt;span class="lnt"&gt;11
&lt;/span&gt;&lt;span class="lnt"&gt;12
&lt;/span&gt;&lt;span class="lnt"&gt;13
&lt;/span&gt;&lt;span class="lnt"&gt;14
&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;
&lt;td class="lntd"&gt;
&lt;pre tabindex="0" class="chroma"&gt;&lt;code class="language-c++" data-lang="c++"&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="k"&gt;class&lt;/span&gt; &lt;span class="nc"&gt;Solution&lt;/span&gt; &lt;span class="p"&gt;{&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="k"&gt;public&lt;/span&gt;&lt;span class="o"&gt;:&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="kt"&gt;int&lt;/span&gt; &lt;span class="n"&gt;tribonacci&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="kt"&gt;int&lt;/span&gt; &lt;span class="n"&gt;n&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="p"&gt;{&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;vector&lt;/span&gt;&lt;span class="o"&gt;&amp;lt;&lt;/span&gt;&lt;span class="kt"&gt;int&lt;/span&gt;&lt;span class="o"&gt;&amp;gt;&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;{&lt;/span&gt;&lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;};&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;n&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;&lt;/span&gt; &lt;span class="mi"&gt;3&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="n"&gt;n&lt;/span&gt;&lt;span class="p"&gt;];&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="k"&gt;for&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="kt"&gt;int&lt;/span&gt; &lt;span class="n"&gt;i&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;2&lt;/span&gt;&lt;span class="p"&gt;;&lt;/span&gt; &lt;span class="n"&gt;i&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;&lt;/span&gt; &lt;span class="n"&gt;n&lt;/span&gt;&lt;span class="p"&gt;;&lt;/span&gt; &lt;span class="o"&gt;++&lt;/span&gt;&lt;span class="n"&gt;i&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="p"&gt;{&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="kt"&gt;int&lt;/span&gt; &lt;span class="n"&gt;temp&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;2&lt;/span&gt;&lt;span class="p"&gt;];&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;2&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt; &lt;span class="o"&gt;+=&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;];&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;];&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;temp&lt;/span&gt;&lt;span class="p"&gt;;&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="p"&gt;}&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;nums&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;2&lt;/span&gt;&lt;span class="p"&gt;];&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt; &lt;span class="p"&gt;}&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;span class="line"&gt;&lt;span class="cl"&gt;&lt;span class="p"&gt;};&lt;/span&gt;
&lt;/span&gt;&lt;/span&gt;&lt;/code&gt;&lt;/pre&gt;&lt;/td&gt;&lt;/tr&gt;&lt;/table&gt;
&lt;/div&gt;
&lt;/div&gt;</description></item></channel></rss>