jaxdem.system#
The simulation configuration and the tools that drive the simulation.
Classes
|
The full simulation configuration. |
- final class jaxdem.system.System(linear_integrator: LinearIntegrator, rotation_integrator: RotationIntegrator, collider: Collider, domain: Domain, force_manager: ForceManager, bonded_force_model: BondedForceModel | None, force_model: ForceModel, mat_table: MaterialTable, dt: jax.Array, time: jax.Array, dim: jax.Array, step_count: jax.Array, key: jax.Array, interact_same_bond_id: jax.Array, user_pre_step_actions: Callable[[State, System], tuple[State, System]] = <PjitFunction of <function _save_state_system>>, user_post_step_actions: Callable[[State, System], tuple[State, System]] = <PjitFunction of <function _save_state_system>>, minimizer: Any = None, target_fn: Callable[[State, System], jax.Array] | None = None)#
Bases:
objectThe full simulation configuration.
Notes:#
The System object supports JIT compilation for efficient execution.
The System dataclass is compatible with
jax.jit(), so every field should remain JAX arrays for best performance.
Example:#
Creating a basic 2D simulation system:
>>> import jaxdem as jdem >>> import jax.numpy as jnp >>> >>> # Create a System instance >>> sim_system = jdem.System.create( >>> state_shape=state.shape, >>> dt=0.001, >>> linear_integrator_type="euler", >>> rotation_integrator_type="spiral", >>> collider_type="naive", >>> domain_type="free", >>> force_model_type="spring", >>> # You can pass keyword arguments to component constructors via '_kw' dicts >>> domain_kw=dict(box_size=jnp.array([5.0, 5.0]), anchor=jnp.array([0.0, 0.0])) >>> ) >>> >>> print(f"System integrator: {sim_system.linear_integrator.__class__.__name__}") >>> print(f"System force model: {sim_system.force_model.__class__.__name__}") >>> print(f"Domain box size: {sim_system.domain.box_size}")
- linear_integrator: LinearIntegrator#
Instance of
jaxdem.LinearIntegratorthat advances the simulation linear state in time.
- rotation_integrator: RotationIntegrator#
Instance of
jaxdem.RotationIntegratorthat advances the simulation angular state in time.
- collider: Collider#
Instance of
jaxdem.Colliderthat performs contact detection and computes inter-particle forces and potential energies.
- domain: Domain#
Instance of
jaxdem.Domainthat defines the simulation boundaries, displacement rules, and boundary conditions.
- force_manager: ForceManager#
Instance of
jaxdem.ForceManagerthat handles per particle forces like external forces and resets forces.
- bonded_force_model: BondedForceModel | None#
Optional instance of
jaxdem.BondedForceModelthat defines bonded interactions by passing a force and energy function to the ForceManager.
- force_model: ForceModel#
Instance of
jaxdem.ForceModelthat defines the physical laws for inter-particle interactions.
- mat_table: MaterialTable#
Instance of
jaxdem.MaterialTableholding material properties and pairwise interaction parameters.
- dt: jax.Array#
The global simulation time step \(\Delta t\).
- time: jax.Array#
Elapsed simulation time.
- dim: jax.Array#
Spatial dimension of the system.
- step_count: jax.Array#
Number of integration steps that have been performed.
- key: jax.Array#
PRNG key for stochastic operations. Always update it with split so each use gets new random numbers.
- interact_same_bond_id: jax.Array#
Boolean scalar controlling interactions between particles with the same
bond_id.If
False(default), colliders mask out these pairs. IfTrue, these pairs interact.
- user_pre_step_actions(system: System) tuple[State, System][source]#
Function called before every step to perform user-defined actions.
- user_post_step_actions(system: System) tuple[State, System][source]#
Function called after every step to perform user-defined actions.
- minimizer: Any = None#
An optax GradientTransformation wrapped in CustomGradientTransformation, used for target_fn minimization.
- target_fn: Callable[[State, System], jax.Array] | None = None#
Optional custom target evaluation function for minimization.
- static create(state_shape: tuple[int, ...] | None = None, *, state: State | None = None, dt: float = 0.005, time: float = 0.0, linear_integrator_type: str | None = 'verlet', rotation_integrator_type: str | None = 'verletspiral', collider_type: str = 'naive', domain_type: str = 'free', bonded_force_model_type: str | None = None, bonded_force_model_kw: dict[str, Any] | None = None, bonded_force_manager_kw: dict[str, Any] | None = None, bonded_force_model: BondedForceModel | None = None, force_model_type: str = 'spring', force_manager_kw: dict[str, Any] | None = None, mat_table: MaterialTable | None = None, linear_integrator: LinearIntegrator | None = None, rotation_integrator: RotationIntegrator | None = None, collider: Collider | None = None, domain: Domain | None = None, force_model: ForceModel | None = None, force_manager: ForceManager | None = None, linear_integrator_kw: dict[str, Any] | None = None, rotation_integrator_kw: dict[str, Any] | None = None, collider_kw: dict[str, Any] | None = None, domain_kw: dict[str, Any] | None = None, force_model_kw: dict[str, Any] | None = None, seed: int = 0, key: jax.Array | None = None, interact_same_bond_id: bool = False, user_pre_step_actions: Callable[[State, System], tuple[State, System]] | None = None, user_post_step_actions: Callable[[State, System], tuple[State, System]] | None = None, minimizer: Any = None, minimizer_kw: dict[str, Any] | None = None, target_fn: Callable[[State, System], jax.Array] | None = None) System[source]#
Factory method to create a
Systeminstance with specified components.Every component slot accepts either a pre-built instance (
linear_integrator,rotation_integrator,collider,domain,force_model,force_manager,bonded_force_model,mat_table) or a registered type string plus keyword dict (<component>_type/<component>_kw). When you provide an instance, the method uses it as-is and ignores the corresponding*_type/*_kwarguments.- Parameters:
state_shape (Tuple, optional) – Shape of the state tensors handled by the simulation. The penultimate dimension corresponds to the number of particles
Nand the last dimension corresponds to the spatial dimensiondim. You can omit it when you providestate.state (State, optional) – The initial simulation state. When provided, the method infers
state_shapefrom it and forwards the state to colliders whoseCreatemethod requires one (e.g."CellList","NeighborList"), socollider_kw={"state": state}is not needed.dt (float, optional) – The global simulation time step.
linear_integrator_type (str or None, optional) – The registered type string for the
jaxdem.integrators.LinearIntegratorused to evolve translational degrees of freedom.None(or the empty string) disables linear integration (no-op integrator).rotation_integrator_type (str or None, optional) – The registered type string for the
jaxdem.integrators.RotationIntegratorused to evolve angular degrees of freedom.None(or the empty string) disables rotational integration (no-op integrator).collider_type (str, optional) – The registered type string for the
jaxdem.Colliderto use.domain_type (str, optional) – The registered type string for the
jaxdem.Domainto use.bonded_force_model_type (str or None, optional) – The registered type string for the
jaxdem.BondedForceModelto use.bonded_force_model_kw (Dict[str, Any] or None, optional) – Keyword arguments forwarded to
BondedForceModel.create.bonded_force_manager_kw (Dict[str, Any] or None, optional) – Deprecated alias of
bonded_force_model_kw(the dict has always been forwarded to the bonded force model, not the manager).force_model_type (str, optional) – The registered type string for the
jaxdem.ForceModelto use.force_manager_kw (Dict[str, Any] or None, optional) – Keyword arguments to pass to the constructor of ForceManager.
mat_table (MaterialTable or None, optional) – An optional pre-configured
jaxdem.MaterialTable. If None, the method creates a default jaxdem.MaterialTable with one generic elastic material and the “harmonic” jaxdem.MaterialMatchmaker.linear_integrator (LinearIntegrator, optional) – Pre-built linear integrator instance. Overrides
linear_integrator_type/linear_integrator_kw.rotation_integrator (RotationIntegrator, optional) – Pre-built rotation integrator instance. Overrides
rotation_integrator_type/rotation_integrator_kw.collider (Collider, optional) – Pre-built collider instance. Overrides
collider_type/collider_kw.domain (Domain, optional) – Pre-built domain instance. Overrides
domain_type/domain_kw.force_model (ForceModel, optional) – Pre-built force model instance. Overrides
force_model_type/force_model_kw.force_manager (ForceManager, optional) – Pre-built force manager instance. Overrides
force_manager_kw. Cannot be combined with a bonded force model (the bonded force functions must already be part of the provided manager).linear_integrator_kw (Dict[str, Any] or None, optional) – Keyword arguments forwarded to the constructor of the selected LinearIntegrator type.
rotation_integrator_kw (Dict[str, Any] or None, optional) – Keyword arguments forwarded to the constructor of the selected RotationIntegrator type.
collider_kw (Dict[str, Any] or None, optional) – Keyword arguments to pass to the constructor of the selected Collider type.
domain_kw (Dict[str, Any] or None, optional) – Keyword arguments to pass to the constructor of the selected Domain type.
force_model_kw (Dict[str, Any] or None, optional) – Keyword arguments to pass to the constructor of the selected ForceModel type.
seed (int, optional) – Integer seed used for random number generation. Defaults to 0. Used only when
keyis not provided.key (jax.Array, optional) – Key for JAX random number generation. When you provide
key, the method ignoresseed.interact_same_bond_id (bool, optional) – Whether particles with the same bond_id interact. Defaults to False.
user_pre_step_actions (Callable, optional) – A function called before every time step to perform user-defined actions.
user_post_step_actions (Callable, optional) – A function called after every time step to perform user-defined actions.
minimizer (Callable, optional) – Optimizer factory used by
System.minimize(). Called asminimizer(**minimizer_kw)and must return an optax-styleGradientTransformation. Defaults to FIRE (jaxdem.minimizers.fire()).minimizer_kw (Dict[str, Any] or None, optional) – Keyword arguments passed to
minimizer. If the minimizer’s signature accepts adtparameter and none is given here, the method passes the systemdtautomatically.target_fn (Callable, optional) – Custom objective
(state, system) -> scalarminimized bySystem.minimize(). WhenNone,System.minimize()uses the total potential energy.
- Returns:
A fully configured System instance ready for simulation.
- Return type:
- Raises:
KeyError – If a specified *_type is not registered in its respective factory, or if the mat_table is missing properties required by the force_model.
TypeError – If constructor keyword arguments are invalid for any component.
ValueError – If the domain_kw ‘box_size’ or ‘anchor’ shapes do not match the dim.
Example
Creating a 3D system with reflective boundaries and a custom dt:
>>> import jaxdem as jdem >>> import jax.numpy as jnp >>> >>> system_reflect = jdem.System.create( >>> state_shape=(N, 3), >>> dt=0.0005, >>> domain_type="reflect", >>> domain_kw=dict(box_size=jnp.array([20.0, 20.0, 20.0]), anchor=jnp.array([-10.0, -10.0, -10.0])), >>> force_model_type="spring", >>> ) >>> print(f"System dt: {system_reflect.dt}") >>> print(f"Domain type: {system_reflect.domain.__class__.__name__}")
Creating a system with a pre-defined MaterialTable:
>>> custom_mat_kw = dict(young=2.0e5, poisson=0.25) >>> custom_material = jdem.Material.create("custom_mat", **custom_mat_kw) >>> custom_mat_table = jdem.MaterialTable.from_materials( ... [custom_material], matcher=jdem.MaterialMatchmaker.create("linear") ... ) >>> >>> system_custom_mat = jdem.System.create( ... state_shape=(N, 2), ... mat_table=custom_mat_table, ... force_model_type="spring" ... )
- static trajectory_rollout(state: State, system: System, *, n: int | None = None, stride: int = 1, strides: jax.Array | None = None, save_fn: Callable[[State, System], Any] = <PjitFunction of <function _save_state_system>>, unroll: int = 2) tuple[State, System, Any][source]#
Roll the system forward while collecting saved outputs at each frame.
The rollout always stores one output per frame via save_fn(state, system). The output of save_fn must be a pytree. Frame spacing can be either: - constant (stride), or - variable (strides jax.Array).
The rollout saves each frame after its integration steps, so it does not store the initial (step-0) state. To record it, save it before the rollout, or pass a leading
0entry in strides.- Parameters:
state (State) – Initial state.
system (System) – Initial system configuration.
n (int, optional) – Number of saved frames. Required when strides is None. The method ignores n when you provide strides.
stride (int, optional) – Constant number of integration steps between consecutive saves. Used only when strides is None. Defaults to 1.
strides (jax.Array, optional) – Integer 1D array of per-frame integration strides. When provided, this overrides stride, and the method infers n from len(strides).
save_fn (Callable[[State, System], Any], optional) – Function called after each saved frame. The rollout stacks its return pytree along axis 0 across frames. Defaults to returning (state, system).
unroll (int, optional) – Unroll factor passed to the outer jax.lax.scan. Defaults to 2.
- Returns:
(final_state, final_system, trajectory_like) where trajectory_like is the stacked output of save_fn.
- Return type:
- Raises:
ValueError – If n is missing while strides is None, or if strides is not 1D.
Example
>>> import jaxdem as jdem >>> import jax.numpy as jnp >>> >>> state = jdem.utils.grid_state(n_per_axis=(1, 1), spacing=1.0, radius=0.1) >>> system = jdem.System.create(state_shape=state.shape, dt=0.01) >>> >>> # Constant stride: n is required >>> final_state, final_system, traj = jdem.System.trajectory_rollout( ... state, system, n=10, stride=5 ... ) >>> >>> # Variable strides: n inferred from len(strides) >>> deltas = jnp.array([1, 2, 4, 8]) >>> final_state, final_system, traj = jdem.System.trajectory_rollout( ... state, system, strides=deltas ... )
- static step(state: State, system: System, *, n: int | jax.Array = 1) tuple[State, System][source]#
Advance the simulation by n integration steps.
- Parameters:
- Returns:
(final_state, final_system) after n steps.
- Return type:
Example
>>> # Advance by 10 steps >>> state_after_10_steps, system_after_10_steps = jdem.System.step(state, system, n=10)
Notes
This method does not check collider overflow, to avoid a host synchronization per step.
- static stack(systems: Sequence[System]) System[source]#
Concatenate a sequence of
Systemsnapshots into a trajectory or batch along axis 0.Use this method to collect simulation snapshots over time into a single System object where the leading dimension represents time, or to prepare a batched system.
- Parameters:
systems (Sequence[System]) – A sequence (e.g., list, tuple) of
Systeminstances to be stacked.- Returns:
A new
Systeminstance where each attribute is a JAX array with an additional leading dimension representing the stacked trajectory. For example, if input pos was (N, dim), output pos will be (T, N, dim).- Return type:
- static unstack(system: System) list[System][source]#
Split a stacked/batched
Systemalong the leading axis into a Python list.This method is the inverse of
System.stack():If stacked = System.stack([sys0, sys1, …]), then System.unstack(stacked) returns [sys0, sys1, …].
Notes
The method splits along axis 0 (the leading axis).
This method cannot split a single snapshot System.
- static minimize(state: State, system: System, *, max_steps: int = 10000, pe_tol: float = 1e-16, pe_diff_tol: float = 1e-16) tuple[State, System, int, float][source]#
Minimize the energy of the system using the configured minimizer.
- Parameters:
state (State) – The state of the simulation.
system (System) – The system configuration.
max_steps (int, optional) – The maximum number of steps to take. Defaults to 10000.
pe_tol (float, optional) – The tolerance for the potential energy. Defaults to 1e-16.
pe_diff_tol (float, optional) – The tolerance for the difference in potential energy. Defaults to 1e-16.
- Returns:
The final state, system, number of steps, and potential energy (per particle when no custom
target_fnis set).- Return type:
Notes
The loop stops as soon as any convergence criterion is met (energy tolerance, relative energy change, or force tolerance) — see
jaxdem.minimizers.minimize()for the full list.