"""Burgers 1D in finite volume: Rusanov, HLL and exact Godunov fluxes. The explicit Euler time integration uses :class:`TimeDiscreteFVscheme`, the same public Butcher-tableau / compiled ``lax.scan`` API as time-dependent DG. The Riemann problem ``u_L=1``, ``u_R=-0.25`` gives a shock travelling at speed ``(u_L + u_R)/2 = 0.375``. It is deliberately asymmetric: Rusanov, HLL and the exact Burgers Godunov flux are then visibly different. """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.linear_approximation.finite_volume import ( FiniteVolumeScheme, TimeDiscreteFVscheme, ) from scimba_jax.linear_approximation.finite_volume.hyperbolic import ( BurgersGodunovFlux, HLLFlux, RusanovFlux, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.linear_approximation.variables.variables_fv import VariablesFV from scimba_jax.nonlinear_approximation.model_class.funcparam_scalar import ( ParamScalarFunction, ) from scimba_jax.physical_models.abstract_conservative_pde import ( AbstractConservativePDE, ) from scimba_jax.physical_models.abstract_physical_conservative_model import ( AbstractPhysicalConservativeModel, ) from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet from scimba_jax.time_discrete.butcher_tableau import build_explicit_euler_tableau N_CELLS = 160 CFL = 0.4 T_FINAL = 0.2 LEFT_STATE = 1.0 RIGHT_STATE = -0.25 class BurgersPDE(AbstractConservativePDE): """``u_t + div(F(u)) = 0`` with ``F(u)=u²/2``.""" def __init__(self): super().__init__(dim=1) def construct_F(self, u): # noqa: N802 return ParamScalarFunction( u.dims, lambda scheme, x: 0.5 * u(scheme, x) ** 2, f_type=u.f_type, ) def construct_A(self, u): # noqa: N802 return None def construct_R(self, u): # noqa: N802 return None def initial_condition(x): return jnp.where(x[0] < 0.5, jnp.array([LEFT_STATE]), jnp.array([RIGHT_STATE])) def exact_solution(x, time): shock_position = 0.5 + 0.5 * (LEFT_STATE + RIGHT_STATE) * time return jnp.where(x < shock_position, LEFT_STATE, RIGHT_STATE) def hll_wave_speeds(_model, u_left, u_right, context): """Burgers characteristic bounds of the normal flux ``F.n``.""" speed_left = context.normal[0] * u_left[0] speed_right = context.normal[0] * u_right[0] return jnp.minimum(speed_left, speed_right), jnp.maximum(speed_left, speed_right) def make_time_scheme(flux, dt): mesh = cartesian_mesh([N_CELLS], quad_order=2) model = AbstractPhysicalConservativeModel(BurgersPDE()) model.add_boundary_condition("west", Dirichlet(lambda _x: jnp.array([LEFT_STATE]))) model.add_boundary_condition("east", Dirichlet(lambda _x: jnp.array([RIGHT_STATE]))) spatial_scheme = FiniteVolumeScheme( model, VariablesFV(mesh), flux, assemble_second_order=False, assemble_reaction=False, assemble_source=False, ) return TimeDiscreteFVscheme(spatial_scheme, build_explicit_euler_tableau(), dt=dt) def solve(flux, dt, nt): """Project the Riemann data and advance it with explicit Euler.""" scheme = make_time_scheme(flux, dt) initial = scheme.initialize(initial_condition) final, history = scheme.solve(initial, t0=0.0, nt=nt) return scheme, final, history def cell_centres(scheme): return jax.vmap(scheme.variables.mesh.cell_centroid)( jnp.arange(scheme.variables.mesh.n_cells_total) ) if __name__ == "__main__": h = 1.0 / N_CELLS dt = CFL * h / max(abs(LEFT_STATE), abs(RIGHT_STATE)) nt = round(T_FINAL / dt) dt = T_FINAL / nt fluxes = { "Rusanov": RusanovFlux(lambda _model, u, _context: jnp.abs(u[0])), "HLL": HLLFlux(hll_wave_speeds), "Godunov exact (Burgers)": BurgersGodunovFlux(), } figure, axis = plt.subplots(figsize=(11, 4.8)) for name, flux in fluxes.items(): scheme, final, _history = solve(flux, dt, nt) centers = cell_centres(scheme)[:, 0] axis.step(centers, final[:, 0], where="mid", lw=2, label=name) mass = jnp.sum(final[:, 0] * scheme.variables.mesh.cell_measures()) print( f"{name:25s} mass={float(mass):.8f} min/max=({float(final.min()):.3f}, {float(final.max()):.3f})" ) x_plot = jnp.linspace(0.0, 1.0, 1000) axis.plot(x_plot, exact_solution(x_plot, T_FINAL), "k--", lw=2, label="choc exact") axis.set( xlabel="x", ylabel="u", title=( "Burgers 1D — FV P0 + Euler explicite " f"(N={N_CELLS}, CFL={CFL}, t={T_FINAL})" ), ) axis.grid(alpha=0.25) axis.legend(ncol=2) figure.tight_layout() plt.show()