MCPcopy Create free account
hub / github.com/NanoComp/meep / _initialize_callable

Method _initialize_callable

python/adjoint/wrapper.py:219–250  ·  view source on GitHub ↗

Initializes the callable JAX function and registers its VJP.

(self)

Source from the content-addressed store, hash-verified

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

Callers 1

__init__Method · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected