FLARE++: Low-rank attention with dynamic attention routing

arXiv:2608.11519 · cs.LG · Submitted 2026-08-12 · Read on arXiv

Carnegie Mellon University

cs.LG

Submitted: 2026-08-12

Updated: 2026-09-26

License: http://creativecommons.org/licenses/by/4.0/

Importance score: 75/100

The gist: FLARE++: Low-Rank Attention with Dynamic Attention Routing Vedant Puri, Yongjie Jessica Zhang & Levent Burak Kara Department of Mechanical Engineering, Carnegie Mellon University arXiv:2608.11519v1

Terminology

Summary

FLARE++: Low-Rank Attention with Dynamic Attention Routing

Vedant Puri, Yongjie Jessica Zhang & Levent Burak Kara

Department of Mechanical Engineering, Carnegie Mellon University

arXiv:2608.11519v1 [cs.LG] 12 Aug 2026

Abstract

Full self-attention (Vaswani et al., 2017) is a strong token mixer for PDE surrogates on irregular domains, but its quadratic cost limits its use on high-resolution problems. Efficient latent-attention models such as the Fast Low-rank Attention Routing Engine (FLARE) (Puri et al., 2026) avoid that cost by routing all N tokens through M ≪ N learned latent queries, but those queries are parameters: once trained, the same learned query templates serve every input. We remove this restriction with FLARE++, a low-rank attention architecture with dynamic token routing. FLARE++ reuses FLARE’s own encoder to build its routing queries: learned latent seeds drive one extra encode call that gathers the N input tokens into M input-conditioned queries, and those queries then determine how the same tokens are compressed and redistributed. This preserves FLARE’s explicit low-rank factorization and linear O(N M) 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++ is competitive across a set of standard PDE surrogate benchmarks, improving on fixed-query FLARE by 24% on average, and it gains 2.3 points of average accuracy on Long Range Arena.

Introduction

Self-attention has become a dominant architecture for PDE surrogates because it lets every discretization point communicate with every other one. In every surrogate we consider, each discretization point is embedded as its own token, so a mesh of N points is a sequence of N tokens; we use the latter term throughout, since it is the entity the mixer acts on. The cost of that generality is an N × N communication matrix and therefore O(N 2) work in the number of tokens (Vaswani et al., 2017), which is out of reach on the meshes engineering problems actually produce.

A range of efficient token mixers communicate instead through M ≪ N latent tokens. We build on FLARE (Puri et al., 2026), which uses latent tokens only for routing: for each attention head, M learned latent queries gather the N input tokens into M latent tokens through one scaled dot-product attention (SDPA) call, and a second call reverses the direction and dispatches those values back to the N input positions. The two calls form an encode–decode factorization inducing, for each fixed input, an explicit input-to-input attention matrix of rank at most M that can be implemented entirely with fused SDPA. Latent-workspace models such as Transolver (Wu et al., 2024) instead process a latent sequence with its own self-attention stage. We build on FLARE because it isolates routing as the token-mixing operation itself, which lets us make that routing input-dependent without introducing a separate latent-processing stage.

Every mixer of this kind, including PerceiverIO (Jaegle et al., 2021a), LNO (Wang & Wang, 2024), the Transolver family (Wu et al., 2024; Luo et al., 2025; Zhou et al., 2026), and FLARE, separates into two objects: a compression template, the M-slot structure that determines how information is gathered and redistributed, and the field that template is applied to. The compressed representation depends on the field in all of these methods, trivially so, and that is not what distinguishes them. What distinguishes them is where the template comes from: in FLARE, the learned queries that define it are parameters, whereas other models use a learned pointwise map or a fixed field-dependent rule to construct their templates. In FLARE, training learns M query templates per layer-head once, and those same templates then serve every geometry and every boundary condition: the routing weights respond to the current keys, but the queries defining the template do not. This is the one part of the operator that never sees the field it is compressing, and it is what fixes how the rank-M bottleneck is spent.

