Skip to content

Exponax Integration

Dataset generation using Exponax spectral PDE solvers.

Exponax provides pseudo-spectral solvers for a range of PDEs on periodic domains. FlowPDE wraps these to generate training data for flow-based models.

All generators subclass ExponaxDatasetGenerator, which provides shared observation augmentation, JAX→torch conversion, statistics, and dataset wrapping. Every dataset emits {'input': condition, 'target': solution}.

Base Generator

ExponaxDatasetGenerator

Base class for concise Exponax problem generators.

Source code in flowpde/datasets/exponax/generator.py
class ExponaxDatasetGenerator:
    """Base class for concise Exponax problem generators."""

    config: GenerationConfig
    dataset_cls = PDEDataset

    def __init__(self, config: Optional[GenerationConfig] = None, **kwargs):
        if config is None:
            config = self.config_cls(**kwargs)
        self.config = config

    def resolve_run(self, num_samples: Optional[int], seed: Optional[int]):
        cfg = self.config
        n = num_samples or cfg.num_samples
        s = seed if seed is not None else cfg.seed
        return cfg, n, s

    def validate_problem(self, problem: str) -> None:
        if problem not in {"forward", "inverse"}:
            raise ValueError("problem must be 'forward' or 'inverse'")

    def to_torch_data(self, arrays: Dict[str, Any]) -> Dict[str, torch.Tensor]:
        return {
            key: jax_to_torch(value, device=self.config.torch_device)
            for key, value in arrays.items()
            if value is not None
        }

    def apply_observation_augmentation(
        self,
        data: Dict[str, torch.Tensor],
        *,
        observation_key: str,
        problem: Literal["forward", "inverse"],
        n: int,
        seed: int,
    ) -> Dict[str, Any]:
        """
        Add observation noise and/or spatial masking, in place.

        Returns:
            A dict to forward to `wrap_dataset`. It carries `'clean_stats'`
            (pass as `extra_stats`) and `'observation_key'` (pass inside
            `extra_metadata`), both empty when no augmentation applied.
        """
        cfg = self.config
        if problem != "inverse":
            return {}
        if cfg.obs_noise_std <= 0.0 and cfg.obs_mask_fraction >= 1.0:
            return {}

        # Statistics must describe the *clean* field. Fitting them on the
        # masked observation deflates the std by roughly
        # sqrt(obs_mask_fraction) and shifts the mean, so the same normalizer
        # would put clean data on a different scale.
        clean_stats = {observation_key: compute_normalization_stats(data[observation_key])}

        if cfg.obs_noise_std > 0.0:
            noise_gen = torch.Generator().manual_seed(seed + 1)
            noise = torch.randn(
                data[observation_key].shape,
                generator=noise_gen,
                dtype=data[observation_key].dtype,
            ).to(device=data[observation_key].device)
            data[observation_key] = (
                data[observation_key] + cfg.obs_noise_std * noise
            )

        if cfg.obs_mask_fraction < 1.0:
            spatial_shape = (cfg.num_points,) * cfg.num_spatial_dims
            mask_gen = torch.Generator().manual_seed(seed)
            obs_mask = (
                torch.rand((n, 1) + spatial_shape, generator=mask_gen)
                < cfg.obs_mask_fraction
            ).float().to(device=cfg.torch_device)
            data[observation_key] = data[observation_key] * obs_mask
            data["obs_mask"] = obs_mask

        return {
            "clean_stats": clean_stats,
            "observation_key": observation_key,
        }

    def wrap_dataset(
        self,
        data: Dict[str, torch.Tensor],
        *,
        problem: Literal["forward", "inverse"],
        extra_stats: Optional[Dict[str, Dict[str, float]]] = None,
        extra_metadata: Optional[Dict[str, Any]] = None,
        **dataset_kwargs,
    ):
        stats = {
            key: compute_normalization_stats(value)
            for key, value in data.items()
            if torch.is_tensor(value) and key != "obs_mask"
        }
        if extra_stats:
            stats.update(extra_stats)

        metadata = {
            "stats": stats,
            "config": self.config.to_dict(),
        }
        # Generation diagnostics (e.g. solver residuals) belong in metadata,
        # not in `data`: they are not fields and must not be normalized.
        if extra_metadata:
            metadata.update(extra_metadata)
        return self.dataset_cls(
            data,
            problem=problem,
            metadata=metadata,
            **dataset_kwargs,
        )

    def print_summary(self, name: str, n: int, details: str = "") -> None:
        cfg = self.config
        spatial = "x".join([str(cfg.num_points)] * cfg.num_spatial_dims)
        suffix = f", {details}" if details else ""
        logger.info(
            "Generated %s %dD dataset: %d samples, %s grid%s",
            name, cfg.num_spatial_dims, n, spatial, suffix,
        )

apply_observation_augmentation(data, *, observation_key, problem, n, seed)

Add observation noise and/or spatial masking, in place.

Returns:

Type Description
Dict[str, Any]

A dict to forward to wrap_dataset. It carries 'clean_stats'

Dict[str, Any]

(pass as extra_stats) and 'observation_key' (pass inside

Dict[str, Any]

extra_metadata), both empty when no augmentation applied.

Source code in flowpde/datasets/exponax/generator.py
def apply_observation_augmentation(
    self,
    data: Dict[str, torch.Tensor],
    *,
    observation_key: str,
    problem: Literal["forward", "inverse"],
    n: int,
    seed: int,
) -> Dict[str, Any]:
    """
    Add observation noise and/or spatial masking, in place.

    Returns:
        A dict to forward to `wrap_dataset`. It carries `'clean_stats'`
        (pass as `extra_stats`) and `'observation_key'` (pass inside
        `extra_metadata`), both empty when no augmentation applied.
    """
    cfg = self.config
    if problem != "inverse":
        return {}
    if cfg.obs_noise_std <= 0.0 and cfg.obs_mask_fraction >= 1.0:
        return {}

    # Statistics must describe the *clean* field. Fitting them on the
    # masked observation deflates the std by roughly
    # sqrt(obs_mask_fraction) and shifts the mean, so the same normalizer
    # would put clean data on a different scale.
    clean_stats = {observation_key: compute_normalization_stats(data[observation_key])}

    if cfg.obs_noise_std > 0.0:
        noise_gen = torch.Generator().manual_seed(seed + 1)
        noise = torch.randn(
            data[observation_key].shape,
            generator=noise_gen,
            dtype=data[observation_key].dtype,
        ).to(device=data[observation_key].device)
        data[observation_key] = (
            data[observation_key] + cfg.obs_noise_std * noise
        )

    if cfg.obs_mask_fraction < 1.0:
        spatial_shape = (cfg.num_points,) * cfg.num_spatial_dims
        mask_gen = torch.Generator().manual_seed(seed)
        obs_mask = (
            torch.rand((n, 1) + spatial_shape, generator=mask_gen)
            < cfg.obs_mask_fraction
        ).float().to(device=cfg.torch_device)
        data[observation_key] = data[observation_key] * obs_mask
        data["obs_mask"] = obs_mask

    return {
        "clean_stats": clean_stats,
        "observation_key": observation_key,
    }

Poisson Generator

Source → solution pairs for the Poisson equation \(\nabla^2 u = f\) (1D/2D/3D).

