Skip to content

Constant Kernel does not work within Numpyro plate #39

Description

@EdwardRaff

See title, bellow code will cause the error. If you remove the Constant kernel it seems to run fine, but errs out that the Constant is not a constant otherwise (but really it is)

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)
    k = k + kernels.Constant(mean)#comment out this line and it runs fine 

    t = jnp.arange(1, 11)/10-1/20
    gp = GaussianProcess(k, t)
    p = sample("gp", gp.numpyro_dist())

Activity

  1. dfm commented on Feb 4, 2022

    @dfm
    Owner

    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.

  2. EdwardRaff commented on Feb 5, 2022

    @EdwardRaff
    Author
  3. dfm commented on Feb 5, 2022

    @dfm
    Owner

    Ok good. Then send a full example and we'll figure it out!

  4. EdwardRaff commented on Feb 11, 2022

    @EdwardRaff
    Author

    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 scalar
    

    Without 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.
    
  5. dfm commented on Feb 11, 2022

    @dfm
    Owner

    Ok, thanks! Can you now describe in words what you expect these models to mean? The Constant kernel 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!

  6. EdwardRaff commented on Feb 11, 2022

    @EdwardRaff
    Author

    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 mean value that I could have simply used directly, but I wouldn't expect it to error out.

  7. dfm commented on Feb 11, 2022

    @dfm
    Owner

    Alright - I'll take a look at this soon. One note: I don't think the Constant kernel is doing what you think it is anyways…, but if using the mean parameter doesn't work either I can see if I can squint at what you're doing and try to provide some advice.

  8. EdwardRaff commented on Feb 11, 2022

    @EdwardRaff
    Author

    Is the Constant kernel not representing $k(x, x') = c$ ?

  9. EdwardRaff commented on Feb 11, 2022

    @EdwardRaff
    Author

    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 scalar
    
  10. dfm commented on Feb 11, 2022

    @dfm
    Owner

    This 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.

  11. dfm commented on Feb 11, 2022

    @dfm
    Owner

    To 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 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.

  12. EdwardRaff commented on Feb 11, 2022

    @EdwardRaff
    Author

    (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.

  13. leeek commented on Jul 22, 2022

    @leeek

    Minor correction to get the example to work: gp.scale_tril doesn't exist and instead should be gp.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!

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions