Skip to content

Block transition matrices break some addition and multiplication of kernels #265

Description

@markfortune

Hi, great work on the package!

The block transition matrices implemented for quasiseparable matrices are a neat optimisation but I've noticed that some operations don't seem to have been modified to deal with them. I give a few minimal examples here, in practice they cause a lot of headaches for some multi-wavelength light curve fitting I'm doing and also for GP optimisations I'm implementing which work with quasiseparable matrices directly. I'm not sure what version these specifically became an issue but it's certainly an issue with the latest versions.

For this example I'm running:

  • tinygp 0.3.1
  • jax, jaxlib 0.9.2
import tinygp
import jax.numpy as jnp

N_t = 100
t = jnp.linspace(-10., 10., N_t)

# k_beat has Block transition matrices
k_beat = tinygp.kernels.quasisep.Cosine(1.) + tinygp.kernels.quasisep.Cosine(2.)
banded_term = tinygp.noise.Banded(diag=jnp.ones(N_t), off_diags=jnp.ones((N_t, 1)))

# addition of two QSM where at least one of them has Block transition matrices fails
gp1 = tinygp.GaussianProcess(k_beat, t, noise=banded_term)  # breaks

k_prod = k_beat * tinygp.kernels.quasisep.Exp(1.)

# product of two kernels where at least one of them has Block transition matrices fails
gp2 = tinygp.GaussianProcess(k_prod, t, diag=jnp.ones(N_t))  # breaks
gp3 = tinygp.GaussianProcess(k_beat * k_beat, t, diag=jnp.ones(N_t))  # breaks

Crash 1: Adding QSMs with Block transition matrices

The first crash gives the error TypeError: Cannot determine dtype of Block(blocks=(f32[2,2], f32[2,2])) which corresponds to:

211 p1, q1, a1 = self
212 p2, q2, a2 = other
213 return StrictLowerTriQSM(
214     p=jnp.concatenate((p1, p2)),
215     q=jnp.concatenate((q1, q2)),
216     a=block_diag(a1, a2),  # <-- fails here
217 )

Crash 2: Product kernel with Block transition matrices

The second crash gives the error TypeError: dot_general requires contracting dimensions to have the same shape, got (0,) and (4,). which corresponds to Quasisep.to_symm_qsm:

95 h = jax.vmap(self.observation_model)(X)
96 q = h
97 p = h @ Pinf  # <-- fails here
98 d = jnp.sum(p * q, axis=1)
99 p = jax.vmap(lambda x, y: x @ y)(p, a)

Crash 3: Product of two Sum kernels

The third crash also happens when building a symmetric QSM:

 89 def to_symm_qsm(self, X: JAXArray) -> SymmQSM:
 90     """The symmetric quasiseparable representation of this kernel"""
 91     Pinf = self.stationary_covariance()  # <-- enters Product.stationary_covariance
 92     a = jax.vmap(self.transition_matrix)(
 93         jax.tree_util.tree_map(lambda y: jnp.append(y[0], y[:-1]), X), X
 94     )
 95     h = jax.vmap(self.observation_model)(X)

which calls into _prod_helper:

273 def stationary_covariance(self) -> JAXArray:
274     return _prod_helper(
275         self.kernel1.stationary_covariance(),
276         self.kernel2.stationary_covariance(),
277     )
639     return a1[i] * a2[j]
640 elif a1.ndim == 2:
641     return a1[i[:, None], i[None, :]] * a2[j[:, None], j[None, :]]  # <-- fails here
642 else:
643     raise NotImplementedError

which ultimately hits:

47 @jax.jit
48 def __mul__(self, other: Any) -> "Block":
49     return Block(*(b * other for b in self.blocks))
    # TypeError: unsupported operand type(s) for *: 'DynamicJaxprTracer' and 'Block'

Ideally I would like if there were an option to turn off the formation of block matrices (as suggested in PR #240), but the computational savings they can offer is useful and I'd imagine there shouldn't be any fundamental issue with updating the addition and multiplication rules of kernels to account for Block transition matrices, so that would be a nicer long-term fix.

Thanks again for all the great work!

Activity

  1. dfm commented on Apr 2, 2026

    @dfm
    Owner

    Thanks for the clear repro! I think this should be fixed by #266, which also adds a (pretty awkward to use) option to disable the blocks. Give it a try!

  2. markfortune commented on Apr 7, 2026

    @markfortune
    ContributorAuthor

    Hi Dan, thanks for the quick response!

    So the option to disable the formation of block matrices is a nice addition, and these edits fix the issue for the sum and product of two quasisep kernels (although it would be nice if it could maintain the computational advantage of keeping the block matrix structure 😊, can understand that might not be a high priority though). However, there is a remaining bug if more than two kernels are added together, I guess because it forms nested Block transition matrices which result in similar issues with addition and multiplication as before.

    For example, in the slightly modified initial example we have:

    import tinygp
    import jax.numpy as jnp
    
    N_t = 100
    t = jnp.linspace(-10., 10., N_t)
    
    # k_Fourier has nested Block transition matrices
    k_Fourier = tinygp.kernels.quasisep.Cosine(1.) + tinygp.kernels.quasisep.Cosine(2.) + tinygp.kernels.quasisep.Cosine(4.)
    banded_term = tinygp.noise.Banded(diag = jnp.ones(N_t), off_diags = jnp.ones((N_t, 1)))
    
    # addition of two QSM where at least one of them has nested Block transition matrices fails
    gp1 = tinygp.GaussianProcess(k_Fourier, t, noise=banded_term) # breaks
    
    k_prod = k_Fourier * tinygp.kernels.quasisep.Exp(1.)
    
    # product of two kernels where at least one of them has nested Block transition matrices fails
    gp2 = tinygp.GaussianProcess(k_prod, t, diag=jnp.ones(N_t)) # breaks
    gp3 = tinygp.GaussianProcess(k_Fourier * k_Fourier, t, diag=jnp.ones(N_t)) # breaks

    we get it crashing in all three cases with the error "TypeError: Cannot interpret 'Block(blocks=(f32[2,2], f32[2,2]))' as a data type". I guess if the nesting behaviour was switched to stacking the blocks together instead that'd be enough to fully close this issue?

    Thanks again!

  3. dfm commented on Apr 7, 2026

    @dfm
    Owner

    D'oh! I'm happy to take another closer look when I get a chance, but if you're up for it, perhaps you could dig a bit more to see if you can contribute a fix? Absolutely no pressure of course.

  4. markfortune commented on Apr 8, 2026

    @markfortune
    ContributorAuthor

    That's fair yeah (clearly seen through my goal of getting out of doing more work!), I can try give it a look either this week or next week and see if I can contribute a fix.

  5. markfortune commented on Apr 10, 2026

    @markfortune
    ContributorAuthor

    Now have the above code working with a small update to avoid nested Block objects in PR #267, haven't touched further optimising it by avoiding the conversion to dense matrices but if it is a consistent problem I might take a look in future.

    I was also considering whether it might be worth implementing the transition matrices as banded rather than in separate blocks that need to be looped over, in order to avoid the long JAX compile times mentioned in PR #240 while keeping most of the performance. If I get a chance I might give it a try at some point. Anyway happy to close if you think the new fix is good!

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