PoissonGenerator

Bases: ExponaxDatasetGenerator

Generate Poisson equation datasets using Exponax.

Workflow
  1. Create simple smooth sine/cosine source terms f
  2. Solve \(\nabla^2 u = f\) via 'exponax.poisson.Poisson'
  3. Convert JAX arrays -> PyTorch tensors
  4. Wrap in a 'PDEDataset'

Parameters:

Name Type Description Default
config Optional[GenerationConfig]

A 'PoissonConfig' instance. All keyword arguments are forwarded to 'PoissonConfig' if config is None.

None
Source code in flowpde/datasets/exponax/poisson.py
class PoissonGenerator(ExponaxDatasetGenerator):
    r"""
    Generate Poisson equation datasets using Exponax.

    Workflow:
        1. Create simple smooth sine/cosine source terms *f*
        2. Solve $\nabla^2 u = f$ via 'exponax.poisson.Poisson'
        3. Convert JAX arrays -> PyTorch tensors
        4. Wrap in a 'PDEDataset'

    Args:
        config: A 'PoissonConfig' instance.  All keyword arguments
                are forwarded to 'PoissonConfig' if *config* is None.
    """

    config_cls = PoissonConfig

    def generate(
        self,
        num_samples: Optional[int] = None,
        seed: Optional[int] = None,
        problem: Literal['forward', 'inverse'] = 'forward',
    ) -> PDEDataset:
        """
        Generate a Poisson dataset.

        Args:
            num_samples: Override 'config.num_samples'.
            seed: Override 'config.seed'.
            problem: 'forward' (data->solution) or 'inverse' (solution->data).

        Returns:
            A `PDEDataset` with keys `'source'` and `'solution'`.
        """
        cfg, n, s = self.resolve_run(num_samples, seed)
        self.validate_problem(problem)

        solver = ex.poisson.Poisson(
            num_spatial_dims=cfg.num_spatial_dims,
            domain_extent=cfg.domain_extent,
            num_points=cfg.num_points,
            order=cfg.order,
        )

        sources = sample_sine_fields(
            jax.random.PRNGKey(s),
            n=n,
            num_spatial_dims=cfg.num_spatial_dims,
            num_points=cfg.num_points,
            domain_extent=cfg.domain_extent,
            num_terms=cfg.source_num_terms,
            max_mode=cfg.source_max_mode,
        )
        solutions = jax.vmap(solver)(sources)

        data = self.to_torch_data({
            'source': sources,
            'solution': solutions,
        })
        augmentation = self.apply_observation_augmentation(
            data,
            observation_key='solution',
            problem=problem,
            n=n,
            seed=s,
        )

        self.print_summary("Poisson", n)
        return self.wrap_dataset(
            data,
            problem=problem,
            extra_stats=augmentation.pop("clean_stats", None),
            extra_metadata=augmentation,
        )

generate(num_samples=None, seed=None, problem='forward')

Generate a Poisson dataset.

Parameters:

Name Type Description Default
num_samples Optional[int]

Override 'config.num_samples'.

None
seed Optional[int]

Override 'config.seed'.

None
problem Literal['forward', 'inverse']

'forward' (data->solution) or 'inverse' (solution->data).

'forward'

Returns:

Type Description
PDEDataset

A PDEDataset with keys 'source' and 'solution'.

Source code in flowpde/datasets/exponax/poisson.py
def generate(
    self,
    num_samples: Optional[int] = None,
    seed: Optional[int] = None,
    problem: Literal['forward', 'inverse'] = 'forward',
) -> PDEDataset:
    """
    Generate a Poisson dataset.

    Args:
        num_samples: Override 'config.num_samples'.
        seed: Override 'config.seed'.
        problem: 'forward' (data->solution) or 'inverse' (solution->data).

    Returns:
        A `PDEDataset` with keys `'source'` and `'solution'`.
    """
    cfg, n, s = self.resolve_run(num_samples, seed)
    self.validate_problem(problem)

    solver = ex.poisson.Poisson(
        num_spatial_dims=cfg.num_spatial_dims,
        domain_extent=cfg.domain_extent,
        num_points=cfg.num_points,
        order=cfg.order,
    )

    sources = sample_sine_fields(
        jax.random.PRNGKey(s),
        n=n,
        num_spatial_dims=cfg.num_spatial_dims,
        num_points=cfg.num_points,
        domain_extent=cfg.domain_extent,
        num_terms=cfg.source_num_terms,
        max_mode=cfg.source_max_mode,
    )
    solutions = jax.vmap(solver)(sources)

    data = self.to_torch_data({
        'source': sources,
        'solution': solutions,
    })
    augmentation = self.apply_observation_augmentation(
        data,
        observation_key='solution',
        problem=problem,
        n=n,
        seed=s,
    )

    self.print_summary("Poisson", n)
    return self.wrap_dataset(
        data,
        problem=problem,
        extra_stats=augmentation.pop("clean_stats", None),
        extra_metadata=augmentation,
    )

PoissonConfig dataclass

Bases: GenerationConfig

Configuration specific to the Poisson equation.

Attributes:

Name Type Description
order int

Order of the Poisson operator (default 2 -> Laplacian).

source_num_terms int

Number of sine/cosine terms per source field.

source_max_mode int

Largest integer wavenumber sampled per dimension.

Source code in flowpde/datasets/exponax/poisson.py
@dataclass
class PoissonConfig(GenerationConfig):
    """
    Configuration specific to the Poisson equation.

    Attributes:
        order: Order of the Poisson operator (default 2 -> Laplacian).
        source_num_terms: Number of sine/cosine terms per source field.
        source_max_mode: Largest integer wavenumber sampled per dimension.
    """
    order: int = 2
    source_num_terms: int = 3
    source_max_mode: int = 3

Burgers Generator

Initial-condition → final-state pairs (optionally full trajectories) for the Burgers equation \(\partial_t u + u \cdot \nabla u = \nu \nabla^2 u\) (1D/2D).

BurgersGenerator

Bases: ExponaxDatasetGenerator

Generate Burgers equation datasets using Exponax.

Workflow
  1. Create simple smooth sine/cosine initial conditions
  2. Step forward using exponax.stepper.Burgers (optionally via exponax.rollout for trajectories)
  3. Convert JAX arrays → PyTorch tensors
  4. Wrap in a PDEDataset

Parameters:

Name Type Description Default
config Optional[GenerationConfig]

A BurgersConfig instance. Keyword arguments are forwarded to BurgersConfig if config is None.

