TL;DR. We provide a hand-crafted derivation of two estimators for the diagonal of the Hessian matrix, which is widely used in modern adaptive optimizers. This post mainly follows https://arxiv.org/abs/2305.14342
<aside> 💡
Hutchinson's unbiased estimator.
$$ {\rm diag} (\nabla^2 \ell(\theta)) = \mathbb{E}_{u \sim \mathcal{N}(0, I_d)}[u \odot \nabla^2 \ell(\theta) u]. $$
</aside>
Proof.
Denote the entries of $u \odot \nabla^2 \ell(\theta) u$ as $\hat{\sigma}_i$,
$$ \begin{aligned} \hat{\sigma}i &= u_i \sum{j=1}^d \left(\nabla^2 \ell(\theta) \right){ij} u_j \\ &= \left(\nabla^2 \ell(\theta) \right){ii} u_i^2 + \sum_{j \neq i} (\nabla^2\ell(\theta))_{ij} u_i u_j \end{aligned} $$
Clearly,
$$ \begin{aligned} \mathbb{E}{u\sim \mathcal{N}(0, I_d)}[\hat{\sigma}i] &= \mathbb{E}{u_i\sim \mathcal{N}(0, 1)}\left[\left(\nabla^2 \ell(\theta) \right){ii} u_i^2\right]+ \mathbb{E}{u\sim \mathcal{N}(0, I_d)}\left[\sum{j \neq i} (\nabla^2\ell(\theta)){ij} u_i u_j\right] \\ &= \left(\nabla^2 \ell(\theta) \right){ii} {\rm Var}(u_i) + \sum_{j\neq i} \left( \nabla^2 \ell(\theta)\right){ij} \mathbb{E}[u_i] \mathbb{E}[u_j] \\ &= \left(\nabla^2 \ell(\theta) \right){ii}. \end{aligned} $$
End of Proof.
Before we derive the Gauss-Newton-Bartlett (GNB) estimator, we first introduce two identities.
<aside> 💡
Setup. Consider the CE loss. Given a pair of examples $(x, y)$, we have
$$ \ell_{\rm CE}(f(\theta, x), y) = - \sum_{k=1}^K p_k \log q_k, \\ q_k = \frac{\exp(f(\theta, x)k)}{\sum{i=1}^K\exp (f(\theta, x)_i)}, \quad p_k = \textbf{1}\{y=k\}. $$
</aside>
<aside> 💡
Bartlett's identity I.
$$ \mathbb{E}{\hat{y}\sim {\rm Cat}(K, q)} [\nabla{\theta} \ell_{\rm CE} (f(\theta, x), \hat{y})] = 0. $$
</aside>
Proof.
$$ \begin{aligned} \nabla_{\theta} \ell_{\rm CE} (f(\theta, x), y) &= \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \frac{\partial \ell_{\rm CE}(f(\theta, x), y)}{\partial f(\theta, x)} \\ &= \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \sum_{k=1}^K \left(\frac{\partial \ell_{\rm CE}(f(\theta, x), y)}{\partial q_k} \frac{\partial q_k}{\partial f(\theta, x)}\right) \\ &= \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \sum_{k=1}^K \left(-\frac{p_k}{q_k} \frac{\partial q_k}{\partial f(\theta, x)}\right) \\ &= \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \sum_{k=1}^K \left(-\frac{p_k}{q_k} \frac{\partial q_k}{\partial f(\theta, x)}\right) \\ \end{aligned} $$
We know that
$$ \begin{aligned} \frac{\partial q_k}{\partial f(\theta, x)_k} &= \frac{\exp(f(\theta, x)k)}{\sum{i=1}^K \exp(f(\theta, x)_i)} - \frac{\exp(f(\theta, x)k)^2}{\left(\sum{i=1}^K \exp(f(\theta, x)_i) \right)^2} \\ &= q_k - q_k^2 \end{aligned} $$
Also,
$$ \begin{aligned} \frac{\partial q_k}{\partial f(\theta, x)_i} &= - \frac{\exp(f(\theta, x)_k) \exp(f(\theta, x)i)}{\left(\sum{i=1}^K \exp(f(\theta, x)_i) \right)^2} \\ &= -q_k q_i \end{aligned} $$
Therefore,
$$ \frac{\partial q_k}{\partial f(\theta, x)} = q_k (e_k - q), \quad \frac{\partial \ell(f(\theta, x), y)}{\partial f(\theta, x)} = \sum_{k=1}^K p_k (q - e_k) = q - p $$
Then,
$$ \begin{aligned} \mathbb{E}{\hat{y}\sim {\rm Cat}(K, q)} [\nabla{\theta} \ell_{\rm CE} (f(\theta, x), \hat{y})] &= \mathbb{E}{\hat{y}\sim {\rm Cat}(K, q)} \left[\frac{\partial f(\theta, x)^{\top}}{\partial \theta} \frac{\partial \ell{\rm CE}(f(\theta, x), y)}{\partial f(\theta, x)} \right] \\ &= \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \mathbb{E}_{\hat{y}\sim {\rm Cat}(K, q)} \left[q-p \right] \\ &= \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \cdot 0 = 0 \end{aligned} $$
End of Proof.
<aside> 💡
Bartlett's identity II.
$$ \mathbb{E}{\hat{y} \sim {\rm Cat}(K, q)} \left[ \frac{\partial^2 \ell{\rm CE}(f, \hat{y})}{\partial^2 f}\right] = \mathbb{E}{\hat{y} \sim {\rm Cat}(K, q)} \left[ \frac{\partial \ell{\rm CE}(f, \hat{y})}{\partial f} \left(\frac{\partial \ell_{\rm CE}(f, \hat{y})}{\partial f}\right)^{\top} \right]. $$
</aside>
Proof.
$$ \begin{aligned} \frac{\partial^2 \ell_{CE}(f, \hat{y})}{\partial^2 f} &= \frac{\partial}{\partial f}\left(\frac{\partial \ell_{CE}(f, \hat{y})}{\partial f^{\top}} \right) \\ &= \frac{\partial (q-p)^{\top}}{\partial f} \\ &= \begin{pmatrix} q_1(e_1-q), \cdots , q_K (e_K-q) \end{pmatrix} \\ &= \text{diag}(q) -qq^{\top} \end{aligned} $$
Then,
$$ \begin{aligned} \frac{\partial \ell_{\rm CE}(f, \hat{y})}{\partial f} \left(\frac{\partial \ell_{\rm CE}(f, \hat{y})}{\partial f}\right)^{\top} &= \left(p - q\right)\left(p - q\right)^{\top} \\ &= p p^{\top} - q p^{\top} - pq^{\top} + qq^{\top} \end{aligned} $$
Taking expectation,
$$ \mathbb{E}{\hat{y} \sim {\rm Cat}(K, q)} \left[ \frac{\partial \ell{\rm CE}(f, \hat{y})}{\partial f} \left(\frac{\partial \ell_{\rm CE}(f, \hat{y})}{\partial f}\right)^{\top} \right] = \mathbb{E}_{\hat{y} \sim {\rm Cat}(K, q)}[pp^\top] - qq^{\top} = \text{diag}(q) - qq^{\top} $$
End of Proof.
Let us consider the famous Gauss-Newton decomposition,
$$ \nabla^2_\theta \ell (f(\theta, x), y) = \frac{\partial f(\theta, x)^{\top}}{\partial \theta} \frac{\partial^2 \ell(f(\theta, x), y)}{\partial^2 f(\theta, x)} \frac{\partial f(\theta, x)}{\partial \theta^{\top}} + \frac{\partial^2 \langle f(\theta,x) , \frac{\partial \ell(f(\theta, x), y)}{\partial f(\theta, x)}\rangle}{\partial^2 \theta} $$
Because the second term is often small compared to the first term, we focus on the first term, namely the Gauss-Newton matrix. Given the two identities, we are able to build an unbiased estimator for the Gauss-Newton matrix.
$$ \begin{aligned} \nabla_\theta^2 \ell(\theta) &= \mathbb{E}{(x, y)}\left[ \nabla^2\theta \ell (f(\theta, x), y)\right] \\ &\approx \frac{1}{B}\sum_{b=1}^B\underbrace{\left( \frac{\partial f(\theta, x_b)^{\top}}{\partial \theta} \frac{\partial^2 \ell(f(\theta, x_b), y_b)}{\partial^2 f(\theta, x_b)} \frac{\partial f(\theta, x_b)}{\partial \theta^{\top}}\right)}{\text{Gauss-Newton Matrix}} \\ &= \frac{1}{B}\sum{b=1}^B\left( \frac{\partial f(\theta, x_b)^{\top}}{\partial \theta} \mathbb{E}{\hat{y}b \sim \text{Cat}(K, q)}\left[\frac{\partial^2 \ell(f(\theta, x_b), \hat{y}b)}{\partial^2 f(\theta, x_b)}\right] \frac{\partial f(\theta, x_b)}{\partial \theta^{\top}}\right) \\ &= \frac{1}{B}\sum{b=1}^B\left( \frac{\partial f(\theta, x_b)^{\top}}{\partial \theta} \mathbb{E}{\hat{y}b \sim {\rm Cat}(K, q)} \left[ \frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial f} \left(\frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial f}\right)^{\top} \right] \frac{\partial f(\theta, x_b)}{\partial \theta^{\top}}\right) \\ &\overset{\text{Idt II}}{=} \frac{1}{B}\sum{b=1}^B\underbrace{\left( \mathbb{E}{\hat{y}b}\left[\frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial \theta} \left(\frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial \theta}\right)^{\top}\right] \right)}{\text{Fisher Information Matrix}}.\\ \end{aligned} $$
Now, we build up the biased estimator for the diagonal Hessian.
<aside> 💡
Gauss-Newton-Bartlett (GNB) estimator.
$$ \begin{aligned} \text{diag}(\nabla_\theta^2\ell(\theta)) &\approx \frac{1}{B}\sum_{b=1}^B\left( \mathbb{E}{\hat{y}b}\left[\frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial \theta} \odot \frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial \theta}\right] \right) \\ & \overset{\text{Idt I}}{=} B\ \mathbb{E}{\hat{y}b}\left[\left(\frac{1}{B}\sum{b=1}^B\frac{\partial \ell{\rm CE}(f, \hat{y}b)}{\partial \theta}\right) \odot \left(\frac{1}{B}\sum{b=1}^B\frac{\partial \ell_{\rm CE}(f, \hat{y}_b)}{\partial \theta}\right)\right]. \end{aligned} $$
</aside>