Skip to content
HN On Hacker News ↗

Exploding variance of means of exponentials: least-squares to the rescue – Machine Learning Research Blog

▲ 65 points • 0 comments • by matt_d • 2w ago • HN discussion ↗

Pangram verdict · v3.3

We believe that this entire text is human-written.

0 %

AI likelihood · overall

Human
100% human-written 0% AI-generated
SEGMENTS · HUMAN 1 of 1
SEGMENTS · AI 0 of 1
WORD COUNT 1,367
PEAK AI % 0% · §1
Analyzed
Sep 27
backend: pangram/v3.3
Segments scanned
1 windows
avg 1367 words each
Distribution
100 / 0%
human / AI fraction
Verdict
Human
Pangram v3.3

Article text · 1,367 words · 1 segments analyzed

Human AI-generated
§1 Human · 0%

A common task in machine learning is to estimate or optimize “log-sum-exp” functions with (potentially continuously) many terms such as $$ \log \Big( \int_{\mathcal{X}} e^{v(x)} dq(x) \Big),$$ where \(v: \mathcal{X} \to \mathbb{R}\) is some potential function, and \(q\) is a probability distribution on the set \(\mathcal{X}\). This has many applications throughout data science, often through the normalization of probabilistic models, but also as a smooth approximation to the maximum, in transformers through its derivatives, or in reinforcement learning when using entropy regularization [19]. Sometimes the set \(\mathcal{X}\) is finite (potentially big) and the integral can be done by explicit summing, but often an exact computation is infeasible, and sampling from the probability distribution \(q\) is used instead. The key difficulty comes from the variance of such estimates, in particular when \(v\) takes large values. In the simplest example, for \(z_1,\dots,z_n \in \mathbb{R}\) independent and normally distributed with mean \(\mu\) and variance \(\sigma^2\), the relative squared error for estimating \(\mathbb{E}[e^z]\) is $$\frac{ {\rm var}\big( \frac{1}{n} \sum_{i=1}^n e^{z_i} \big) }{( \mathbb{E}[ e^{z} ])^2} = \frac{1}{n} \frac{ {\rm var}(e^z) }{( \mathbb{E}[ e^{z} ])^2} = \frac{ e^{\sigma^2}-1}{n}.$$ It converges to zero when \(n\) grows (as can be expected from the law of large numbers), but explodes exponentially when \(\sigma\) grows. Even taking the logarithm does not change the exploding variance, that is, \({\rm var}\big( \log \big( \frac{1}{n} \sum_{i=1}^n e^{z_i} \big)\big)\) also can be shown to grow asymptotically similarly in \(\frac{ e^{\sigma^2}-1}{n}\) (when \(n\) is large, as can be obtained from the delta method). While difficult to estimate, the log-sum-exp function comes with many nice properties (and that’s why people love it); I particularly like the fact that (1) it is a smooth approximation to the maximum (see, e.g., this earlier post), and (2) it is a way to normalize probabilistic models that is adapted to maximum likelihood estimation, in particular in hierarchical probabilistic models, where (conditional) independence assumptions lead to separability of associated loss functions (as thoroughly used in probabilistic graphical models). The main question I try to answer in this post is: Can we keep the advantages of optimizing log-sum-exp functions while being less exposed to their computational / statistical disadvantages? At the other end of the spectrum sits least-squares regression, with essentially the exact opposite features: On the positive side, we obtain computational and statistical simplicity in various forms, e.g., it leads to closed-form estimation for linear models through linear algebra, it is based on computing moments with fixed controlled variance, and it leads to sharp analyses in various setups (acceleration, stochastic gradient descent, etc.). See, e.g., this post on acceleration, and this one on averaging. On the negative side, using least-squares regression for all prediction problems, particularly with discrete outputs, creates some artefacts. The traditional example is classification with Gaussian class-conditional data (with identical covariance matrices), where least-squares on the one-hot encoded outputs has problems, such as “masking” (see [13, Section 2.4] and the example below), or high approximation error compared with using multinomial logistic regression (a.k.a. softmax regression), because then the log conditional probabilities are affine. Can we reconcile them? In other words, is least-squares really all I need? (my colleagues sometimes mock me for my love of least-squares). Note that there is another (classic) attempt at seeing the world through least-squares: doing it in series through Newton’s method, leading in this context to iteratively reweighted least-squares, but this is for computations only, with no statistical improvement. What we are aiming at is stronger: can we get least-squares-based closed-form estimators for maximum-likelihood problems that typically require optimization of a convex function (such as logistic or softmax regression)? Interestingly, my new attempt can be summarized in one integral equation $$ t \log t\, – t + 1 = \int_0^1 \!\! \frac{ (t-1)^2}{\rho t + 1-\rho} (1-\rho) d\rho,$$ which can be checked by usual integration tricks. Let’s see why and how! Relative density estimation as a testbed In this post, I look at a simple fundamental problem where we can study and compare various estimation frameworks, noting that it can be extended in several ways (in particular, through mutual information, see below). We consider two probability distributions \(p\) and \(q\) on \(\mathcal{X}\); our goal is to estimate the logarithm of the relative density \(\log \big(\frac{dp}{dq}(x)\big)\). This turns out to be equivalent to estimating the Kullback-Leibler (KL) divergence because of the variational formulation [1] $${\rm KL}(p\|q) = \int_{\mathcal{X}} \log \big(\frac{dp}{dq}(x)\big) dp(x) = \sup_{v: \mathcal{X} \to \mathbb{R}} \int_{\mathcal{X}} v(x) dp(x) + 1\, – \int_{\mathcal{X}} e^{v(x)} dq(x). \tag{1} $$ This is one particularly important instance of an \(f\)-divergence (see, e.g., [2]), with the following definition and variational formulation based on the Fenchel conjugate \(f^\ast\) of \(f\): $$D(p\|q) = \int_{\mathcal{X}} f \big( \frac{dp}{dq}(x) \big) dq(x)= \sup_{v: \mathcal{X} \to \mathbb{R}} \int_{\mathcal{X}} v(x) dp(x) \, – \int_{\mathcal{X}} f^\ast(v(x)) dq(x),$$ the representation being a consequence of \(f(t) = \sup_{ u \in \mathbb{R}} ut-f^\ast(u)\) applied to each \(t = \frac{dp}{dq}(x)\). The KL divergence corresponds to \(f(t) = t \log t \, – t + 1\) and \(f^\ast(u) = e^u \, – 1\). Note that for the particular case of the KL divergence, when optimizing with respect to a constant on top of \(v\), we obtain the Donsker-Varadhan representation [3] $${\rm KL}(p\|q) = \sup_{v: \mathcal{X} \to \mathbb{R}} \int_{\mathcal{X}} v(x) dp(x)\, – \log \Big( \int_{\mathcal{X}} e^{v(x)} dq(x) \Big). \tag{2}$$ We see the log-sum-exp function appearing explicitly. To estimate the potential \(v\) from i.i.d. samples from \(p\) and \(q\), the traditional variational approach corresponds to replacing integrals with empirical averages. For \(q\), this leads to a potentially unstable empirical average when only samples are available. The goal of this post is to explore another way (I present the main principles behind this new framework; see [7] for more details). The framing through \(f\)-divergences is really key to the new approach, as other divergences will be instrumental in the definition. Before we move on, we state another (equivalent) variational formulation with two potentials \(v\) and \(w\), which we will need later: \(D(p\|q)\) is equal to $$\sup_{v,w: \mathcal{X} \to \mathbb{R}} \int_{\mathcal{X}} v(x) dp(x) + \int_{\mathcal{X}} w(x) dq(x) \mbox{ such that } \forall x \in \mathcal{X}, w(x) \leqslant -f^\ast(v(x)). \quad \tag{3}$$ At optimum, we get \(w(x) = -f^\ast(v(x))\), and we recover Eq. (1) in the KL case. The constraint is convex, but for \(f(t) = t \log t – t + 1\), it is far from what traditional convex optimization methods typically allow. This formulation appears in [25, Theorem 4.4] and has the nice property of preserving the symmetry of the problem (that is, if \(p\) and \(q\) are swapped, this is equivalent to replacing \(f\) by \(t \mapsto t f(1/t)\), and this corresponds to swapping \(v\) and \(w\).) In what follows, we will obtain candidates for functions \(v\) and \(w\) that satisfy the constraint \(\forall x \in \mathcal{X}, \ w(x) \leqslant -f^\ast(v(x))\), typically without equality. Weighted chi-square divergences Another relevant function \(f\) for \(f\)-divergences is, for a parameter \(\rho \in [0,1]\), $$ f(t) = \frac{1}{2} \frac{ (t-1)^2}{ \rho t + 1-\rho}. $$ It leads to a weighted chi-square divergence $$ D(p\|q) = \frac{1}{2} \int_{\mathcal{X}} \frac{ \big(\frac{dp}{dq}(x)-1 \big)^2}{ \rho \frac{dp}{dq}(x) + 1-\rho} dq(x).$$ It has been used in various areas of applied mathematics [4, 5], and comes under several names for special cases, such as Pearson chi-square divergence for \(\rho =0\), or Neyman chi-square (or reverse Pearson) for \(\rho=1\), or Le Cam divergence for \(\rho=1/2\). The function \(f\) above, which is of the form “quadratic over affine” has the variational representation through “quadratic plus affine” functions: $$ \frac{1}{2} \frac{ (t-1)^2}{ \rho t + 1-\rho} = \sup_{u \in \mathbb{R}} \ (t-1) u \, – \frac{1}{2} ( \rho t + 1 – \rho) u^2, $$ with the optimal \(u = \frac{t-1}{\rho t + 1 \, – \rho}\) (this, by the way, is not the Fenchel representation). Thus, applying this for each \(x \in \mathcal{X}\) to \(t = \frac{dp}{dq}(x)\), for this function \(f\), we have $$ D(p\|q) = \!\! \sup_{u(\rho,\cdot):\mathcal{X} \to \mathbb{R}} \int_{\mathcal{X}} \Big\{ \big(\frac{dp}{dq}(x) -1\big) u(\rho,x) \, – \frac{1}{2} \big( \rho \frac{dp}{dq}(x) + 1 – \rho\big) u(\rho,x)^2 \Big\} dq(x). $$ This is exactly a quadratic variational problem since the function \(u(\rho,\cdot): \mathcal{X} \to \mathbb{R}\) only appears quadratically. The optimal variational function is then \(\displaystyle u(\rho,x) = \frac{ \frac{dp}{dq}(x) \, – 1}{ \rho \frac{dp}{dq}(x) + 1-\rho}.\)