None
Source code in flowpde/datasets/exponax/burgers.py
class BurgersGenerator(ExponaxDatasetGenerator):
    """
    Generate Burgers equation datasets using Exponax.

    Workflow:
        1. Create simple smooth sine/cosine initial conditions
        2. Step forward using `exponax.stepper.Burgers` (optionally via
           `exponax.rollout` for trajectories)
        3. Convert JAX arrays → PyTorch tensors
        4. Wrap in a `PDEDataset`

    Args:
        config: A `BurgersConfig` instance.  Keyword arguments are
                forwarded to `BurgersConfig` if *config* is None.
    """

    config_cls = BurgersConfig

    def generate(
        self,
        num_samples: Optional[int] = None,
        seed: Optional[int] = None,
        problem: Literal['forward', 'inverse'] = 'forward',
    ) -> PDEDataset:
        """
        Generate a Burgers dataset.

        Args:
            num_samples: Override `config.num_samples`.
            seed: Override `config.seed`.
            problem: `'forward'` (IC→final) or `'inverse'`.

        Returns:
            A `PDEDataset` with keys `'initial'` and `'final'`
            (and optionally `'trajectory'`).
        """
        cfg, n, s = self.resolve_run(num_samples, seed)
        self.validate_problem(problem)

        key = jax.random.PRNGKey(s)
        field_key, nu_key = jax.random.split(key, 2)

        nus = log_uniform(
            nu_key,
            n=n,
            min_value=cfg.diffusivity_min,
            max_value=cfg.diffusivity_max,
        )
        ics = sample_sine_fields(
            field_key,
            n=n,
            num_spatial_dims=cfg.num_spatial_dims,
            num_points=cfg.num_points,
            domain_extent=cfg.domain_extent,
            num_terms=cfg.ic_num_terms,
            max_mode=cfg.ic_max_mode,
        )

        # --- Time integration (per-sample ν) ---
        # The stepper is constructed inside the vmapped function so that JAX
        # traces through the ETDRK coefficient computation with each sample's
        # individual nu value.
        def step_one(ic, nu):
            stepper = ex.stepper.Burgers(
                num_spatial_dims=cfg.num_spatial_dims,
                domain_extent=cfg.domain_extent,
                num_points=cfg.num_points,
                dt=cfg.dt,
                diffusivity=nu,
                convection_scale=cfg.convection_scale,
                single_channel=cfg.single_channel,
            )
            return ex.RepeatedStepper(stepper, cfg.num_steps)(ic)

        def rollout_one(ic, nu):
            stepper = ex.stepper.Burgers(
                num_spatial_dims=cfg.num_spatial_dims,
                domain_extent=cfg.domain_extent,
                num_points=cfg.num_points,
                dt=cfg.dt,
                diffusivity=nu,
                convection_scale=cfg.convection_scale,
                single_channel=cfg.single_channel,
            )
            return ex.rollout(stepper, cfg.num_steps, include_init=True)(ic)

        if cfg.store_trajectory:
            trajectories = jax.vmap(rollout_one)(ics, nus)  # (N, T+1, C, *spatial)
            finals = trajectories[:, -1]
        else:
            finals = jax.vmap(step_one)(ics, nus)  # (N, C, *spatial)
            trajectories = None

        data = self.to_torch_data({
            'initial': ics,
            'final': finals,
            'diffusivity': nus,
            'trajectory': trajectories,
        })
        augmentation = self.apply_observation_augmentation(
            data,
            observation_key='final',
            problem=problem,
            n=n,
            seed=s,
        )

        self.print_summary(
            "Burgers",
            n,
            details=(
                f"{cfg.num_steps} steps (dt={cfg.dt}), "
                f"nu ~ LogUniform({cfg.diffusivity_min:.0e}, "
                f"{cfg.diffusivity_max:.0e})"
            ),
        )
        return self.wrap_dataset(
            data,
            problem=problem,
            extra_stats=augmentation.pop("clean_stats", None),
            extra_metadata=augmentation,
        )

generate(num_samples=None, seed=None, problem='forward')

Generate a Burgers dataset.

Parameters:

Name Type Description Default
num_samples Optional[int]

Override config.num_samples.

None
seed Optional[int]

Override config.seed.

None
problem Literal['forward', 'inverse']

'forward' (IC→final) or 'inverse'.

'forward'

Returns:

Type Description
PDEDataset

A PDEDataset with keys 'initial' and 'final'

PDEDataset

(and optionally 'trajectory').

Source code in flowpde/datasets/exponax/burgers.py
def generate(
    self,
    num_samples: Optional[int] = None,
    seed: Optional[int] = None,
    problem: Literal['forward', 'inverse'] = 'forward',
) -> PDEDataset:
    """
    Generate a Burgers dataset.

    Args:
        num_samples: Override `config.num_samples`.
        seed: Override `config.seed`.
        problem: `'forward'` (IC→final) or `'inverse'`.

    Returns:
        A `PDEDataset` with keys `'initial'` and `'final'`
        (and optionally `'trajectory'`).
    """
    cfg, n, s = self.resolve_run(num_samples, seed)
    self.validate_problem(problem)

    key = jax.random.PRNGKey(s)
    field_key, nu_key = jax.random.split(key, 2)

    nus = log_uniform(
        nu_key,
        n=n,
        min_value=cfg.diffusivity_min,
        max_value=cfg.diffusivity_max,
    )
    ics = sample_sine_fields(
        field_key,
        n=n,
        num_spatial_dims=cfg.num_spatial_dims,
        num_points=cfg.num_points,
        domain_extent=cfg.domain_extent,
        num_terms=cfg.ic_num_terms,
        max_mode=cfg.ic_max_mode,
    )

    # --- Time integration (per-sample ν) ---
    # The stepper is constructed inside the vmapped function so that JAX
    # traces through the ETDRK coefficient computation with each sample's
    # individual nu value.
    def step_one(ic, nu):
        stepper = ex.stepper.Burgers(
            num_spatial_dims=cfg.num_spatial_dims,
            domain_extent=cfg.domain_extent,
            num_points=cfg.num_points,
            dt=cfg.dt,
            diffusivity=nu,
            convection_scale=cfg.convection_scale,
            single_channel=cfg.single_channel,
        )
        return ex.RepeatedStepper(stepper, cfg.num_steps)(ic)

    def rollout_one(ic, nu):
        stepper = ex.stepper.Burgers(
            num_spatial_dims=cfg.num_spatial_dims,
            domain_extent=cfg.domain_extent,
            num_points=cfg.num_points,
            dt=cfg.dt,
            diffusivity=nu,
            convection_scale=cfg.convection_scale,
            single_channel=cfg.single_channel,
        )
        return ex.rollout(stepper, cfg.num_steps, include_init=True)(ic)

    if cfg.store_trajectory:
        trajectories = jax.vmap(rollout_one)(ics, nus)  # (N, T+1, C, *spatial)
        finals = trajectories[:, -1]
    else:
        finals = jax.vmap(step_one)(ics, nus)  # (N, C, *spatial)
        trajectories = None

    data = self.to_torch_data({
        'initial': ics,
        'final': finals,
        'diffusivity': nus,
        'trajectory': trajectories,
    })
    augmentation = self.apply_observation_augmentation(
        data,
        observation_key='final',
        problem=problem,
        n=n,
        seed=s,
    )

    self.print_summary(
        "Burgers",
        n,
        details=(
            f"{cfg.num_steps} steps (dt={cfg.dt}), "
            f"nu ~ LogUniform({cfg.diffusivity_min:.0e}, "
            f"{cfg.diffusivity_max:.0e})"
        ),
    )
    return self.wrap_dataset(
        data,
        problem=problem,
        extra_stats=augmentation.pop("clean_stats", None),
        extra_metadata=augmentation,
    )

BurgersConfig dataclass

Bases: GenerationConfig

Configuration specific to the Burgers equation.

Attributes:

Name Type Description
dt float

Time-step size for the ETDRK stepper.

num_steps int

Number of time steps to advance.

diffusivity_min float

