Initializes the callable JAX function and registers its VJP.
(self)
| 217 | ) |
| 218 | |
| 219 | def _initialize_callable(self) -> Callable[[List[jnp.ndarray]], jnp.ndarray]: |
| 220 | """Initializes the callable JAX function and registers its VJP.""" |
| 221 | |
| 222 | @jax.custom_vjp |
| 223 | def simulate(design_variables: List[jnp.ndarray]) -> jnp.ndarray: |
| 224 | monitor_values, _ = self._run_fwd_simulation(design_variables) |
| 225 | return monitor_values |
| 226 | |
| 227 | def _simulate_fwd(design_variables): |
| 228 | """Runs forward simulation, returning monitor values and fields.""" |
| 229 | monitor_values, self.fwd_design_region_monitors = self._run_fwd_simulation( |
| 230 | design_variables |
| 231 | ) |
| 232 | design_variable_shapes = [x.shape for x in design_variables] |
| 233 | return monitor_values, (design_variable_shapes) |
| 234 | |
| 235 | def _simulate_rev(res, monitor_values_grad): |
| 236 | """Runs adjoint simulation, returning VJP of design wrt monitor values.""" |
| 237 | design_variable_shapes = res |
| 238 | self.adj_design_region_monitors = self._run_adjoint_simulation( |
| 239 | monitor_values_grad |
| 240 | ) |
| 241 | vjps = self._calculate_vjps( |
| 242 | self.fwd_design_region_monitors, |
| 243 | self.adj_design_region_monitors, |
| 244 | design_variable_shapes, |
| 245 | ) |
| 246 | return ([jnp.asarray(vjp) for vjp in vjps],) |
| 247 | |
| 248 | simulate.defvjp(_simulate_fwd, _simulate_rev) |
| 249 | |
| 250 | return simulate |