Repository navigation
Constant Kernel does not work within Numpyro plate #39
Description
Activity
I assume that you want to change the mean, right? Not add a constant to the kernel everywhere. In that case, you'll use the mean parameter:
with numpyro.plate('Clusters', max_clusters): mean = numpyro.sample("mean", dist.Normal(1.0, 5)) expo_len = numpyro.sample("expo_len", dist.HalfNormal(5)) k = kernels.ExpSquared(expo_len) t = jnp.arange(1, 11)/10-1/20 gp = GaussianProcess(k, t, mean=mean) p = sample("gp", gp.numpyro_dist())
If not, please provide a full working example (it needs imports, etc.) and I can explain from there.
- Both result in the same error - and I’m interested in each case. I’ve actually been having a lot of trouble getting the GP to work within a plate…Sent from my iPhoneOn Feb 4, 2022, at 6:51 PM, Dan Foreman-Mackey ***@***.***> wrote: I assume that you want to change the mean, right? Not add a constant to the kernel everywhere. In that case, you'll use the mean parameter: with numpyro.plate('Clusters', max_clusters): mean = numpyro.sample("mean", dist.Normal(1.0, 5)) expo_len = numpyro.sample("expo_len", dist.HalfNormal(5)) k = kernels.ExpSquared(expo_len) t = jnp.arange(1, 11)/10-1/20 gp = GaussianProcess(k, t, mean=mean) p = sample("gp", gp.numpyro_dist()) If not, please provide a full working example (it needs imports, etc.) and I can explain from there. — Reply to this email directly, view it on GitHub, or unsubscribe. Triage notifications on the go with GitHub Mobile for iOS or Android. You are receiving this because you authored the thread.
Ok good. Then send a full example and we'll figure it out!
With imports
from tinygp import kernels, GaussianProcess import numpyro from numpyro.infer import DiscreteHMCGibbs, MCMC, NUTS import jax import jax.numpy as jnp
Both of the following error out:
def model(max_clusters): with numpyro.plate('Clusters', max_clusters): mean = numpyro.sample("mean", numpyro.distributions.Normal(1.0, 5)) # expo_len = numpyro.sample("expo_len", numpyro.distributions.HalfNormal(5)) # k = kernels.ExpSquared(expo_len) k = kernels.Constant(mean) t = jnp.linspace(0, 1) gp = GaussianProcess(k, t) p = numpyro.sample("gp", gp.numpyro_dist()) return SEED = 1337 max_clusters = 3 args = { 'num_warmup': 250, 'num_samples':450, 'num_chains':1, } rng_key, rng_key_predict = jax.random.split(jax.random.PRNGKey(SEED)) mcmc = MCMC( NUTS(model), num_warmup=args['num_warmup'], num_samples=args['num_samples'], num_chains=args['num_chains'], progress_bar=True, ) mcmc.run(rng_key, max_clusters)
gives
[/usr/local/lib/python3.7/dist-packages/tinygp/kernels.py](https://localhost:8080/#) in __init__(self, value) 154 def __init__(self, value: JAXArray): 155 if jnp.ndim(value) != 0: --> 156 raise ValueError("The value of a constant kernel must be a scalar") 157 self.value = value 158 ValueError: The value of a constant kernel must be a scalarWithout the plate a different error
def model(max_clusters): mean = numpyro.sample("mean", numpyro.distributions.Normal(1.0, 5)) k = kernels.Constant(mean) t = jnp.linspace(0, 1) gp = GaussianProcess(k, t) p = numpyro.sample("gp", gp.numpyro_dist()) return SEED = 1337 max_clusters = 3 args = { 'num_warmup': 250, 'num_samples':450, 'num_chains':1, } rng_key, rng_key_predict = jax.random.split(jax.random.PRNGKey(SEED)) mcmc = MCMC( NUTS(model), num_warmup=args['num_warmup'], num_samples=args['num_samples'], num_chains=args['num_chains'], progress_bar=True, ) mcmc.run(rng_key, max_clusters)
[/usr/local/lib/python3.7/dist-packages/numpyro/distributions/distribution.py](https://localhost:8080/#) in __init__(self, batch_shape, event_shape, validate_args) 176 raise ValueError( 177 "{} distribution got invalid {} parameter.".format( --> 178 self.__class__.__name__, param 179 ) 180 ) ValueError: MultivariateNormal distribution got invalid scale_tril parameter.Ok, thanks! Can you now describe in words what you expect these models to mean? The
Constantkernel doesn't really make sense on its own, but more fundamentally, since you're not actually conditioning on anything I'm not sure what you're expecting these models to describe!If you are asking about my ultimate goals, I don't want my model to be just a mean (using the mean=BLA also causes issues), but I do have prior knowledge for my models where I want to encode a mean explicitly so that I can treat the other parts as an offset to the mean for interpretation/ analysis purposes.
If you are asking about this specific MVE, I was trying to get this down to just what causes the error - nothing extraneous beyond the computational issue. I would expect this code to run without error and produce a GP that is effectively redundant - returning the
meanvalue that I could have simply used directly, but I wouldn't expect it to error out.Alright - I'll take a look at this soon. One note: I don't think the
Constantkernel is doing what you think it is anyways…, but if using themeanparameter doesn't work either I can see if I can squint at what you're doing and try to provide some advice.Is the
Constantkernel not representing$k(x, x') = c$ ?Beyond that, if I really crudely modify the doc example to occur within a plate - the same error occurs (I know this is a useless model ATM, but it can't get to a useful point because of the error)
random = np.random.default_rng(203618) x = np.linspace(-3, 3, 20) true_log_rate = 2 * np.cos(2 * x) y = random.poisson(np.exp(true_log_rate)) def model(x, y=None): with numpyro.plate('Clusters', 2): # The parameters of the GP model mean = numpyro.sample("mean", dist.Normal(0.0, 2.0)) sigma = numpyro.sample("sigma", dist.HalfNormal(3.0)) rho = numpyro.sample("rho", dist.HalfNormal(10.0)) # Set up the kernel and GP objects kernel = sigma**2 * kernels.Matern52(rho) gp = GaussianProcess(kernel, x, diag=1e-5, mean=mean) # This parameter has shape (num_data,) and it encodes our beliefs about # the process rate in each bin log_rate = numpyro.sample("log_rate", gp.numpyro_dist()) log_rate = log_rate.sum(axis=-1) print(log_rate.shape) # Finally, our observation model is Poisson numpyro.sample("obs", dist.Poisson(jnp.exp(log_rate)), obs=y) # Run the MCMC nuts_kernel = numpyro.infer.NUTS(model, target_accept_prob=0.9) mcmc = numpyro.infer.MCMC( nuts_kernel, num_warmup=10, num_samples=10, num_chains=2, progress_bar=False, ) rng_key = jax.random.PRNGKey(55873) mcmc.run(rng_key, x, y=y) samples = mcmc.get_samples()
produces
ValueError Traceback (most recent call last) [<ipython-input-24-88e45b64736a>](https://localhost:8080/#) in <module>() 35 ) 36 rng_key = jax.random.PRNGKey(55873) ---> 37 mcmc.run(rng_key, x, y=y) 38 samples = mcmc.get_samples() 13 frames [/usr/local/lib/python3.7/dist-packages/tinygp/kernels.py](https://localhost:8080/#) in __init__(self, value) 154 def __init__(self, value: JAXArray): 155 if jnp.ndim(value) != 0: --> 156 raise ValueError("The value of a constant kernel must be a scalar") 157 self.value = value 158 ValueError: The value of a constant kernel must be a scalarThis example is useful! I now see the issues, I think, and I'll be able to give some tips. Let me play around with it a bit and report back.
Reacted by EdwardRaffTo get this last example to work, here's how you would do it:
def model(x, y=None): with numpyro.plate('Clusters', 2): mean = numpyro.sample("mean", dist.Normal(0.0, 2.0)) sigma = numpyro.sample("sigma", dist.HalfNormal(3.0)) rho = numpyro.sample("rho", dist.HalfNormal(10.0)) def build_gp(sigma, rho, mean): kernel = sigma**2 * kernels.Matern52(rho) gp = GaussianProcess(kernel, x, diag=1e-5, mean=mean) return gp.loc, gp.scale_tril loc, scale_tril = jax.vmap(build_gp)(sigma, rho, mean) log_rate = numpyro.sample("log_rate", dist.MultivariateNormal(loc, scale_tril=scale_tril)) log_rate = log_rate.sum(axis=0) numpyro.sample("obs", dist.Poisson(jnp.exp(log_rate)), obs=y)
(In this particular case, it would probably make sense to just model it as a single GP, since a sum of GPs is a GP where the kernel is the sum of kernels, but I know that you meant this as an artificial example!)
The key point is that
tinygpisn't going to know which axes tovmap, so you'll need to do that manually. We could certainly add a feature to specify a batch dimension, but I'm not so keen on that since I think it's generally useful to be explicit with such things.(In this particular case, it would probably make sense to just model it as a single GP, since a sum of GPs is a GP where the kernel is the sum of kernels, but I know that you meant this as an artificial example!)
Totally, was more about the mechanical code & process than a useful example - useful examples I'm allowed to share are much harder.
The key point is that tinygp isn't going to know which axes to vmap, so you'll need to do that manually. We could certainly add a feature to specify a batch dimension, but I'm not so keen on that since I think it's generally useful to be explicit with such things.
This makes sense! IMO Adding this to the tutorial would be super helpful for others so people don't run into the same confusion on the plates.
Minor correction to get the example to work:
gp.scale_trildoesn't exist and instead should begp.solver.scale_tril.def build_gp(sigma, rho, mean): kernel = sigma**2 * kernels.Matern52(rho) gp = GaussianProcess(kernel, x, diag=1e-5, mean=mean) return gp.loc, gp.solver.scale_tril
Thanks for the helpful example!
See title, bellow code will cause the error. If you remove the
Constantkernel it seems to run fine, but errs out that the Constant is not a constant otherwise (but really it is)