from __future__ import annotations
from collections.abc import Callable
import jax
from fdtdx.config import SimulationConfig
from fdtdx.core.jax.default_key import default_key
from fdtdx.fdtd.container import ArrayContainer, ObjectContainer, SimulationState
from fdtdx.fdtd.execution import NativeExecutionUnavailable, select_engine
from fdtdx.fdtd.fdtd import adjoint_fdtd, checkpointed_fdtd, reversible_fdtd
from fdtdx.fdtd.native import native_scene_support, run_native_fdtd
from fdtdx.fdtd.stop_conditions import StoppingCondition
[docs]
def run_fdtd(
arrays: ArrayContainer,
objects: ObjectContainer,
config: SimulationConfig,
key: jax.Array | None = None,
stopping_condition: StoppingCondition | None = None,
show_progress: bool = True,
progress_callback: Callable[[int, int], None] | None = None,
) -> SimulationState:
key = default_key(key)
if select_engine(config) == "cuda":
support = native_scene_support(
arrays, objects, config, stopping_condition=stopping_condition
)
if support.supported:
return run_native_fdtd(
arrays=arrays,
objects=objects,
config=config,
key=key,
stopping_condition=stopping_condition,
show_progress=show_progress,
progress_callback=progress_callback,
)
if config.execution_mode == "cuda":
raise NativeExecutionUnavailable(
"Native CUDA scene is unsupported: " + "; ".join(support.reasons)
)
if stopping_condition is not None:
if config.gradient_config is not None:
raise NotImplementedError(
"Custom stopping conditions are not yet compatible with gradient computation. "
"Set config.gradient_config to None or use default time-based stopping by "
"setting stopping_condition=None."
)
if config.gradient_config is None:
# only forward simulation, use standard while loop of checkpointed fdtd
return checkpointed_fdtd(
arrays=arrays,
objects=objects,
config=config,
key=key,
stopping_condition=stopping_condition,
show_progress=show_progress,
progress_callback=progress_callback,
)
if config.gradient_config.method == "reversible":
return reversible_fdtd(
arrays=arrays,
objects=objects,
config=config,
key=key,
show_progress=show_progress,
progress_callback=progress_callback,
)
elif config.gradient_config.method == "adjoint":
return adjoint_fdtd(
arrays=arrays,
objects=objects,
config=config,
key=key,
show_progress=show_progress,
progress_callback=progress_callback,
)
elif config.gradient_config.method == "checkpointed":
return checkpointed_fdtd(
arrays=arrays,
objects=objects,
config=config,
key=key,
stopping_condition=stopping_condition,
show_progress=show_progress,
progress_callback=progress_callback,
)
elif config.gradient_config.method == "auto":
raise RuntimeError(
"automatic gradient backend is unresolved; call "
"fdtdx.configure_gradient_backend(arrays, objects, config) once after placement"
)
else:
raise Exception(f"Unknown gradient computation method: {config.gradient_config.method}")