Lower bound of the log-uniform per-sample viscosity ν. Each sample gets an independently drawn ν from LogUniform(diffusivity_min, diffusivity_max), so the dataset spans multiple physical regimes (smooth ↔ shock-dominated).

diffusivity_max float

Upper bound of the log-uniform per-sample viscosity.

convection_scale float

Scaling of the nonlinear convection term.

single_channel bool

If True, use a single-channel formulation (scalar Burgers) regardless of spatial dimension.

ic_num_terms int

Number of sine/cosine terms per initial condition.

ic_max_mode int

Largest integer wavenumber sampled per dimension.

store_trajectory bool

If True, keep the full rollout trajectory in the dataset (key 'trajectory').

Source code in flowpde/datasets/exponax/burgers.py
@dataclass
class BurgersConfig(GenerationConfig):
    """
    Configuration specific to the Burgers equation.

    Attributes:
        dt: Time-step size for the ETDRK stepper.
        num_steps: Number of time steps to advance.
        diffusivity_min: Lower bound of the log-uniform per-sample viscosity
            ν.  Each sample gets an independently drawn ν from
            LogUniform(diffusivity_min, diffusivity_max), so the dataset
            spans multiple physical regimes (smooth ↔ shock-dominated).
        diffusivity_max: Upper bound of the log-uniform per-sample viscosity.
        convection_scale: Scaling of the nonlinear convection term.
        single_channel: If True, use a single-channel formulation
            (scalar Burgers) regardless of spatial dimension.
        ic_num_terms: Number of sine/cosine terms per initial condition.
        ic_max_mode: Largest integer wavenumber sampled per dimension.
        store_trajectory: If True, keep the full rollout trajectory
            in the dataset (key `'trajectory'`).
    """
    num_spatial_dims: int = 1
    domain_extent: float = 1.0
    dt: float = 0.001
    num_steps: int = 400
    diffusivity_min: float = 1e-2
    diffusivity_max: float = 1e-2
    convection_scale: float = 1.0
    single_channel: bool = False
    ic_num_terms: int = 3
    ic_max_mode: int = 4
    store_trajectory: bool = False

Darcy Generator

\((\kappa, f) \rightarrow u\) for variable-coefficient Poisson \(-\nabla \cdot (\kappa \nabla u) = f\). generate() accepts inverse_mode in {'both', 'coefficient', 'source'}.

DarcyGenerator

Bases: ExponaxDatasetGenerator

Generate Darcy-flow / variable-coefficient Poisson datasets.

Workflow:

1. Draw per-sample log-normal $\kappa$ fields from a GRF.
2. Draw per-sample smooth Fourier source fields.
3. Solve $-\nabla\cdot(\kappa\,\nabla u)=f$ via FD + fixed-step CG for each sample.
4. Optionally apply additive noise and/or spatial masking to the
   solution field (for inverse-problem datasets).
5. Convert to PyTorch tensors and return a `DarcyDataset`.

Parameters:

Name Type Description Default
config Optional[GenerationConfig]

A DarcyConfig instance. Keyword arguments are forwarded to DarcyConfig if config is None.

None

Example

gen = DarcyGenerator(num_points=64)
train = gen.generate(num_samples=1000, seed=0)
test  = gen.generate(num_samples=200,  seed=1)
Source code in flowpde/datasets/exponax/darcy.py
class DarcyGenerator(ExponaxDatasetGenerator):
    r"""
    Generate Darcy-flow / variable-coefficient Poisson datasets.

    Workflow:

        1. Draw per-sample log-normal $\kappa$ fields from a GRF.
        2. Draw per-sample smooth Fourier source fields.
        3. Solve $-\nabla\cdot(\kappa\,\nabla u)=f$ via FD + fixed-step CG for each sample.
        4. Optionally apply additive noise and/or spatial masking to the
           solution field (for inverse-problem datasets).
        5. Convert to PyTorch tensors and return a `DarcyDataset`.

    Args:
        config: A `DarcyConfig` instance.  Keyword arguments are forwarded
                to `DarcyConfig` if *config* is `None`.

    **Example**

    ```python
    gen = DarcyGenerator(num_points=64)
    train = gen.generate(num_samples=1000, seed=0)
    test  = gen.generate(num_samples=200,  seed=1)
    ```
    """

    config_cls = DarcyConfig
    dataset_cls = DarcyDataset

    def generate(
        self,
        num_samples: Optional[int] = None,
        seed: Optional[int] = None,
        problem: Literal['forward', 'inverse'] = 'forward',
        inverse_mode: Literal['both', 'coefficient', 'source'] = 'both',
    ) -> DarcyDataset:
        """
        Generate a Darcy-flow dataset.

        Args:
            num_samples: Override `config.num_samples`.
            seed:        Override `config.seed`.
            problem:     `'forward'` (κ,f → u) or `'inverse'`.
            inverse_mode: Inverse mapping to use:
                `'both'` for u → (κ,f), `'coefficient'` for (u,f) → κ,
                or `'source'` for (u,κ) → f.

        Returns:
            A `DarcyDataset`.
        """
        cfg, n, s = self.resolve_run(num_samples, seed)
        self.validate_problem(problem)
        N   = cfg.num_points
        d   = cfg.num_spatial_dims
        L   = cfg.domain_extent

        if d not in (1, 2):
            raise ValueError(
                f"DarcyGenerator supports num_spatial_dims ∈ {{1, 2}}; got {d}."
            )

        # Vertex-centred grid: N points on [0, L], spacing h = L/(N-1)
        h = L / (N - 1)

        # Split all keys upfront for full reproducibility
        key = jax.random.PRNGKey(s)
        key, kappa_key, f_key = jax.random.split(key, 3)

        # ── 1. κ fields: per-sample log-normal GRF ────────────────────────
        kappa_sample_keys = jax.random.split(kappa_key, n)

        def make_kappa(k):
            grf = (
                _grf_1d(k, N, cfg.kappa_alpha, cfg.kappa_tau) if d == 1
                else _grf_2d(k, N, cfg.kappa_alpha, cfg.kappa_tau)
            )
            # Standardise then scale → kappa_scale controls log-contrast
            grf   = (grf - grf.mean()) / (grf.std() + 1e-6) * cfg.kappa_scale
            kappa = jnp.maximum(jnp.exp(grf), cfg.kappa_min)
            return kappa[jnp.newaxis]               # (1, *spatial)

        kappas = jax.vmap(make_kappa)(kappa_sample_keys)   # (n, 1, *spatial)

        # ── 2. f fields: smooth fixed-cutoff Fourier series ───────────────
        sources = sample_fourier_fields(
            f_key,
            n=n,
            cfg=cfg,
            field=FourierFieldConfig(
                cutoff=cfg.f_cutoff,
                amplitude_min=cfg.f_amplitude_min,
                amplitude_max=cfg.f_amplitude_max,
                normalize=True,
            ),
        )

        # ── 3. Solve for each sample ───────────────────────
        if d == 1:
            def solve_one(kappa, f):
                return _solve_one_1d(kappa, f, h, N, cfg.cg_steps)
        else:
            def solve_one(kappa, f):
                return _solve_one_2d(kappa, f, h, N, cfg.cg_steps)

        solutions, cg_residuals = jax.vmap(solve_one)(kappas, sources)
        # solutions: (n, 1, *spatial);  cg_residuals: (n,)

        worst_residual = float(jnp.max(cg_residuals))
        if cfg.cg_tolerance is not None and worst_residual > cfg.cg_tolerance:
            raise RuntimeError(
                f"Darcy CG did not converge: worst relative residual "
                f"{worst_residual:.3e} over {n} samples exceeds "
                f"cg_tolerance={cfg.cg_tolerance:.1e} after "
                f"{cfg.cg_steps} iterations. These solutions do not solve the "
                f"PDE and must not be used as ground truth. Raise cg_steps "
                f"(unpreconditioned CG needs more of them as the grid or the "
                f"contrast in kappa grows), or relax cg_tolerance if you "
                f"genuinely accept this error level."
            )

        # ── 4. Convert, optionally augment inverse observations, and wrap ──
        data = self.to_torch_data({
            'kappa': kappas,
            'source': sources,
            'solution': solutions,
        })
        augmentation = self.apply_observation_augmentation(
            data,
            observation_key='solution',
            problem=problem,
            n=n,
            seed=s,
        )

        spatial_str = 'x'.join([str(N)] * d)
        logger.info(
            "Generated Darcy %dD dataset: %d samples, %s grid, "
            "κ ~ LogNormal(α=%s, τ=%s, scale=%s), CG steps=%s, "
            "worst CG residual=%.2e",
            d, n, spatial_str, cfg.kappa_alpha, cfg.kappa_tau,
            cfg.kappa_scale, cfg.cg_steps, worst_residual,
        )

        return self.wrap_dataset(
            data,
            problem=problem,
            inverse_mode=inverse_mode,
            extra_stats=augmentation.pop("clean_stats", None),
            extra_metadata={
                **augmentation,
                "cg_residual_max": worst_residual,
                "cg_residual_mean": float(jnp.mean(cg_residuals)),
                "cg_steps": cfg.cg_steps,
            },
        )

