Mixed Boundary Conditions with PINNs

In this tutorial, we solve a 2D Poisson equation with mixed boundary conditions using Physics-Informed Neural Networks (PINNs).

We enforce Dirichlet conditions on the bottom boundary and Neumann conditions on the top boundary of a disk domain.

Problem Statement

We solve:

\[\begin{split}\begin{align*} -\mu \Delta u &= f \quad\text{in } \Omega \times M \\ u &= g_D \quad\text{on } \Gamma_D \\ \frac{\partial u}{\partial n} &= g_N \quad\text{on } \Gamma_N \end{align*}\end{split}\]

where:

  • \(\Omega\) is a disk centered at (0, 0) with radius 1

  • \(\Gamma_D\) is the bottom semicircle (Dirichlet boundary)

  • \(\Gamma_N\) is the top semicircle (Neumann boundary)

  • \(\mu \in M = [1, 2]\) is a parameter

  • The exact solution is: \(u(x_1, x_2, \mu) = \log(1 + x_1^2 + x_2^2 + x_1 x_2)\)

Preliminary imports and definition

[1]:
import jax
import jax.numpy as jnp

def exact_solution(xy, mu):
    """Exact solution for comparison."""
    x, y = xy[..., 0:1], xy[..., 1:2]
    return jnp.log(1 + x**2 + y**2 + x * y)

def f_rhs(xy, mu):
    """Right-hand side of the PDE."""
    x, y = xy[..., 0:1], xy[..., 1:2]
    numer = 4 - x**2 - y**2 - 4 * x * y
    denom = (1 + x**2 + y**2 + x * y) ** 2
    return -mu[..., 0:1] * numer / denom

Define the Domain with Mixed Boundaries

To enforce mixed boundary conditions, we need to:

  1. Create separate boundary domains for each condition type

  2. Define a boundary condition RHS function for each boundary

  3. Map boundaries to their conditions in a dictionary

Here, we create a disk with:

  • Bottom semicircle (bc south) for Dirichlet condition

  • Top semicircle (bc north) for Neumann condition

[2]:
from scimba_jax.domains.meshless_domains.domains_2d import ArcCircle2D, Disk2D

domain_mu = [(1.0, 2.0)]  # parameter domain

# Main domain: unit disk
domain_x = Disk2D((0.0, 0.0), 1, is_main_domain=True)

# Boundary domains: top and bottom semicircles
disk_up = ArcCircle2D((0.0, 0.0), 1, (0.0, jnp.pi), label_str='bc north')
disk_do = ArcCircle2D((0.0, 0.0), 1, (jnp.pi, 2 * jnp.pi), label_str='bc south')

domain_x.add_bc_domain(disk_up)
domain_x.add_bc_domain(disk_do)

# Map boundary labels to condition types
domain_x.set_boundaries_dict({'S': ['bc south'], 'N': ['bc north']})

Next we define the righthand-sides for the boundary conditions:

[3]:
def f_bc_dirichlet(x, n, mu):
    """Dirichlet boundary condition (bottom semicircle)."""
    x, y = x[..., 0:1], x[..., 1:2]
    return jnp.log(2 + x * y)

def f_bc_neumann(x, n, mu):
    """Neumann boundary condition (top semicircle)."""
    x, y = x[..., 0:1], x[..., 1:2]
    return 2 * (1 + x * y) / (2 + x * y)

# Dictionary mapping boundary labels to their RHS functions
f_bc_rhs = {'S': f_bc_dirichlet, 'N': f_bc_neumann}

Define the Model

For mixed boundary conditions, the model’s physical_residuals dictionary must contain:

  • An interior residual for the PDE in the domain

  • A Dirichlet residual for the boundary with Dirichlet condition

  • A Neumann residual for the boundary with Neumann condition

Each residual is associated with its corresponding domain using the boundary labels defined above.

[4]:
from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel
from scimba_jax.physical_models.boundary_residuals import DirichletResidual, NeumannResidual
from scimba_jax.physical_models.elliptic_pde.laplacians import ParametricLaplacianResidual

class MixedBCLaplacian(AbstractPhysicalModel):
    def __init__(self, main_domain, f_rhs, f_bc_rhs):
        super().__init__(main_domain=main_domain)
        self.physical_residuals = {
            'interior': ParametricLaplacianResidual(main_domain, f_rhs),
            'S': DirichletResidual(domain=self.boundaries['S'], f_rhs=f_bc_rhs['S']),
            'N': NeumannResidual(domain=self.boundaries['N'], f_rhs=f_bc_rhs['N']),
        }

And finally instantiate the problem:

[5]:
model = MixedBCLaplacian(domain_x, f_rhs, f_bc_rhs)

Define and train a PINN

[6]:
from scimba_jax.nonlinear_approximation.integration.monte_carlo import DomainSampler, TensorizedSampler
from scimba_jax.nonlinear_approximation.integration.monte_carlo_parameters import UniformParametricSampler

sampler = TensorizedSampler(
    [
        DomainSampler(domain_x),
        UniformParametricSampler(domain_mu),
    ],
    bc=True,  # sample boundary points
)
[7]:
from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ApproximationSpace
from scimba_jax.nonlinear_approximation.networks.mlp import MLP
from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector

key = jax.random.PRNGKey(0)
nn = MLP(in_size=3, out_size=1, hidden_sizes=[20] * 3, key=key)
space = ApproximationSpace({'x': 2, 'mu': 1}, [(nn, 'scalar', 1)], model_type='x_mu')

# Loss weights (higher for boundary conditions to ensure they are satisfied)
weights = {'interior': [1.0], 'S': [30.0], 'N': [30.0]}

pinn = Projector(model, space, sampler, matrix_reguarization=5e-4, weights=weights)
[8]:
import timeit

N_COLLOC = 2000  # interior collocation points
N_BC_COLLOC = 3000  # boundary collocation points
N_EPOCHS = 50  # training epochs

start = timeit.default_timer()
key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC, N_BC_COLLOC)
end = timeit.default_timer()

print('Best loss:', pinn.best_loss)
print(f'Time for {N_EPOCHS} epochs: {end - start:.4f} seconds')
Training: 100%|||||||||||||||||||| 50/50[00:05<00:00] , loss: 5.4e+01 -> 1.8e-08
Best loss: {'N': Array([7.841326e-11], dtype=float64), 'S': Array([6.66803892e-11], dtype=float64), 'interior': Array([1.37123828e-08], dtype=float64), 'total': Array(1.80651923e-08, dtype=float64)}
Time for 50 epochs: 5.6441 seconds
[9]:
import matplotlib.pyplot as plt

from scimba_jax.plots.plots_nd import plot_abstract_approx_space

plot_abstract_approx_space(
    pinn.space,
    domain_x,
    domain_mu,
    loss=pinn.losses,
    residual=pinn.model,
    draw_contours=True,
    n_drawn_contours=20,
    solution=exact_solution,
)
plt.show()
../_images/tutorials_jax_mixed_boundary_conditions_16_0.png