FLARE++: Low-rank attention with dynamic attention routing
Full self-attention is a strong token mixer for PDE surrogates on irregular domains, but its quadratic cost limits its use on large problems. Latent-space attention methods such as PerceiverIO, Transolver, and FLARE (Fast Low-rank Attention Routing Engine) avoid that cost by routing attention among $N$ tokens through $M\ll N$ learned latents. They compress by dot-product matching of the input tokens against $M$ learned query tokens or projection weights: once trained, the same learned templates serve every input. We remove this restriction with FLARE++, a low-rank attention architecture with input-conditioned routing queries. FLARE++ uses FLARE's own encoder to map the $N$ input tokens to $M$ query tokens, which correct the learned queries. The adapted queries then determine how that same input is compressed and redistributed. This preserves FLARE's explicit low-rank factorization and linear $\mathcal O(NM)$ complexity, and expresses the complete routing operation with standard scaled dot-product attention (SDPA) calls alone. We also provide a multi-GPU context-parallel implementation that shards input tokens across devices without ever gathering the full token sequence on one of them. FLARE++ reduces FLARE's error by $25\%$ on average across five standard PDE benchmarks, achieving the lowest errors among the efficient models compared. The gains persist on industrial-scale DrivAerML aerodynamics and on Long Range Arena, where average accuracy rises by $3.5\%$ over fixed-query FLARE. Code is available at https://github.com/vpuri3/FLARE.py.