generate(num_samples=None, seed=None, problem='forward', inverse_mode='both')

Generate a Darcy-flow dataset.

Parameters:

Name Type Description Default
num_samples Optional[int]

Override config.num_samples.

None
seed Optional[int]

Override config.seed.

None
problem Literal['forward', 'inverse']

'forward' (κ,f → u) or 'inverse'.

'forward'
inverse_mode Literal['both', 'coefficient', 'source']

Inverse mapping to use: 'both' for u → (κ,f), 'coefficient' for (u,f) → κ, or 'source' for (u,κ) → f.

'both'

Returns:

Type Description
DarcyDataset

A DarcyDataset.

Source code in flowpde/datasets/exponax/darcy.py
def generate(
    self,
    num_samples: Optional[int] = None,
    seed: Optional[int] = None,
    problem: Literal['forward', 'inverse'] = 'forward',
    inverse_mode: Literal['both', 'coefficient', 'source'] = 'both',
) -> DarcyDataset:
    """
    Generate a Darcy-flow dataset.

    Args:
        num_samples: Override `config.num_samples`.
        seed:        Override `config.seed`.
        problem:     `'forward'` (κ,f → u) or `'inverse'`.
        inverse_mode: Inverse mapping to use:
            `'both'` for u → (κ,f), `'coefficient'` for (u,f) → κ,
            or `'source'` for (u,κ) → f.

    Returns:
        A `DarcyDataset`.
    """
    cfg, n, s = self.resolve_run(num_samples, seed)
    self.validate_problem(problem)
    N   = cfg.num_points
    d   = cfg.num_spatial_dims
    L   = cfg.domain_extent

    if d not in (1, 2):
        raise ValueError(
            f"DarcyGenerator supports num_spatial_dims ∈ {{1, 2}}; got {d}."
        )

    # Vertex-centred grid: N points on [0, L], spacing h = L/(N-1)
    h = L / (N - 1)

    # Split all keys upfront for full reproducibility
    key = jax.random.PRNGKey(s)
    key, kappa_key, f_key = jax.random.split(key, 3)

    # ── 1. κ fields: per-sample log-normal GRF ────────────────────────
    kappa_sample_keys = jax.random.split(kappa_key, n)

    def make_kappa(k):
        grf = (
            _grf_1d(k, N, cfg.kappa_alpha, cfg.kappa_tau) if d == 1
            else _grf_2d(k, N, cfg.kappa_alpha, cfg.kappa_tau)
        )
        # Standardise then scale → kappa_scale controls log-contrast
        grf   = (grf - grf.mean()) / (grf.std() + 1e-6) * cfg.kappa_scale
        kappa = jnp.maximum(jnp.exp(grf), cfg.kappa_min)
        return kappa[jnp.newaxis]               # (1, *spatial)

    kappas = jax.vmap(make_kappa)(kappa_sample_keys)   # (n, 1, *spatial)

    # ── 2. f fields: smooth fixed-cutoff Fourier series ───────────────
    sources = sample_fourier_fields(
        f_key,
        n=n,
        cfg=cfg,
        field=FourierFieldConfig(
            cutoff=cfg.f_cutoff,
            amplitude_min=cfg.f_amplitude_min,
            amplitude_max=cfg.f_amplitude_max,
            normalize=True,
        ),
    )

    # ── 3. Solve for each sample ───────────────────────
    if d == 1:
        def solve_one(kappa, f):
            return _solve_one_1d(kappa, f, h, N, cfg.cg_steps)
    else:
        def solve_one(kappa, f):
            return _solve_one_2d(kappa, f, h, N, cfg.cg_steps)

    solutions, cg_residuals = jax.vmap(solve_one)(kappas, sources)
    # solutions: (n, 1, *spatial);  cg_residuals: (n,)

    worst_residual = float(jnp.max(cg_residuals))
    if cfg.cg_tolerance is not None and worst_residual > cfg.cg_tolerance:
        raise RuntimeError(
            f"Darcy CG did not converge: worst relative residual "
            f"{worst_residual:.3e} over {n} samples exceeds "
            f"cg_tolerance={cfg.cg_tolerance:.1e} after "
            f"{cfg.cg_steps} iterations. These solutions do not solve the "
            f"PDE and must not be used as ground truth. Raise cg_steps "
            f"(unpreconditioned CG needs more of them as the grid or the "
            f"contrast in kappa grows), or relax cg_tolerance if you "
            f"genuinely accept this error level."
        )

    # ── 4. Convert, optionally augment inverse observations, and wrap ──
    data = self.to_torch_data({
        'kappa': kappas,
        'source': sources,
        'solution': solutions,
    })
    augmentation = self.apply_observation_augmentation(
        data,
        observation_key='solution',
        problem=problem,
        n=n,
        seed=s,
    )

    spatial_str = 'x'.join([str(N)] * d)
    logger.info(
        "Generated Darcy %dD dataset: %d samples, %s grid, "
        "κ ~ LogNormal(α=%s, τ=%s, scale=%s), CG steps=%s, "
        "worst CG residual=%.2e",
        d, n, spatial_str, cfg.kappa_alpha, cfg.kappa_tau,
        cfg.kappa_scale, cfg.cg_steps, worst_residual,
    )

    return self.wrap_dataset(
        data,
        problem=problem,
        inverse_mode=inverse_mode,
        extra_stats=augmentation.pop("clean_stats", None),
        extra_metadata={
            **augmentation,
            "cg_residual_max": worst_residual,
            "cg_residual_mean": float(jnp.mean(cg_residuals)),
            "cg_steps": cfg.cg_steps,
        },
    )