We present FLARE++, which constructs the compression template from the input tokens instead of fixing it. This dynamic query construction reuses FLARE’s own encode mechanism: learned latent seeds act as the queries of one extra encode call, which gathers the N input tokens into M vectors, and those M vectors are then used as the routing queries of the encode–decode pair that compresses and redistributes the field. Because the queries are produced by the same encode call FLARE already performs, their construction inherits its efficient fused SDPA implementation and needs no new kernel. The change is confined to the mixer, and within the mixer to the construction of the template: the induced routing matrix remains rank at most M for each fixed input, and the complexity remains O(N M). The residual stream is untouched, carrying the same number of blocks at the same width with the same residual updates, so nothing reported below is bought by making the network deeper or wider. As Figure 1 shows, FLARE++ replaces a static query parameter with one additional SDPA call and changes nothing else.

We evaluate FLARE++ against FLARE and the Transolver family under a matched backbone on standard PDE surrogate benchmarks, and find that FLARE++ attains the lowest relative L2 error on all five (Table 1), reducing fixed-template FLARE’s error by 24% on average and Transolver-3’s by 31%. Joint ablations against FLARE on latent budget M and residual depth B find that dynamic routing improves on fixed-template FLARE in every configuration we measured (Section 5.2). Additionally, FLARE++ continues improving over the measured latent-budget range, whereas FLARE saturates in that range. Furthermore, dynamic routing substitutes for depth, with FLARE++ reaching a lower error than FLARE at a shallower residual depth.

Outside PDE surrogates, the same substitution improves every Long Range Arena task and lifts the average of FLARE by 2.3 points.

Dynamic routing is not free, costing 1.3–1.5× FLARE’s step time at matched depth and latent budget (Section C.3), but it recovers part of that by reaching a given accuracy with fewer blocks. To alleviate that cost, we provide an exact token-sharded implementation that shards input tokens across devices without ever gathering the full token sequence on one of them, and find that parallel efficiency stays at or near unity in both time and memory.

Contributions

  • Low-rank self-attention via a synthesized compression template. FLARE++ uses SDPA to construct M routing queries from the input, so the input tokens determine how they are themselves gathered into and redistributed from a compact latent representation. It preserves FLARE’s explicit rank-M encode–decode operator, independent per-head pathways, O(N M) complexity, and fused-SDPA implementation, and it adds no depth or width to the residual stream.

  • Multi-GPU context parallelism. An exact token-sharded implementation distributes pointwise activations and attention work across accelerators, communicating only latent outputs and softmax statistics, so the collective payload is independent of the number of input tokens and decoding needs no all-gather. Over four ranks, parallel efficiency stays at or near unity in both time and memory.

  • Evaluation on accuracy and on cost. We compare dynamic routing against fixed-query FLARE, Transolver, and full self-attention under a matched backbone on standard PDE benchmarks and on Long Range Arena, and measure what each mixer costs in single-GPU time and memory over three orders of magnitude in the number of tokens, and in multi-GPU parallel efficiency. We report both where the mechanism pays and where it does not.

Method

Preliminaries. We consider PDE surrogate models that operate on fields discretized at N spatial points. Each point carries problem-dependent input quantities, such as coordinates, boundary conditions, material parameters, forcing terms, or an initial state. A pointwise input projection embeds these quantities into a sequence X = [x1,..., xN]⊤ ∈ RN ×C, where C is the hidden width. Every discretization point thus becomes exactly one token. A stack of residual token-mixing and feedforward (FFN) blocks communicates information between the N tokens, and a pointwise output projection decodes the requested solution field.

Full self-attention. Learned projections construct Q = XWQ, K = XWK, V = XWV with WQ, WK, WV ∈ RC×C, split into H heads of width D = C/H. For head h, scaled dot-product attention computes Sh = Qh Kh⊤ / √D ∈ RN ×N, Yh = softmax(Sh) Vh, and the head outputs are concatenated and projected, Y = [Y1,..., YH]WO. Both sides of Sh are functions of the current input, so no fixed latent bottleneck is imposed on the token–token routing matrix; the cost is O(N 2 D) work per head.

FLARE: fixed-query latent routing. FLARE (Puri et al., 2026) replaces the dense matrix with an encode–decode factorization through M ≪ N latent routes. It projects only the keys and values from the current input, while each head owns an independent learned query set: Kh = XWK,h ∈ RN ×D, Vh = XWV,h ∈ RN ×D, Qh ∈ RM ×D is learned. The encoder and decoder are standard SDPA calls, Zh = SDPA(Qh, Kh, Vh), Yh = SDPA(Kh, Qh, Zh), which, writing Sh = Qh Kh⊤ / √D ∈ RM ×N, Wenc,h = softmax(Sh) and Wdec,h = softmax(Sh), gather the N input values into M latent values and scatter them back: Yh = Weff,h Vh, Weff,h = Wdec,h Wenc,h ∈ RN ×N. Because Weff,h factors through M latent routes its rank is at most M, and it is never materialized: the two SDPA calls apply its factors sequentially at O(N M D) cost. The learned queries Qh that define those routes, however, remain fixed across samples.

