New to numPyro and Bayesian deep learning. I was building a Bayesian deep learning network using TraceMeanField_ELBO as the loss function. I found that TraceMeanField_ELBO does not scale prior loss. Here is an example to reproduce this behaviour,
import jax
import jax.numpy as jnp
import numpyro
import numpyro.distributions as dist
from numpyro.infer import TraceMeanField_ELBO, Trace_ELBO
jax.config.update('jax_enable_x64', True)
numpyro.enable_x64(True)
def model(x, y=None, beta=1.0):
# simple linear regression model
with numpyro.handlers.scale(scale=beta):
# scale my prior distribution, control its contribution to ELBO loss
a = numpyro.sample('a', dist.Normal(0.0, 1.0))
b = numpyro.sample('b', dist.Normal(0.0, 1.0))
sigma = numpyro.sample('sigma', dist.Exponential(1.0))
numpyro.sample('obs', dist.Normal(a*x + b, sigma), obs=y)
# fixed init value for reproducibility
guide = numpyro.infer.autoguide.AutoNormal(model,
init_loc_fn=numpyro.infer.init_to_value(values={
'a': jnp.ones((1)),
'b': jnp.ones((1)),
'sigma': jnp.ones((1))
}))
optimizer = numpyro.optim.ClippedAdam(step_size=1e-3)
svi1 = numpyro.infer.SVI(model, guide, optimizer, TraceMeanField_ELBO())
svi2 = numpyro.infer.SVI(model, guide, optimizer, Trace_ELBO())
svi1_state = svi1.init(jax.random.key(0), jnp.zeros((10)), jnp.zeros((10)))
svi2_state = svi2.init(jax.random.key(0), jnp.zeros((10)), jnp.zeros((10)))
for beta in jnp.linspace(1, 100, 5, dtype=jnp.float64):
_, loss = svi1.update(svi1_state, jnp.zeros((10)), jnp.zeros((10)), beta)
print(f'Using TraceMeanField_ELBO, beta={beta:5.1f}, loss={loss:9.5f}')
_, loss = svi2.update(svi2_state, jnp.zeros((10)), jnp.zeros((10)), beta)
print(f'Using Trace_ELBO, beta={beta:5.1f}, loss={loss:9.5f}')
The output is,
Using TraceMeanField_ELBO, beta= 1.0, loss=21.04159
Using Trace_ELBO, beta= 1.0, loss=21.63456
Using TraceMeanField_ELBO, beta= 25.8, loss=21.04159
Using Trace_ELBO, beta= 25.8, loss=89.67968
Using TraceMeanField_ELBO, beta= 50.5, loss=21.04159
Using Trace_ELBO, beta= 50.5, loss=157.72480
Using TraceMeanField_ELBO, beta= 75.2, loss=21.04159
Using Trace_ELBO, beta= 75.2, loss=225.76991
Using TraceMeanField_ELBO, beta=100.0, loss=21.04159
Using Trace_ELBO, beta=100.0, loss=293.81503
When using TraceMeanField_ELBO, loss is always the same despite scale values.
The reason being, in the numPyro source code, the elbo.py file, TraceMeanField_ELBO() class, lines 429 to 437, we have
guide_site = guide_trace[name]
try:
kl_qp = kl_divergence(guide_site["fn"], model_site["fn"])
kl_qp = scale_and_mask(kl_qp, scale=guide_site["scale"]) # here, guide_site["scale"] always return None!
_elbo_particle[name] = -jnp.sum(kl_qp)
except NotImplementedError:
_elbo_particle[name] = _get_log_prob_sum(
model_site
) - _get_log_prob_sum(guide_site)
The guide_site[“scale”] always returns None, whereas model_site[“scale”] returns the correct scale values. With the scale term set to None, numPyro DOES NOT scale the prior loss at all.
Is this a bug, or is it intended behaviour?
I am using Python 3.14.6 with numpyro 0.21.0.