From acf8d82a9b9ab4b417cb49a7fd97b49582ed317c Mon Sep 17 00:00:00 2001 From: Matthew Keeler Date: Thu, 20 Aug 2026 02:33:22 +0200 Subject: [PATCH] fix(photonics2d): honor lambda1/lambda2 passed per call The solver reads omega1/omega2, which v0 only computes in __init__, so wavelengths passed in config were ignored. v1 now recomputes them in _setup_simulation. Fixes #274. --- engibench/problems/photonics2d/v1.py | 12 ++++++++++++ tests/test_photonics2d.py | 24 ++++++++++++++++++++++++ 2 files changed, 36 insertions(+) diff --git a/engibench/problems/photonics2d/v1.py b/engibench/problems/photonics2d/v1.py index cf0060a5..89efd781 100644 --- a/engibench/problems/photonics2d/v1.py +++ b/engibench/problems/photonics2d/v1.py @@ -57,6 +57,7 @@ from engibench.problems.photonics2d.backend import insert_mode from engibench.problems.photonics2d.backend import mode_overlap from engibench.problems.photonics2d.backend import poly_ramp +from engibench.problems.photonics2d.backend import wavelength_to_frequency from engibench.problems.photonics2d.v0 import Photonics2D as Photonics2D_v0 @@ -103,6 +104,17 @@ def __init__(self, *args: Any, **kwargs: Any) -> None: # ------------------------------------------------------------------ helpers + def _setup_simulation(self, config: dict[str, Any] | None = None) -> dict[str, Any]: + """Set up the simulation and update the frequencies from ``lambda1`` / ``lambda2``. + + v0 only computes ``omega1`` / ``omega2`` in ``__init__``, so wavelengths passed in ``config`` + were ignored. + """ + conditions = super()._setup_simulation(config) + self.omega1 = wavelength_to_frequency(conditions["lambda1"]) + self.omega2 = wavelength_to_frequency(conditions["lambda2"]) + return conditions + def _check_design(self, design: npt.NDArray, config: dict[str, Any] | None) -> None: """Validate ``design`` (and any overridden config) against the declared constraints. diff --git a/tests/test_photonics2d.py b/tests/test_photonics2d.py index 83572635..f1505ad4 100644 --- a/tests/test_photonics2d.py +++ b/tests/test_photonics2d.py @@ -12,6 +12,8 @@ These tests use a small grid and few optimization steps to stay fast. """ +import dataclasses + import numpy as np import pytest @@ -157,3 +159,25 @@ def test_design_to_epsr_is_linear_scaling(problem: Photonics2D) -> None: design_mask = problem._design_region.astype(bool) assert np.allclose(epsr_zeros[design_mask], problem._epsr_min) assert np.allclose(epsr_ones[design_mask], problem._epsr_max) + + +# --------------------------------------------------------------------- conditions + + +def test_simulate_honors_wavelengths_passed_per_call(problem: Photonics2D, start_design: np.ndarray) -> None: + """Wavelengths passed per call must give the same result as wavelengths passed to the constructor.""" + wavelengths = {"lambda1": 0.7, "lambda2": 1.4} + built = Photonics2D(num_elems_x=NUM_X, num_elems_y=NUM_Y, config=wavelengths) + per_call = float(problem.simulate(start_design, config=wavelengths)[0]) + at_construction = float(built.simulate(start_design)[0]) + assert per_call == pytest.approx(at_construction, rel=1e-9) + + +def test_dataset_row_reproduces_its_own_objective() -> None: + """Simulating a stored design under its own conditions returns the objective stored beside it.""" + problem = Photonics2D() + for row in problem.dataset["test"].select(range(2)): + conditions = {f.name: row[f.name] for f in dataclasses.fields(Photonics2D.Conditions)} + design = np.asarray(row["optimal_design"], dtype=float) + got = float(problem.simulate(design, config=conditions)[0]) + assert got == pytest.approx(row["total_overlap"], rel=1e-6)