FLARE++: dynamic token routing. The queries Qh are the compression templates which decide, for a given head, which parts of the input token set each of the M latent slots draws from and returns to. FLARE++ preserves FLARE’s M routing slots but changes where that template comes from. Rather than using a learned set directly as the routing queries, FLARE++ applies FLARE’s own encoder a second time, with learned seeds Q̃h ∈ RM ×D as its queries, and takes its M outputs as the sample-specific routing queries Qh (X). Separate projections of X give K̃h = X W̃K,h, Ṽh = X W̃V,h. One SDPA call synthesizes the routing queries, Qh (X) = SDPA(Q̃h, K̃h, Ṽh) = softmax(Q̃h K̃h⊤ / √D) Ṽh ∈ RM ×D. This is exactly the FLARE encoder of equation 5, with the learned seeds Q̃h in place of Qh and its own key and value projections; the only difference is what its output is used for. Instead of being the transported latent values Zh, the M gathered vectors become the routing queries of the encode–decode pair that follows. Because they are recomputed from the current block input X, the routing-query set changes with the sample and at every layer. The synthesized queries are then used in both factors of the FLARE operator: Kh = XWK,h, Vh = XWV,h, Zh = SDPA(Qh (X), Kh, Vh), Yh = SDPA(Kh, Qh (X), Zh). Writing the resulting routing factors as Wenc,h (X) and Wdec,h (X) gives Yh = Wdec,h Wenc,h Vh, rank(Wdec,h Wenc,h) ≤ M. The synthesized queries Qh (X) define both factors of a conditional routing matrix whose rank is at most M, so the field participates in constructing the template by which it is itself compressed. Query synthesis, gathering, and dispatch all use standard SDPA, with no custom attention kernel or explicit sequence-projection matrix. The complete mixer is three fused SDPA calls.

What dynamic routing does not change. Every modification above is confined to the construction of the routing queries. The residual stream is untouched: a token passes through the same B blocks at the same width C, with the same normalization, residual additions, and pointwise input and output projections as in FLARE. The query-synthesis branch of equation 8 sits outside that stream, since its output is consumed as queries and never added back into the token representation, so FLARE++ adds neither residual depth nor a latent-processing stage of the kind a latent-workspace model introduces. The two extra C × C projections in equation 7 do add parameters, so the mixers are not parameter-matched (Section C.2); what is matched is the depth and width of the representation being mixed, and any accuracy difference is therefore attributable to how the M routes are chosen.

Computational complexity. Both mixers are built from the same two primitives: a pointwise projection, costing O(N C 2) time and O(N C) space, and a token–latent SDPA call, costing O(N M C) time and O(N C) space. The latter is linear rather than quadratic in space because a fused kernel never materializes the N × M score matrix, and the latent tensors are O(M C) with M ≪ N. The mixers differ only in how many of each they use: FLARE performs three projections (K, V, and the output) and two SDPA calls, giving O(N (3C 2 + 2M C)) per block, whereas query synthesis adds K̃, Ṽ, and one more call, giving FLARE++ O(N (5C 2 + 3M C)). Both are linear in N in time and O(N C) in space, and depth multiplies each by B; full self-attention instead needs O(N 2 C) time. FLARE++ is therefore not a free improvement: at matched (M, B, C) it performs roughly 1.6× the mixer arithmetic of FLARE, so dynamic routing is preferable on cost only if it reaches a given accuracy at a smaller latent budget or depth. These are operation counts and not running times, and the two primitives carry very different constants: a dense projection is a large matrix multiplication near the arithmetic peak of the device, whereas a token–latent SDPA call at M ≪ N is a short memory-bound reduction far below it. We therefore treat the counts as a scaling argument and measure the wall-clock consequence in Section C.3.

