TraceMeanField_ELBO class behaviour with scaling prior distribution

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.