DarcyConfig dataclass

Bases: GenerationConfig

Configuration for the Darcy-flow / variable-coefficient Poisson generator.

Inherits all base fields from GenerationConfig (num_points, num_samples, seed, torch_device, obs_noise_std, obs_mask_fraction).

Attributes:

Name Type Description
num_spatial_dims int

Spatial dimension (1 or 2). Default 2.

domain_extent float

Side length of the square/interval domain. Default 1.0 ([0,1]^d), which is the standard Darcy benchmark domain.

kappa_alpha float

Spectral exponent \(\alpha\) for the \(\kappa\) GRF power spectrum \(S(\lvert k\rvert) \propto (\tau^2 + \lvert k\rvert^2)^{-\alpha}\). Higher → smoother permeability. \(\alpha = 2.0\) matches the original FNO Darcy benchmark.

kappa_tau float

GRF inverse correlation length \(\tau\). Higher → more oscillatory \(\kappa\) with smaller features.

kappa_scale float

Standard deviation of \(\log\kappa\) before exponentiation. Scale 1.0 gives \(\kappa \in [e^{-2}, e^{2}] \approx [0.14,\, 7.4]\) roughly. Increase for higher contrast between high- and low-permeability regions.

kappa_min float

Hard lower bound on \(\kappa\) (positivity / ellipticity floor).

f_cutoff int

Fourier cutoff for the random source term.

f_amplitude_min float

Minimum per-sample amplitude scaling for f.

f_amplitude_max float

Maximum per-sample amplitude scaling for f.

cg_steps int

Fixed number of conjugate-gradient iterations used to solve the linear system. The solver is unpreconditioned, so the count needed grows with the grid and with the contrast in κ. Measured on a 64×64 grid at the default κ: 100 steps leaves a relative residual of 2e-1 (the "solutions" are 20 % wrong), 500 leaves 8e-4, and 2000 converges to machine precision. Do not lower this for speed without checking cg_tolerance still passes.

cg_tolerance Optional[float]

Largest acceptable CG relative residual, checked over every generated sample. generate() raises if it is exceeded, so an under-converged dataset fails loudly instead of quietly becoming wrong ground truth. Set to None to skip the check.

Source code in flowpde/datasets/exponax/darcy.py
@dataclass
class DarcyConfig(GenerationConfig):
    r"""
    Configuration for the Darcy-flow / variable-coefficient Poisson generator.

    Inherits all base fields from `GenerationConfig` (`num_points`,
    `num_samples`, `seed`, `torch_device`, `obs_noise_std`,
    `obs_mask_fraction`).

    Attributes:
        num_spatial_dims: Spatial dimension (1 or 2).  Default 2.
        domain_extent: Side length of the square/interval domain.  Default
            1.0 ([0,1]^d), which is the standard Darcy benchmark domain.
        kappa_alpha: Spectral exponent $\alpha$ for the $\kappa$
            GRF power spectrum
            $S(\lvert k\rvert) \propto (\tau^2 + \lvert k\rvert^2)^{-\alpha}$.
            Higher → smoother permeability.
            $\alpha = 2.0$ matches the original FNO Darcy benchmark.
        kappa_tau: GRF inverse correlation length $\tau$.  Higher →
            more oscillatory $\kappa$ with smaller features.
        kappa_scale: Standard deviation of $\log\kappa$ before
            exponentiation.  Scale 1.0 gives
            $\kappa \in [e^{-2}, e^{2}] \approx [0.14,\, 7.4]$
            roughly.  Increase for higher contrast between high- and
            low-permeability regions.
        kappa_min: Hard lower bound on $\kappa$
            (positivity / ellipticity floor).
        f_cutoff: Fourier cutoff for the random source term.
        f_amplitude_min: Minimum per-sample amplitude scaling for f.
        f_amplitude_max: Maximum per-sample amplitude scaling for f.
        cg_steps: Fixed number of conjugate-gradient iterations used to solve
            the linear system.  The solver is unpreconditioned, so the count
            needed grows with the grid and with the contrast in κ.  Measured
            on a 64×64 grid at the default κ: 100 steps leaves a relative
            residual of 2e-1 (the "solutions" are 20 % wrong), 500 leaves
            8e-4, and 2000 converges to machine precision.  Do not lower this
            for speed without checking `cg_tolerance` still passes.
        cg_tolerance: Largest acceptable CG relative residual, checked over
            every generated sample.  `generate()` raises if it is exceeded,
            so an under-converged dataset fails loudly instead of quietly
            becoming wrong ground truth.  Set to `None` to skip the check.
    """
    # Domain — override base defaults for the standard Darcy setting
    num_spatial_dims: int   = 2
    domain_extent:    float = 1.0

    # κ field
    kappa_alpha: float = 2.0
    kappa_tau:   float = 3.0
    kappa_scale: float = 1.0
    kappa_min:   float = 0.1

    # f field
    f_cutoff:        int   = 8
    f_amplitude_min: float = 0.1
    f_amplitude_max: float = 5.0

    # Solver
    cg_steps: int = 2000
    cg_tolerance: Optional[float] = 1e-6

Datasets

PDEDataset

Bases: Dataset

PyTorch Dataset wrapping Exponax-generated PDE data.

Returns samples as {'input': ..., 'target': ...} where the semantics of input and target depend on the PDE type and the chosen problem direction:

  • Poisson (static): input=source, target=solution (forward) or input=solution, target=source (inverse).
  • Burgers (time-dependent): input=initial condition, target=final state (forward) or vice-versa (inverse).

When partial observations are enabled (obs_mask_fraction < 1.0), __getitem__ appends the observation mask to the conditioning input and additionally returns 'obs_mask': a float tensor of shape (1, *spatial) with 1 at observed locations and 0 elsewhere.

The dataset also stores normalization statistics and generation config for reference.