Multi-GPU context parallelism. Linear complexity does not by itself make a high-resolution mesh fit on one accelerator, because pointwise activations still grow with N. We therefore shard the token dimension across R ranks, so that rank r holds Xr ∈ RB×Nr ×C with Σr Nr = N, writing B for the batch size to keep B for the number of blocks. Pointwise projections, normalization, feed-forward layers, and residual updates are then local. The only globally coupled primitive is the encoder that gathers the sharded input tokens into M replicated latent tokens. Each rank runs a fused SDPA primitive on its local keys and values Kr, Vr, exposing a local latent output Or and its rowwise log-normalizer Lr, and the exact global result is recovered by L = logsumexpR−1 r=0 (Lr), Z = ΣR−1 r=0 exp(Lr − L) Or. This is algebraically identical to applying the encoder to the concatenated token sequence, but communicates only latent outputs and softmax statistics. Decoding, Yr = SDPA(Kr, Q, Z), is entirely local once the routing queries Q and the latent values Z are replicated, so no all-gather over tokens is ever required. Each encoder therefore reduces a per-rank payload of O(BHM (D + 1)) values, independent of N, while token-dependent storage and attention work divide across ranks. FLARE invokes the encoder once per mixer; FLARE++ invokes it twice, once to synthesize Q(X) and once to gather physical values.

Measured scaling. On meshes of 5 × 105 and 106 points, sharded over up to four ranks, parallel efficiency stays at or near unity on both axes: it never falls below 0.92 in time and 0.95 in memory, where unity means the step time and the per-rank peak memory both divide by the rank count. The collective therefore costs little of the time it saves, and the largest mesh a given machine can train grows almost linearly with the rank count, since no stage of the forward or backward pass reconstructs an N-token tensor. Both effects are insensitive to the latent budget, as the N-independent payload predicts.

Experiments

Standard PDE surrogate benchmarks. We evaluate on the Elasticity, Darcy, Airfoil, Pipe, and DrivAerML-40K benchmarks studied by FLARE (Puri et al., 2026). These problems span structured and unstructured discretizations with approximately 1K–40K points per sample. Baselines include full self-attention (Vaswani et al., 2017), the Transolver family (Transolver (Wu et al., 2024), Transolver++ (Luo et al., 2025), and Transolver-3 (Zhou et al., 2026)), and FLARE (Puri et al., 2026) under a shared backbone, with matched channel width, head dimension, depth, and latent count wherever applicable. Full self-attention is not a practical PDE surrogate architecture at these resolutions and serves only as a reference for how much a token mixer gives up by imposing a low-rank bottleneck. We therefore separate it from the other entries in every table and exclude it from best-result rankings. We also include PerceiverIO (Jaegle et al., 2021a), Set Transformer (Lee et al., 2019), GNOT (Hao et al., 2023), and LNO (Wang & Wang, 2024). These four are evaluated as complete architectures rather than as token mixers in a shared backbone, because an identical-backbone mixer swap is not well defined for them (Section B.2). All models are trained in FP32 precision.

Results and discussion. FLARE++ records the lowest error in Table 1 on all five benchmarks. It reduces FLARE’s error by 9–41%, averaging 24%, and Transolver-3’s by 18–44%, averaging 31%. Full self-attention is affordable only on the three smallest benchmarks, where FLARE++ is more accurate on Elasticity and Airfoil, and less accurate on Darcy. The Transolver variants do not separate under this backbone: Transolver-3 is ahead of Transolver on three benchmarks by at most 6% and behind on two by 16–22%, so its published gains do not reproduce at matched width, depth, and latent budget. Model family does not determine the ranking either, since Set Transformer is the strongest non-FLARE model on two benchmarks and leads every Transolver variant on the largest one. The latent-workspace models PerceiverIO and LNO, which process a latent sequence with their own self-attention stage, trail FLARE++ on every benchmark, by up to 4.1×.

Long Range Arena benchmark. PDE surrogate modeling is the empirical focus of this paper, but the routing mechanism is not specific to it. We therefore also evaluate FLARE and FLARE++ on Long Range Arena (Tay et al., 2021b) under an identical-backbone protocol, against full self-attention, Transolver, and a broad set of established efficient-attention methods. Because FLARE and FLARE++ assume no canonical token order, we compare against efficient-attention architectures rather than fixed-order sequence models such as S4 or Mamba (Gu et al., 2021; Gu & Dao, 2024). FLARE++ attains the strongest average accuracy in this comparison, raising the FLARE average from 58.08 to 60.36 while preserving an O(N M) token mixer. The gain is not carried by one task: dynamic routing improves on fixed-query FLARE on all five, by 0.2 points on Retrieval and by 5.2 and 3.5 points on Image and Pathfinder-32. It also places the low-rank mixer above the full self-attention row (60.36 against 57.51), which fixed-query FLARE does not manage.

