It seems like building a context tree with a convex branch cross attention estimator then using branch and bound to prune the tree while descending to get exact cross attention when it's above a threshold would work pretty well, assuming the cross attention matrix actually is very sparse and the trouble is just accurately guessing the non-sparse elements.