Explicit programs and objectives
ADEPT is introducing a logging-free numerical API alongside ergoExo. It is opt-in:
existing solver entry points and output dictionaries are unchanged. The currently
registered solvers are tf-1d, electrostatic pic-1d, vfp-2d, vlasov-1d,
and farsight-1d.
Supported forward ergoExo calls also use prepared execution through the
compatibility façade.
The numerical boundary has five explicit values:
program(params, state, inputs, key) -> RawResult
paramscontains only the leaves selected for differentiation.stateis the evolving initial state.inputscontains fixed runtime forcing, scenarios, and targets.keyis the PRNG key for this run.RawResultcontains the final state, named device observations, their named coordinates, status, and solver statistics. See Observation planning and result materialization for bounded retention and explicit host transfer.
Keeping fixed arrays in state or inputs prevents them from becoming accidental
gradient targets. None of these calls start MLflow, write files, or mutate their
arguments.
Forward-only execution
Prepare from an existing configuration, then pass the numerical fields individually through the transform:
import equinox as eqx
import jax
from adept import SimulationSpec, solver_registry
prepared = solver_registry.prepare(
SimulationSpec.from_legacy_config(config),
key=42,
)
def run(program, params, state, inputs, key):
return program(params, state, inputs, key)
result = eqx.filter_jit(run)(
prepared.program,
prepared.params,
prepared.state,
prepared.inputs,
jax.random.key(42),
)
The manifest, analyzer, configuration model, paths, and tracking clients remain on the
host and must not be passed to jit, grad, or vmap.
Differentiated execution
An Objective returns a scalar loss plus stable metric and auxiliary PyTrees. A scalar
callable can be adapted with CallableObjective:
import jax.numpy as jnp
from adept import CallableObjective, partition_parameters, value_and_grad
runtime_values = eqx.combine(prepared.params, prepared.inputs)
selector = jax.tree.map(lambda _: False, runtime_values)
selector = eqx.tree_at(
lambda tree: tree["drivers"]["ex"]["0"]["a0"],
selector,
True,
)
partition = partition_parameters(runtime_values, selector)
params = partition.trainable
inputs = partition.frozen
objective = CallableObjective(
lambda result, params, inputs: jnp.mean(
result.observations["x"]["electron"]["u"][-1] ** 2
),
metric_name="electron_flow_energy",
)
run = eqx.filter_jit(value_and_grad)(
prepared.program,
objective,
params,
prepared.state,
inputs,
jax.random.key(42),
)
loss = run.objective.loss
metrics = run.objective.metrics
gradients = run.gradients
raw_result = run.simulation
value_and_grad differentiates only its params argument. Solver-specific vg()
methods are not needed on this path.
Selecting parameters and freezing other values
The pilot builders keep driver values in inputs by default. Select trainable leaves
with a boolean PyTree of the same structure. The complementary output keeps every
unselected value fixed:
from adept import partition_parameters
runtime_values = eqx.combine(prepared.params, prepared.inputs)
selector = jax.tree.map(lambda _: False, runtime_values)
selector = eqx.tree_at(
lambda tree: tree["drivers"]["ex"]["0"]["a0"],
selector,
True,
)
partition = partition_parameters(runtime_values, selector)
params = partition.trainable
inputs = partition.frozen
Here a0 is the only differentiable leaf. Other arrays—including w0, the driver
envelope, initial state, and any target data—stay explicit but frozen. Replacing a
selected or frozen value with another value of the same shape and dtype reuses the
compiled executable.
Objectives compose without putting logging in the JAX graph:
from adept import L2Penalty, WeightedSumObjective
objective = WeightedSumObjective(
(data_objective, L2Penalty()),
weights=(1.0, 1e-4),
names=("fit", "regularization"),
)
Vlasov1D
Use the same preparation call with a solver: vlasov-1d configuration. The builder
shares initialization and timestep operators with BaseVlasov1D, including
multispecies grids, supported advection and field solvers, Fokker–Planck/Krook
collisions, transverse waves, and stochastic longitudinal forcing. Enable JAX x64
before preparation. The specialized vlasov-1d-iaw module keeps its legacy path.
The driver controls are normalized EMDriverSet objects, so select an amplitude with:
runtime_values = eqx.combine(prepared.params, prepared.inputs)
selector = jax.tree.map(lambda _: False, runtime_values)
selector = eqx.tree_at(lambda tree: tree["drivers"].ex[0].a0, selector, True)
partition = partition_parameters(runtime_values, selector)
objective = CallableObjective(
lambda result, params, inputs: jnp.mean(result.final_state["e"] ** 2)
)
The initial state, numerical grids, and collision operators are fixed unless explicitly
replaced. Point-source location is a discrete cell selection; it does not provide a
useful position gradient. Density noise and precomputed stochastic-driver realizations
retain their configuration seeds (noise_seed and drivers.ex_stochastic.seed); the
preparation key is recorded in the manifest and is not substituted for those seeds.
Observations retain legacy names such as fields, electron.main, default, and
diag-fp-dfdt. Each stream has its own physical times. Between-step samples interpolate
the state before evaluating diagnostics, matching the legacy Diffrax stepper. The final
state always includes the complete last timestep, even if a distribution stream stops
earlier. As in the legacy solver, integration starts at zero and the grid rounds the
requested duration up to a complete step.
run_prepared returns in-memory fields, dists, and scalars xarray datasets in
completed.report.result, plus metrics for the final scalar observation. It performs
no plot or netCDF writes. Existing ergoExo calls retain their artifact pipeline.
The builder supports the existing grid.parallel single-host sharding; multi-host
execution and batched runs are not advertised.
Batched execution
For a builder that advertises prepared.capabilities.batchable, batch the explicit
runtime values rather than closing over scenarios:
if not prepared.capabilities.batchable:
raise ValueError("This solver has not declared batched execution support")
batched_run = eqx.filter_jit(
eqx.filter_vmap(run, in_axes=(None, 0, 0, 0, 0))
)
results = batched_run(prepared.program, params_batch, state_batch, inputs_batch, keys)
The two initial solver builders do not yet advertise general batching. The generic
program contract is tested under vmap; each solver must validate its own state and
observation layout before enabling the capability.
Migrating trainable_modules and vg() callers
LegacyVGAdapter temporarily preserves the numerical call shape expected by older
optimization loops:
from adept import LegacyVGAdapter
adapter = LegacyVGAdapter(
prepared.program,
objective,
prepared.state,
inputs,
jax.random.key(42),
)
(loss, output), gradients = eqx.filter_jit(adapter.vg)(params)
The adapter emits a migration warning and returns output["solver result"] alongside
the structured objective. It does not create an MLflow run or emulate mutation of
captured module attributes. New code should call value_and_grad directly. The first
compatibility-façade slice now routes supported ergoExo forward runs through these
contracts. See Legacy API compatibility for the current
routing and fallback rules.
For host-side execution, tracking, failure policies, and verified artifact storage, see Host-side tracking and artifacts.