Model Analysis and Ablations

Efficiency and scaling. Section C.3 reports wall-clock time and peak memory on a single NVIDIA H100 for complete models that differ only in the token mixer, swept over N from 103 to 106. Full self-attention separates from the low-rank mixers in time rather than in memory, and FLARE++ costs a constant factor of 1.3–1.5× FLARE, flat in N. Beyond one device, sharding the token dimension divides activation storage and attention work across ranks while leaving the collective payload independent of N; Table 2 reports the measured efficiency and per-rank memory.

Fixed versus dynamic routing across the latent budget. We sweep the latent budget M and the depth B jointly for FLARE and FLARE++, holding everything else at the values of Section B.1 for the elasticity and darcy benchmarks. The two mixers differ only in how the M routing queries are obtained, so any difference in the grid is attributable to that choice. Dynamic routing wins in every cell of the grid. FLARE++ is more accurate than FLARE in all 21 matched (M, B) cells, nine on Elasticity and twelve on Darcy, by between 21% and 53% relative (Figure 3). The two benchmarks differ in how that margin behaves with depth. On Elasticity it is flat, at 45%, 44%, and 44% for B = 2, 4, 8 averaged over the latent budget, whereas on Darcy it decays, from 38% to 28% to 22% at M = 128. Depth substitutes for dynamic routing on the rank-limited benchmark and does not on the low-rank one, which is the first sign that the two mixers use additional capacity differently.

Fixed queries saturate in M; input-conditioned queries do not. Plotting the same runs against the latent budget separates the two mechanisms, and the two benchmarks respond differently because they demand different routing ranks. Puri et al. (2026) report that global communication on Elasticity is fundamentally low-rank, so accuracy there stops improving with M almost immediately, whereas Darcy is rank-limited and keeps benefiting from additional latents over most of the range. On Elasticity, enlarging the latent budget from M = 32 to M = 128 makes FLARE monotonically worse at B = 2 and B = 4 (1.63 → 1.91 and 0.90 → 0.95) and leaves it unchanged at B = 8, while FLARE++ improves over the identical grid (1.02 → 0.90 and 0.40 → 0.35). On Darcy the effect is milder but the same in kind: both mixers convert additional routes into accuracy over most of the range, and both flatten past M = 128 at the largest depth, where doubling the budget changes FLARE by +0.5% and FLARE++ by −1.2%. The difference between the two mechanisms is therefore where saturation sets in and how much has been extracted by then, not whether it happens at all. Enlarging a fixed template adds routes that are largely redundant across inputs, and past some budget they cost more than they contribute, whereas a template built from the current field keeps using the routes it is given.

Dynamic routing substitutes for depth. FLARE++ cannot undercut FLARE by shrinking M, because its two extra projections floor the per-block cost independently of the latent budget (Section C.2). The grid shows the trade running the other way, and without exception: at B = 4, FLARE++ is more accurate than FLARE at B = 8 (half the residual depth) at every one of the seven latent budgets measured across the two benchmarks. The same substitution one level down, B = 2 against B = 4, holds in only two of those seven, so halving the depth is supported at the depths we swept and not below them. Dynamic routing buys back its per-block cost by needing fewer blocks, which is the opposite of the mechanism we had anticipated. The trade is stated in operation counts rather than in measured training time, but the measured per-block penalty of 1.3–1.5× (Section C.3) is smaller than the arithmetic 1.6×, so halving the depth would be expected to reduce wall-clock cost as well as FLOPs; end-to-end training-time savings are not measured here. We conclude that dynamic routing matters most where a small number of latent routes must serve a heterogeneous input, and least where fixed queries already saturate the achievable accuracy.

Conclusion

We introduced FLARE++, an efficient token mixer that synthesizes its routing queries from the current input tokens before gathering and redistributing information. This makes FLARE’s low-rank communication scaffold adaptive to each sample and layer while preserving linear complexity in the number of tokens at a fixed latent budget. Under a matched backbone, dynamic routing gives the lowest error on a set of standard PDE benchmarks and raises the Long Range Arena average of the same architecture. None of this is bought with depth or width: the residual stream carries the same number of blocks at the same channel count as fixed-query FLARE, and only the construction of the compression template differs.

