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:
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:
Create separate boundary domains for each condition type
Define a boundary condition RHS function for each boundary
Map boundaries to their conditions in a dictionary
Here, we create a disk with:
Bottom semicircle (
bc south) for Dirichlet conditionTop 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()