Source code in flowpde/datasets/exponax/base.py
class PDEDataset(Dataset):
    """
    PyTorch Dataset wrapping Exponax-generated PDE data.

    Returns samples as {'input': ..., 'target': ...} where the
    semantics of *input* and *target* depend on the PDE type and the
    chosen problem direction:

    * **Poisson** (static): input=source, target=solution (forward)
      or input=solution, target=source (inverse).
    * **Burgers** (time-dependent): input=initial condition,
      target=final state (forward) or vice-versa (inverse).

    When partial observations are enabled (`obs_mask_fraction < 1.0`),
    `__getitem__` appends the observation mask to the conditioning input
    and additionally returns `'obs_mask'`: a float tensor of shape
    `(1, *spatial)` with 1 at observed locations and 0 elsewhere.

    The dataset also stores normalization statistics and generation
    config for reference.
    """

    def __init__(
        self,
        data: Dict[str, torch.Tensor],
        problem: Literal['forward', 'inverse'] = 'forward',
        metadata: Optional[Dict[str, Any]] = None,
        normalizer: Optional[FieldNormalizer] = None,
    ):
        """
        Args:
            data: Dictionary with at least two tensor entries whose
                  keys identify the PDE fields (e.g. 'source' / 'solution'
                  for Poisson, 'initial' / 'final' for Burgers).
            problem: 'forward' maps natural data -> solution;
                     'inverse' reverses the mapping.
            metadata: Optional dict with 'stats', 'config', etc.
            normalizer: Optional `FieldNormalizer` applied to each PDE
                field on access.  Fit it on the training split and share
                the same instance with validation/test splits.
        """
        if problem not in {'forward', 'inverse'}:
            raise ValueError("problem must be 'forward' or 'inverse'")

        self.data = data
        self.problem = problem
        self.metadata = metadata or {}
        self.stats = self.metadata.get('stats', {})
        self.config = self.metadata.get('config', {})
        # Which raw field carries the (possibly masked) observation, if any.
        self.observation_key = self.metadata.get('observation_key')
        self.normalizer = normalizer

        self._setup_keys()


    def _setup_keys(self):
        """Determine input / target keys from the available data."""
        #Static PDEs (Poisson)
        if 'source' in self.data and 'solution' in self.data:
            if self.problem == 'forward':
                self.input_key = 'source'
                self.target_key = 'solution'
            else:
                self.input_key = 'solution'
                self.target_key = 'source'

        #Time-dependent PDEs (Burgers)
        elif 'initial' in self.data and 'final' in self.data:
            if self.problem == 'forward':
                self.input_key = 'initial'
                self.target_key = 'final'
            else:
                self.input_key = 'final'
                self.target_key = 'initial'

        else:
            raise ValueError(
                f"Unrecognized data format. Keys: {list(self.data.keys())}. "
                "Expected ('source', 'solution') or ('initial', 'final')."
            )

    # Normalization
    def set_normalizer(self, normalizer: Optional[FieldNormalizer]) -> 'PDEDataset':
        """
        Attach (or clear) a field normalizer applied on every access.

        Pass the *training* split's normalizer to validation and test splits
        so all data is standardized with the same statistics.

        Returns:
            `self`, for chaining.
        """
        self.normalizer = normalizer
        return self

    def _field(self, name: str, idx: int) -> torch.Tensor:
        """Fetch one raw field, standardized when a normalizer is attached."""
        value = self.data[name][idx]
        if self.normalizer is not None:
            value = self.normalizer.normalize(name, value)
            # Re-apply the mask after standardizing: normalization subtracts
            # the field mean, which would turn the "not observed" zeros into a
            # nonzero constant and contradict the obs_mask channel.
            if name == self.observation_key:
                obs_mask = self.data.get('obs_mask')
                if obs_mask is not None:
                    value = value * obs_mask[idx]
        return value

    @property
    def input_fields(self) -> List[str]:
        """Raw field names composing `sample['input']`, in channel order.

        The mask is appended to the input as an extra channel, so it has to
        appear here too or the names no longer line up with the channels --
        which is exactly what `denormalize_channels` splits on.  It has no
        statistics, so it passes through denormalization unchanged.
        """
        fields = [self.input_key]
        if self.data.get('obs_mask') is not None:
            fields.append('obs_mask')
        return fields

    @property
    def target_fields(self) -> List[str]:
        """Raw field names composing `sample['target']`, in channel order."""
        return [self.target_key]

    #Dataset interface
    def __len__(self) -> int:
        return len(self.data[self.input_key])

    def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
        inp = self._field(self.input_key, idx)
        obs_mask = self.data.get('obs_mask')
        if obs_mask is not None:
            # The mask stays binary: it is an indicator channel, not a field.
            inp = torch.cat([inp, obs_mask[idx]], dim=0)

        sample = {
            'input': inp,
            'target': self._field(self.target_key, idx),
        }
        if obs_mask is not None:
            sample['obs_mask'] = obs_mask[idx]
        return sample

    # Getters for metadata
    def get_stats(self) -> Dict[str, Dict[str, float]]:
        """Get normalization statistics for all fields."""
        return self.stats

    def get_config(self) -> Dict[str, Any]:
        """Get generation configuration."""
        return self.config

    def get_raw_data(self) -> Dict[str, Any]:
        """Get the raw underlying data dictionary."""
        return self.data

input_fields property

Raw field names composing sample['input'], in channel order.

The mask is appended to the input as an extra channel, so it has to appear here too or the names no longer line up with the channels -- which is exactly what denormalize_channels splits on. It has no statistics, so it passes through denormalization unchanged.

target_fields property

Raw field names composing sample['target'], in channel order.

__init__(data, problem='forward', metadata=None, normalizer=None)

Parameters:

Name Type Description Default
data Dict[str, Tensor]

Dictionary with at least two tensor entries whose keys identify the PDE fields (e.g. 'source' / 'solution' for Poisson, 'initial' / 'final' for Burgers).

required
problem Literal['forward', 'inverse']

'forward' maps natural data -> solution; 'inverse' reverses the mapping.

'forward'
metadata Optional[Dict[str, Any]]

Optional dict with 'stats', 'config', etc.

None
normalizer Optional[FieldNormalizer]

Optional FieldNormalizer applied to each PDE field on access. Fit it on the training split and share the same instance with validation/test splits.

None
Source code in flowpde/datasets/exponax/base.py
def __init__(
    self,
    data: Dict[str, torch.Tensor],
    problem: Literal['forward', 'inverse'] = 'forward',
    metadata: Optional[Dict[str, Any]] = None,
    normalizer: Optional[FieldNormalizer] = None,
):
    """
    Args:
        data: Dictionary with at least two tensor entries whose
              keys identify the PDE fields (e.g. 'source' / 'solution'
              for Poisson, 'initial' / 'final' for Burgers).
        problem: 'forward' maps natural data -> solution;
                 'inverse' reverses the mapping.
        metadata: Optional dict with 'stats', 'config', etc.
        normalizer: Optional `FieldNormalizer` applied to each PDE
            field on access.  Fit it on the training split and share
            the same instance with validation/test splits.
    """
    if problem not in {'forward', 'inverse'}:
        raise ValueError("problem must be 'forward' or 'inverse'")

    self.data = data
    self.problem = problem
    self.metadata = metadata or {}
    self.stats = self.metadata.get('stats', {})
    self.config = self.metadata.get('config', {})
    # Which raw field carries the (possibly masked) observation, if any.
    self.observation_key = self.metadata.get('observation_key')
    self.normalizer = normalizer

    self._setup_keys()

set_normalizer(normalizer)

Attach (or clear) a field normalizer applied on every access.

Pass the training split's normalizer to validation and test splits so all data is standardized with the same statistics.

Returns:

Type Description
PDEDataset

self, for chaining.

Source code in flowpde/datasets/exponax/base.py
def set_normalizer(self, normalizer: Optional[FieldNormalizer]) -> 'PDEDataset':
    """
    Attach (or clear) a field normalizer applied on every access.

    Pass the *training* split's normalizer to validation and test splits
    so all data is standardized with the same statistics.

    Returns:
        `self`, for chaining.
    """
    self.normalizer = normalizer
    return self

get_stats()

Get normalization statistics for all fields.

Source code in flowpde/datasets/exponax/base.py
def get_stats(self) -> Dict[str, Dict[str, float]]:
    """Get normalization statistics for all fields."""
    return self.stats

get_config()

Get generation configuration.