Synthesizing the template is not free, and the measurements say what it costs. A FLARE++ block runs at 1.3–1.5× FLARE’s step time and 1.18× its peak memory, flat in the number of tokens and below the 1.6× its operation count predicts, because the added work is dense projection rather than attention. That cost is recovered through depth rather than through a smaller latent budget, since FLARE++ reaches a lower error than FLARE at a shallower residual depth. Both models remain linear in the number of tokens where full self-attention does not, and they separate from it in time rather than in memory: at 5 × 105 tokens the unrestricted operator is two orders of magnitude slower while using less storage. An exact token-sharded implementation extends both beyond a single accelerator with a collective payload that does not grow with the number of tokens, at parallel efficiency at or near unity in both time and memory over four ranks.

Improvements for AI systems

Improvements to AI systems:

  1. Input-conditioned latent routing for efficient attention: Replace fixed learned query templates in low-rank attention with dynamically synthesized queries derived from the current input tokens. This allows the model to adapt its compression strategy per sample and per layer, improving accuracy by 24% on PDE surrogates and 2.3 points on Long Range Arena without adding depth or width.

  2. Three-call SDPA mixer for adaptive token mixing: Implement a mixer that uses three standard scaled dot-product attention calls—one to synthesize routing queries from input-conditioned seeds, and two to encode and decode through those queries. This preserves O(N·M) complexity and rank-M factorization while making the routing matrix input-dependent.

  3. Token-sharded context parallelism for long sequences: Distribute input tokens across multiple GPUs without ever gathering the full sequence on one device. Communicate only latent outputs and softmax statistics (payload independent of token count), achieving parallel efficiency ≥0.92 in time and ≥0.95 in memory over four ranks.

  4. Depth-substitution mechanism for cost efficiency: Use dynamic routing to reach target accuracy with fewer residual blocks (e.g., FLARE++ at B=4 outperforms FLARE at B=8), recovering the 1.3–1.5× per-block cost increase through reduced depth.

  5. Latent-budget scalability: Enable continued accuracy improvements as latent budget M increases, where fixed-query methods saturate. Dynamic routing converts additional routes into accuracy gains (e.g., on Elasticity, FLARE++ improves 1.02→0.90 error from M=32 to M=128, while FLARE degrades).

What the improved AI system can do:

  • Process high-resolution PDE surrogates (up to 106 mesh points) with linear complexity in token count, achieving state-of-the-art accuracy on Elasticity, Darcy, Airfoil, Pipe, and DrivAerML-40K benchmarks.

  • Handle long-range sequence tasks (up to 4K tokens in Long Range Arena) with better accuracy than full self-attention (60.36 vs 57.51 average) while using O(N·M) memory.

  • Scale to multi-GPU training on meshes of 5×105–106 points with near-linear speedup and memory division, without reconstructing full N-token tensors.

  • Adapt compression templates to heterogeneous inputs (e.g., varying boundary conditions, geometries) where fixed templates fail, improving accuracy by 21–53% across all tested latent budgets and depths.

  • Achieve target accuracy with fewer layers, reducing wall-clock training time despite higher per-block cost, particularly on low-rank communication problems like Elasticity.

Abstract

Full self-attention is a strong token mixer for PDE surrogates on irregular domains, but its quadratic cost limits its use on high-resolution problems. Efficient latent-attention models such as the Fast Low-rank Attention Routing Engine (FLARE) avoid that cost by routing all N tokens through M << N learned latent queries, but those queries are parameters: once trained, the same learned query templates serve every input. We remove this restriction with FLARE++, a low-rank attention architecture with dynamic token routing. FLARE++ reuses FLARE's own encoder to build its routing queries: learned latent seeds drive one extra encode call that gathers the N input tokens into M input-conditioned queries, and those queries then determine how the same tokens are compressed and redistributed. This preserves FLARE's explicit low-rank factorization and linear 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++ is competitive across a set of standard PDE surrogate benchmarks, improving on fixed-query FLARE by 24% on average, and it gains 2.3 points of average accuracy on Long Range Arena.

Sources

Related papers