Skip to content
AI.info

Research

PDE-JEPA: Predictive Representation Learning of Latent Dynamics Modeling for Parametric PDEs

PDE-JEPA: Predictive Representation Learning of Latent Dynamics Modeling for Parametric PDEs Authors: Zhentao Tan, Jianrong Zhang, Ruijie Quan, Yi Yang (Zhejiang University) arXiv: 2609.34715v1 [cs.AI

PDE-JEPA: Predictive Representation Learning of Latent Dynamics Modeling for Parametric PDEs
arXiv
2609.34715
Published
2026-09-28
Authors
Zhentao Tan, Jianrong Zhang, Ruijie Quan, Yi Yang

AI summary

PDE-JEPA: Predictive Representation Learning of Latent Dynamics Modeling for Parametric PDEs

Authors: Zhentao Tan, Jianrong Zhang, Ruijie Quan, Yi Yang (Zhejiang University) arXiv: 2609.34715v1 [cs.AI], 28 Sep 2026 | License: CC BY 4.0

Overview

Research area: Machine learning for scientific computing — specifically representation learning and latent-space dynamics modeling for families of partial differential equations (PDEs) whose governing parameters vary.

Technical level: Advanced. The paper assumes familiarity with neural operators, joint-embedding predictive architectures (JEPAs), masked-latent self-supervision, latent ODE/autoregressive rollout, and relative L2 rollout error as an evaluation metric.

Scope (one sentence): The paper asks what makes a learned latent state space suitable for forecasting and extrapolating parametric PDE dynamics, and answers it with a two-stage framework (geometry alignment plus a physics-structured predictor) built on a frozen JEPA encoder, evaluated on nine PDE benchmarks.

What This Paper Is About

Most latent models for PDEs learn their latent state by reconstructing observed physical fields. The authors argue that reconstruction fidelity does not guarantee that physically relevant information is easily accessible, nor that the resulting latent space is well organized for repeated temporal evolution — especially when governing conditions shift. PDE-JEPA instead starts from predictive (JEPA-style) pretraining, then explicitly reshapes the latent geometry for evolution and builds a predictor whose structure mirrors how parameters enter a PDE.

Key Contributions

  1. First systematic study of JEPA for parametric PDEs. The authors report that informative predictive representations alone do not ensure accurate autoregressive rollout, and use this diagnosis to motivate the proposed framework.

  2. Physics-Aligned Latent Geometry (PAG). A lightweight residual geometry projector that aligns latent trajectory geometry with the evolution geometry of physical fields, with an anchor loss to preserve the information encoded by the pretrained representation.

  3. Physics-Structured Latent Predictor (PSP). A predictor that decomposes latent dynamics into a parameter-independent evolution term and parameter-dependent response terms, encoding PDE-formulation inductive bias to improve extrapolation to unseen governing conditions.

  4. Broad empirical evaluation across nine parametric PDE benchmarks spanning transport, diffusion, reaction–diffusion, wave propagation, and fluid dynamics, reporting an average 33.4% in-distribution improvement and an average 51.4% improvement when extrapolating to unseen governing parameters.

Main Findings

  • Predictive learning gives more physically informative representations. On frozen-encoder probes, JEPA features outperform reconstruction-based features on both a local-state probe (recovering instantaneous physical fields) and a parameter probe (inferring governing conditions). In the Burgers dataset, the JEPA representation shows two branches associated with different parameter ranges.

  • Informativeness does not equal evolvability. Despite stronger probing, a vanilla JEPA-based model produces higher in-distribution rollout errors than a reconstruction-based baseline on both Wave-2D and Vorticity. Under parameter shifts, JEPA achieves lower out-of-distribution error on Wave-2D and comparable performance on Vorticity.

  • Latent trajectories are geometrically "zig-zaggy." The mean turning angle increases from 31.3° to 60.5° on Vorticity and from 18.2° to 59.4° on Burgers relative to physical trajectories.

  • PAG substantially improves rollout. Adding PAG reduces Vorticity error from 0.086 to 0.040 in-distribution and from 0.491 to 0.397 out-of-distribution, corresponding to 53.4% and 19.1% improvements. On Wave-2D it reduces ID error from .363 to .143 and OOD error from .502 to .321.

  • Geometry metrics confirm the alignment. Alignment reduces Angle MAE, Sym. Acc. MAE, and Lag-1 Cos. MAE: Vorticity 29.1→6.3 (−78%), .46→.10 (−77%), .36→.05 (−84%); Wave-2D 35.5→10.1 (−71%), .46→.14 (−68%), .50→.14 (−74%); Burgers 41.4→20.8 (−49%), .73→.37 (−49%), .43→.17 (−59%); Gray–Scott 28.8→20.4 (−29%), .39→.27 (−30%), .39→.26 (−32%). These are computed directly on representations without a dynamics predictor.

  • PSP improves out-of-distribution extrapolation. Adding PSP on top of PAG reduces OOD rollout error from .397 to .288 on Vorticity and from .321 to .157 on Wave-2D. Under matched initial conditions, parameter-conditioned responses improve in both field and latent space (e.g., Vorticity latent cosine .913→.924 and amplitude 1.142→1.043; field cosine .626→.708 and amplitude .846→.987).

  • In-distribution results. PDE-JEPA achieves the lowest relative L2 error on eight of nine benchmarks and ranks second on Advection (0.0074 vs. 0.0068 for CoDA). Reported errors: Advect 0.0074, Burgers 0.0428, Heat 0.0274, Wave-B 0.0350, Combined 0.0074, Wave-2D 0.1140, Vorticity 0.0348, HeterNS 0.0089, GS 0.0284. Reported relative improvements: −8.8%, 50.7%, 70.6%, 68.0%, 12.9%, 44.9%, 41.2%, 9.2%, 12.1%. The authors attribute the Advection gap to translation-dominated dynamics where reconstruction ability matters most, and note several methods already fall below 10⁻² there.

  • Out-of-distribution results. The method reports the lowest error on all five OOD benchmarks: Combined 0.008 (77.9% improvement), Wave-2D 0.157 (74.2%), Vorticity 0.288 (9.7%), HeterNS 0.011 / 0.103 (69.5% / 1.5%), GS 0.033 (59.7%). ID-to-OOD degradation remains limited: 0.0074→0.0084, 0.1140→0.157, and 0.0284→0.0337.

Methodology in Plain English

  1. Pretrain an encoder with masked-latent prediction (JEPA). Instead of rebuilding the physical field, the encoder and predictor are trained to predict representations of masked content in latent space, using an EMA target encoder. The governing parameter ξ is not given to the encoder or predictor during pretraining, encouraging parameter-agnostic representations. The encoder is then frozen and reused.

  2. Align the latent geometry (PAG). A token-wise residual projector computes q = z + G(z), where the final layer is zero-initialized so q = z at the start — it acts as a coordinate correction, not a re-learned representation. The projector is trained so that the directions between consecutive latency steps match the directions between consecutive physical-field steps, using SmoothL1 loss over lags {1, 2, 4}, with an anchor loss keeping q close to z and an auxiliary dynamics loss keeping the aligned coordinates predictable.

  3. Build a structured predictor (PSP). Rather than letting the parameter interact arbitrarily with the latent state throughout the network, the latent vector field is written as a shared evolution term plus a sum over parameter-dependent response terms, each scaled by a normalized physical parameter. The Navier–Stokes vorticity equation, where viscosity ν scales the viscous term, is given as the motivating example. The authors state they do not require either term to recover exact analytical PDE operators.

  4. Propagate in continuous time. Latent states are advanced with a fixed-step fourth-order Runge–Kutta (RK4) ODE solver over one observation interval, and long-horizon predictions come from recursively integrating the predicted states.

  5. Evaluate. Nine benchmarks: Advection, Burgers, Heat, Wave-B, Combined Equation, Vorticity, Wave-2D, Gray–Scott, and HeterNS. Advection, Wave-B, Combined, Wave-2D, and Vorticity follow Zebra; Burgers and Heat are generated following MP-PDE with an added forcing coefficient; Gray–Scott is adopted from ENMA; HeterNS from UniSolver. Baselines span parametric solvers (FNO, CAPE, CoDA, GEPS), in-context solvers (ViT-in-context, [CLS]ViT, Zebra), foundation models (UniSolver, MPP, DPOT-S, Poseidon-T), and latent solvers (LE-PDE, LNS, MAE-PDE, ENMA). The metric is relative L2 error over the full rollout trajectory, averaged over test trajectories.

Dataset scale (as reported): Most datasets use 12,000 training trajectories with 120 validation and 120 test trajectories; Vorticity, Wave-2D, and Gray–Scott use 12,000 / 1,200 / 1,200 plus 120 OOD; HeterNS uses 15,000 / 1,500 / 1,500 / 1,500. Trajectory shapes are 140×256 (Advection, Combined), 250×256 (Burgers, Heat, Wave-B), 1×128×128×30 (Vorticity), 2×64×64×30 (Wave-2D), 2×32×32×20 (Gray–Scott), and 1×64×64×20 (HeterNS). OOD parameter ranges include Vorticity ν from [10⁻³, 10⁻²] to [10⁻⁵, 10⁻⁴], Wave-2D c from [100, 500] to [500, 550] and k from [0, 50] to [50, 60], Combined α Out-D [1.0682, 1.7581], and Gray–Scott F from 𝒰([0.023, 0.045]) to 𝒰([0.045, 0.0467]).

Why This Matters

Impact on research. The paper reframes a common assumption in latent PDE modeling: that a representation good at recovering or probing physical fields is automatically good for recursive forecasting. It supplies evidence that these two capabilities can diverge, and offers a concrete recipe for adapting a pretrained predictive representation rather than replacing it. It also introduces JEPA-style self-supervision to the parametric PDE setting, connecting world-modeling ideas to scientific machine learning, and provides quantitative geometry diagnostics (turning angle, second-order variation, lag-1 directional consistency) that are computed without any dynamics predictor.

Real-world applications (domains represented by the paper's benchmarks):

  • Transport and advection problems, where a quantity moves through a domain with varying speed.
  • Diffusion and forced reaction–diffusion systems, relevant to heat transport and pattern-forming chemical or biological processes (Heat, Gray–Scott).
  • Wave propagation, including boundary-condition and celerity variations (Wave-B, Wave-2D).
  • Fluid dynamics, including incompressible flow and heterogeneous Navier–Stokes regimes (Vorticity, HeterNS).

Industry relevance. The paper motivates learning-based solvers by the substantial computational cost of classical numerical methods and notes that practical applications involve variations in coefficients, forcing terms, and boundary conditions across physical environments. That maps directly onto simulation-heavy workflows — engineering design sweeps, digital twins, and surrogate models that must remain reliable when operating conditions move outside the training range. The reported OOD behavior (e.g., 0.0074→0.0084 on Combined) speaks to the cost of retraining or re-running a solver when parameters change.

Future Directions

  • Extending the structured decomposition to more parameter families. PSP is demonstrated with a shared evolution term plus parameter-scaled responses (formulated generally over M parameters, and instantiated with a single viscosity-like scalar for Navier–Stokes). How well the decomposition scales to settings with many interacting governing parameters is an open question the paper does not resolve.

  • Testing beyond the evaluated regimes. The reported OOD splits move parameters moderately outside training ranges (e.g., Wave-2D c from [100, 500] to [500, 550]). Whether the gains hold for far larger distribution shifts, and how the ID-to-OOD degradation behaves at greater distances, is not reported.

  • Broadening the pretraining corpus. The encoder is pretrained and frozen per setup in this study; whether a single large-scale multi-physics predictive pretraining run could serve as a reusable state space across PDE families (analogous to foundation-model approaches) is not explored.

  • Characterizing when reconstruction-oriented learning remains preferable. The authors attribute the Advection gap to translation-dominated dynamics where shape change is minimal. A principled criterion for choosing between predictive and reconstruction-based latent learning per dynamical regime is left open.

Note: the paper does not contain a dedicated future-work section; the items above are open questions raised by its results and framing.

Target Audience

Researchers and graduate students in scientific machine learning, neural operators, and AI for physical simulation; practitioners building surrogate or latent-dynamics models who need forecasting that survives changes in governing parameters; and readers interested in self-supervised representation learning (JEPA-style methods) applied outside vision and video. The paper is best suited to readers comfortable with latent variable models, ODE integration, and PDE terminology — beginners will find the method sections accessible in outline but the experimental framing demanding.

Authors’ abstract

Physical trajectories contain more than snapshots of a system: they also reveal how its states evolve under governing conditions. However, representation learning for parametric partial differential equations (PDEs) has largely relied on reconstruction-based objectives that emphasize recovering observed physical fields. In this paper, we investigate predictive representation pretraining as an alternative to reconstruction-based learning. We find that predictive representations preserve rich physical information, yet this advantage alone does not ensure accurate field evolution. Based on these observations, we introduce PDE-JEPA for parametric PDE dynamics. Specifically, we first train an encoder using a masked-latent prediction to capture the underlying regularities of PDE dynamics. To explicitly adapt the pretrained representation toward a more dynamics-aligned state space, we then introduce a geometry projector that aligns latent trajectory geometry with the evolution geometry of physical fields. Finally, building on this geometry-aligned latent space, we further develop a physics-structured latent predictor that decomposes the dynamics into parameter-independent evolution and parameter-dependent response components. Extensive experiments on nine widely used PDE benchmarks demonstrate that our framework outperforms existing state-of-the-art methods by an average of 33.4\% in-distribution, while achieving an average improvement of 51.4\% when extrapolating to unseen governing parameters. The project page is available \href{https://tanpig-x.github.io/PDE-JEPA/}{here}.

Read the original paper