Source code in flowpde/datasets/exponax/base.py
def get_config(self) -> Dict[str, Any]:
    """Get generation configuration."""
    return self.config

get_raw_data()

Get the raw underlying data dictionary.

Source code in flowpde/datasets/exponax/base.py
def get_raw_data(self) -> Dict[str, Any]:
    """Get the raw underlying data dictionary."""
    return self.data

GenerationConfig dataclass

Configuration for PDE data generation via Exponax.

Attributes:

Name Type Description
num_spatial_dims int

Number of spatial dimensions (1, 2, or 3)

num_points int

Number of grid points per spatial dimension

domain_extent float

Physical size of the periodic domain

num_samples int

Number of samples to generate

seed int

Random seed for reproducibility

torch_device str

Target PyTorch device for converted tensors

obs_noise_std float

Optional additive Gaussian noise applied to the inverse-problem observation field. Set to 0.0 to disable.

obs_mask_fraction float

Fraction of spatial grid points that are observed in inverse-problem datasets. Each sample independently draws a random Bernoulli mask with this probability; unobserved locations are zeroed out in the observation field. The binary mask (1 = observed, 0 = unobserved) is stored as 'obs_mask' in the dataset and returned by __getitem__. 1.0 (default) = full observations; 0.1 = only 10 % of points visible. Has no effect when problem='forward' or when left at 1.0.

Source code in flowpde/datasets/exponax/base.py
@dataclass
class GenerationConfig:
    """
    Configuration for PDE data generation via Exponax.

    Attributes:
        num_spatial_dims: Number of spatial dimensions (1, 2, or 3)
        num_points: Number of grid points per spatial dimension
        domain_extent: Physical size of the periodic domain
        num_samples: Number of samples to generate
        seed: Random seed for reproducibility
        torch_device: Target PyTorch device for converted tensors
        obs_noise_std: Optional additive Gaussian noise applied to the
            inverse-problem observation field.  Set to 0.0 to disable.
        obs_mask_fraction: Fraction of spatial grid points that are *observed*
            in inverse-problem datasets.  Each sample independently draws a
            random Bernoulli mask with this probability; unobserved locations
            are zeroed out in the observation field.  The binary mask
            (1 = observed, 0 = unobserved) is stored as `'obs_mask'` in the
            dataset and returned by `__getitem__`.  1.0 (default) = full
            observations; 0.1 = only 10 % of points visible.
            Has no effect when `problem='forward'` or when left at 1.0.
    """
    num_spatial_dims: int = 2
    num_points: int = 64
    domain_extent: float = 10.0
    num_samples: int = 10000
    seed: int = 42
    torch_device: str = 'cpu'
    obs_noise_std: float = 0.0
    obs_mask_fraction: float = 1.0

    def to_dict(self) -> Dict[str, Any]:
        return asdict(self)

Utilities

utilities

Exponax Dataset Utilities

Utility helpers for Exponax-backed datasets, including JAX↔PyTorch conversions and normalization stats.

jax_to_torch(jax_array, device='cpu', dtype=torch.float32)

Convert a JAX array to a PyTorch tensor.

Uses NumPy as an intermediate format for maximum compatibility across different JAX/PyTorch device configurations.

Parameters:

Name Type Description Default
jax_array Any

JAX array to convert

required
device DeviceType

Target PyTorch device ('cpu', 'cuda', etc.)

'cpu'
dtype Optional[dtype]

Target PyTorch dtype (default: torch.float32)

float32

Returns:

Type Description
Tensor

PyTorch tensor on the specified device

Source code in flowpde/datasets/exponax/utilities.py
def jax_to_torch(
    jax_array: "Any",
    device: DeviceType = 'cpu',
    dtype: Optional[torch.dtype] = torch.float32,
) -> torch.Tensor:
    """
    Convert a JAX array to a PyTorch tensor.

    Uses NumPy as an intermediate format for maximum compatibility
    across different JAX/PyTorch device configurations.

    Args:
        jax_array: JAX array to convert
        device: Target PyTorch device ('cpu', 'cuda', etc.)
        dtype: Target PyTorch dtype (default: torch.float32)

    Returns:
        PyTorch tensor on the specified device
    """
    np_array = np.array(jax_array, copy=True)
    tensor = torch.from_numpy(np_array)
    if dtype is not None:
        tensor = tensor.to(dtype=dtype)
    tensor = tensor.to(device=device)
    return tensor

compute_normalization_stats(tensor)

Compute normalization statistics for a tensor.

Parameters:

Name Type Description Default
tensor Tensor

Input tensor of any shape

required

Returns:

Type Description
dict

Dictionary with 'mean', 'std', 'min', and 'max' values.

Source code in flowpde/datasets/exponax/utilities.py
def compute_normalization_stats(tensor: torch.Tensor) -> dict:
    """
    Compute normalization statistics for a tensor.

    Args:
        tensor: Input tensor of any shape

    Returns:
        Dictionary with 'mean', 'std', 'min', and 'max' values.
    """
    return {
        'mean': tensor.mean().item(),
        'std': tensor.std().item(),
        'min': tensor.min().item(),
        'max': tensor.max().item(),
    }

sample_sine_fields(key, *, n, num_spatial_dims, num_points, domain_extent, num_terms, max_mode)

Generate simple smooth sine fields with random modes, weights, and phases.

Source code in flowpde/datasets/exponax/utilities.py
def sample_sine_fields(
    key,
    *,
    n: int,
    num_spatial_dims: int,
    num_points: int,
    domain_extent: float,
    num_terms: int,
    max_mode: int,
):
    """Generate simple smooth sine fields with random modes, weights, and phases."""
    if num_terms < 1:
        raise ValueError("num_terms must be at least 1")
    if max_mode < 1:
        raise ValueError("max_mode must be at least 1")

    import jax
    import jax.numpy as jnp

    coords = [
        jnp.linspace(0.0, domain_extent, num_points, endpoint=False)
        for _ in range(num_spatial_dims)
    ]
    grids = jnp.meshgrid(*coords, indexing='ij') if num_spatial_dims > 1 else coords
    stacked_grid = jnp.stack(grids, axis=0)

    _, mode_key, coeff_key, phase_key = jax.random.split(key, 4)
    modes = jax.random.randint(
        mode_key,
        shape=(n, num_terms, num_spatial_dims),
        minval=1,
        maxval=max_mode + 1,
    )
    coeffs = jax.random.normal(coeff_key, shape=(n, num_terms))
    phases = jax.random.uniform(
        phase_key,
        shape=(n, num_terms),
        minval=0.0,
        maxval=2.0 * jnp.pi,
    )

    def sample_one(sample_modes, sample_coeffs, sample_phases):
        field = jnp.zeros([num_points] * num_spatial_dims)
        for i in range(num_terms):
            angle = (
                2.0
                * jnp.pi
                * (
                    sample_modes[i].reshape(
                        (num_spatial_dims,) + (1,) * num_spatial_dims
                    )
                    * stacked_grid
                ).sum(axis=0)
                / domain_extent
                + sample_phases[i]
            )
            field = field + sample_coeffs[i] * jnp.sin(angle)
        field = field / jnp.sqrt(float(num_terms))
        return field[jnp.newaxis]

    return jax.vmap(sample_one)(modes, coeffs, phases)