Repository navigation
Block transition matrices break some addition and multiplication of kernels #265
Description
Activity
- added a commit that references this issue
on Apr 2, 2026 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!
Reacted by Mark Fortune- added a commit that references this issue
on Apr 2, 2026 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!
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.
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.
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!
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:
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: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 toQuasisep.to_symm_qsm:Crash 3: Product of two Sum kernels
The third crash also happens when building a symmetric QSM:
which calls into
_prod_helper:which ultimately hits:
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!