Unifying md-system and cg_jaxmd: calculator & builder schema¶
Architecture history, not a current support matrix
This document preserves decisions, experiments, and phased implementation
notes. Statements such as “currently,” “not yet,” and cluster observations
are historical unless they carry an [evidence: ...] marker. Use the
maintained capability and validation pages for current behavior.
Status: In progress — the shared stack is built and tested; both
front-end entrypoint swaps are landed (opt-in for md-system); apocharmm and
RDF validation remain (see §0 and §11).
Scope: How to split the calculators and system builders shared by
mmml/cli/run/md_system.py and
examples/cg_jaxmd.py so the two can run on one
architecture.
0. Implementation status¶
The design below is largely realized in the mmml/md/ package. What exists,
end-to-end from a config:
config ──lowering──▶ RunConfig ──assemble──▶ builder → HybridEnergy → JaxmdDriver
| Layer | Module | State |
|---|---|---|
| Topology / FF state | mmml/md/system.py (MolecularSystem, FFParams) |
✅ tested |
| Config | mmml/md/config.py (RunConfig, EnsembleSpec) |
✅ |
| Energy terms (all 6) | mmml/md/energy/terms/ |
✅ parity-tested |
| Capacity / dtype policy | mmml/md/energy/capacity.py |
✅ tested |
| Builders + FF bridge | mmml/md/builders/ |
✅ CHARMM-integration tested |
| Driver incl. NPT | mmml/md/drivers/jaxmd.py |
✅ tested |
| Assembly glue | mmml/md/assemble.py |
✅ tested |
| Lowering adapters | mmml/md/lowering.py |
✅ tested |
| Neighbor-list factory | mmml/md/neighbors.py |
✅ tested |
Rigid-body Sampler (MC) |
mmml/md/samplers/rigid.py |
✅ tested |
Energy terms: ml_intra, ml_pep_water, mm_nonbonded, vdw_core, smd,
dihedral — each registered, box-aware where relevant, and validated against
the cg_jaxmd originals / the reference nonbonded / the example ML checkpoint.
End-to-end: assemble_and_run runs the full pipeline as real dynamics —
mm_nonbonded declares an intermolecular NeighborRequest, assemble_and_run
auto-builds the padded neighbor list (make_intermolecular_neighbor_fn), and
the JaxmdDriver propagates it (NVE smoke test on a multi-molecule box).
Cross-cutting: the cross-platform libcharmm loader (pycharmm/lib.py)
now self-discovers setup/charmm on both .dylib/.so; import mmml.md
stays free of jax/CHARMM (heavy deps are lazy inside make() / run()); the
whole tests/unit/test_md_*.py glob (incl. pre-existing legacy md-system
tests) is green (333 passed, 1 skipped).
Bug fix found via end-to-end validation: JaxmdDriver's fixed-box
NVE/NVT/FIRE path called space.periodic_general(box) without
fractional_coordinates=False. jax-md defaults that to True, so real-space
(Å) positions were silently treated as fractional coordinates and wrapped
modulo 1 by the integrator's shift_fn on the very first step — a discontinuous,
unphysical jump invisible to energy-only checks at R but catastrophic once any
step was taken. Found by running the full assemble_and_run pipeline
(PeptideWaterSystemBuilder + the example ML checkpoint + ml_intra +
mm_nonbonded) end-to-end: FIRE minimization exploded to ~1e15 eV within a few
steps. Fixed by passing fractional_coordinates=False for the fixed-box branch
(the NPT branch was already correct, since it explicitly bridges fractional↔real
itself). Regression-tested in
tests/unit/test_md_jaxmd_driver.py::test_fixed_box_pbc_keeps_real_space_positions_unwrapped
(verified to fail without the fix, pass with it). With the fix, a real
CHARMM-built peptide-water system + ML checkpoint runs a stable 50-step FIRE
minimization end-to-end (energy monotonically decreasing, no divergence).
Second bug fix found the same way: assemble_and_run's auto-wired
neighbor list did not pass peptide_water_ml=True to
make_intermolecular_neighbor_fn when the ml_pep_water term was selected, so
mm_nonbonded still included core–solvent pairs even though ml_pep_water
already scores them — a silent double-count of the peptide-water interaction
(classical LJ/Coulomb and the ML dimer term). Found while validating the
peptide_water_ml toggle path in examples/cg_jaxmd_unified.py: energy jumped
to ~2e4 eV instead of the expected MD/MM scale. Fixed in assemble.py by
checking "ml_pep_water" in config.terms. Regression-tested by comparing the
excluded vs. unexcluded pair list and energy on the same geometry (see
tests/unit/test_cg_jaxmd_unified.py::test_end_to_end_peptide_water_ml_no_double_counting) —
note that this test asserts the relative with/without-exclusion difference
rather than an absolute energy threshold, since CHARMM's persistent global
state across builds in one process makes absolute geometry (and thus energy)
depend on prior test execution order within the same session.
Third fix, same technique: validating --jaxmd-unified end-to-end found
that build_packmol_composition_cluster (the packmol composition builder used
by md-system) neither writes a PSF file nor bootstraps CHARMM itself — both
assumptions the legacy md_pbc_suite/jaxmd.py path could rely on implicitly
because its caller already did so, but which a fresh, explicit front-end
cannot. mmml/cli/run/md_system_unified.py makes both explicit: it writes the
live CHARMM PSF to a scratch file to resolve FFParams, and calls
ensure_pycharmm_loaded() before building.
Fourth fix, found by a cross-backend sweep rather than one entrypoint:
building workflows/unified_backend_sweep/ (a Snakemake workflow exercising
every backend/ensemble reachable through assemble_and_run — FIRE/NVE/NVT/NPT
via JaxmdDriver, plus RigidBodySampler) found that RigidBodySampler had no
way to supply a neighbor list, so any term with an intermolecular pair
requirement (mm_nonbonded) raised TracerArrayConversionError under the
sampler's jax.jit. Fixed by giving RigidBodySampler the same neighbor_fn
mechanism as JaxmdDriver (rebuilt once per N sweeps, not per proposed move),
and extending assemble_and_run's auto-wiring to cover the rigid path too. The
sweep also confirmed NPT is markedly slower to JIT-compile with a real ML model
than the fixed-box ensembles.
Fifth+sixth fix, found submitting that sweep to a real cluster (pc-studix,
SLURM): (1) the vendored Packmol binary isn't in git (platform-specific;
needs a one-time scripts/rebuild_packmol.sh), and (2) concurrent
CPU-JIT-heavy jobs on the same node hit a transient XLA compile race, mitigated
with retries: 2. Also added a CPU-partition submission path
(config.cpu.yaml + profiles/slurm-cpu/) since this cluster's .venv has no
CUDA jaxlib, so gpu-partition jobs were already silently CPU-only.
Decided: lr_solver in mm_nonbonded. The term now supports mic
(default, jit-compatible, the only solver any sweep backend uses) plus
jax_pme / nvalchemiops_pme / scafacos — but only through the ASE face;
these evaluators are host-orchestrated (e.g. jax-pme builds an ASE Atoms +
its own neighbor list on the host) and cannot be traced inside JaxmdDriver /
RigidBodySampler's jax.jit. The jax face raises NotImplementedError
immediately for non-mic solvers rather than silently falling back to mic.
Remaining (§11): the apocharmm driver (needs a GPU build) and the RDF
validation of rigid sampling (needs a real force-field production run) are the
two items still gated on hardware/scope rather than more unit-testable code.
The --jaxmd-unified flag also doesn't yet cover the pyxtal/template-PDB
builders or handoff/continuation — see §11 for the exact gaps.
1. Motivation¶
Today there are two parallel MD stacks that solve ~80% of the same problem but share almost nothing above the leaf-helper level.
md-system — the orchestrator¶
- File:
mmml/cli/run/md_system.py(~3.3k lines). - Pure dispatcher:
parse_args → build_command() → run_backend()hands off in-process to one of three backend modules: mmml/cli/run/md_pbc_suite/ase.py(ASE dynamics)mmml/cli/run/md_pbc_suite/jaxmd.py(jax-md)mmml/cli/run/md_pbc_suite/pycharmm_mlpot.py(CHARMM dynamics)- Speaks in ASE
Calculatorobjects (MonomerSumCalculator+JAXIntermolecularCalculator, seemmml/interfaces/calculators/hybrid.py). - Generic liquid/crystal builder (packmol, pyxtal, composition, box sizing) plus a campaign / handoff / manifest layer.
- Setups:
{free,pbc}_{nve,nvt,thermalize},pbc_npt,pycharmm_*,lambda_ti.
cg_jaxmd — the research script¶
- File:
examples/cg_jaxmd.py(~2.6k lines), driven byworkflows/cg_jaxmd_ala_water_sweep(Snakemake). - Monolithic, single-system (trialanine in a water box).
- Speaks in a jax-md
energy_fn(R)hand-composed from many terms: ML intramolecular, ML peptide–water dimers, MM nonbonded, SMD bias, φ/ψ restraints, vdW core; plus CHARMM rescue/repair, per-molecule ML charges, and inline NHC / NVE / FIRE loops. - Config = JSON emitted by
scripts/run_setting.py, swept via Snakemake. - Not integrated with
md-system; reuses only leaf helpers (jaxmdInterface/hybrid_energy.py,calculators/simple_inference.py, the pycharmm builders).
The goal: a shared architecture where cg_jaxmd is just a specific
builder + term selection + driver, and md-system --backend jaxmd uses the
same driver and term registry.
2. The core tension¶
There are two irreconcilable energy contracts. This is the one place a single abstraction should not be forced — instead we bridge them.
| ASE side | jax-md side | |
|---|---|---|
| Contract | Calculator.calculate(atoms) → results["energy"/"forces"] |
energy_fn(R, neighbor, box, **kw) pure & jittable |
| State | stateful, numpy / float64 | pure, static shapes, padded pairs |
| Gradients | finite-diff or per-calc | jax.grad, autodiff |
| Neighbors | ASE / rebuilt each call | jax-md neighbor_list, capacity-padded |
Everything below the energy layer (topology, builders) and around it (config, drivers, term selection) can and should be shared. The schema splits exactly along those seams.
3. Design constraints¶
- Two energy faces must both stay first-class (ASE
Calculator⇄ jaxenergy_fn). Bridge, don't merge. - CHARMM is stateful, global, CPU, non-jittable — usable only for building (PSF / topology / FF params) and rescue (minimize / repair), never inside a jitted loop.
- jax-md needs static shapes → dynamic intermolecular / peptide–water
pairs require padding + masks (
cg_jaxmdalready does this with_pad_pairs,_pad_peptide_water_slots). - Energy composition is per-experiment — SMD, restraints, ML-vs-MM intramolecular are toggles. Needs a pluggable term registry, not a hardcoded sum.
- ML models vary (physnet vs spooky:
charges/spinskwargs) — already probed via signature inspection inhybrid_energy.py; keep that behind the term factory. - Ensembles / space are orthogonal to energy: {free, PBC} × {NVE, NVT (NHC | Langevin), NPT} × {minimize, FIRE}.
- Config comes from two mouths (argparse CLI + Snakemake JSON) — both must
lower to one internal
RunConfig.
4. Available backends¶
- Integrator / driver backends: ASE, jax-md, PyCHARMM (Fortran/pyCHARMM), and — proposed — apocharmm (GPU-only CHARMM, C++/CUDA + pybind11).
- Sampler backends (orthogonal to MD integrators): standard MD, and — proposed — rigid-body sampling (rigid monomers: translation + rotation DOF, SHAKE/SETTLE constraints, MC or rigid MD moves).
- Builder backends: packmol (
packmol_placement,tip3_liquid_box,dcm_liquid_box), pyxtal,peptide_builder/protein_charmm_build,trialanine_water_box, template-PDB,setupBox/setupRes. - Energy backends: physnetjax ML (intramolecular monomer, peptide–water
dimer), CHARMM MM nonbonded (
nonbonded_energy_and_forces, jax), CHARMM bonded (cgenff_bonded), biases (SMD, flat-bottom, φ/ψ, COM restraint). - QC / reference backends (eval only, out of scope for the MD loop): orca, pyscf, molpro.
apocharmm (GPU CHARMM) — interface target¶
apocharmm is a GPU-only CHARMM MD
package (C++ ~27%, CUDA ~20%, pybind11 Python bindings). It is a natural
fourth Driver backend, entered at the C++/CUDA level rather than through
the Fortran/pyCHARMM path. Concrete entry points (from include/):
- Container:
CharmmContext(holdsCharmmPSF,CharmmParameters,CharmmCrd,Coordinates) — the object aMolecularSystemwould lower to. - Integrators:
CudaLangevinPistonIntegrator(NPT),CudaLangevinThermostatIntegrator/CudaNoseHooverThermostatIntegrator(NVT),CudaVelocityVerletIntegrator,CudaLeapFrogIntegrator,CudaBAOABIntegrator. - Forces / neighbor lists:
CudaNeighborList,CudaPMEDirectForce,CudaPMEReciprocalForce,CudaBondedForce— native PBC + PME on device. - Free energy:
EDSForceManager,FEPEIForceManager,MBARForceManager,DualTopologyForceManager,PertForceManager— directly relevant to the existinglambda_dynamics/lambda_mbar/lambda_ticode. - I/O:
Subscriberfamily (DcdSubscriber,NetCDFSubscriber,RestartSubscriber,StateSubscriber,CheckpointSubscriber). - Restraints/constraints:
GeometricRestraintForce,HarmonicRestraintForce,Constraints(SHAKE) — the constraint machinery a rigid-body sampler would build on.
Committed direction (decided): the ML forces must enter apocharmm's
device force loop as a custom ForceManager contribution — forces stay on
device, no per-step host round-trips. This is the deep C++/CUDA integration,
not host-side evaluation between steps. The two interface levels below are
therefore a staging order, not an either/or:
- Python (pybind11), host-side ML — validation stepping stone only. Drive
CharmmContext+Cuda*Integratorfrom anApoCharmmDriver, evaluate jax ML forces on the host and add them per step. Used to pin down correctness (energy/force parity) before touching CUDA. Not the shipping path. - C++/CUDA, device-side ML — the target. MMML's jax ML forces enter the
device force loop as a custom
ForceManagerso forces never leave the GPU. The buffer boundary crosses via DLPack capsules (decided, §10) — not raw device pointers — to get shape/dtype/device metadata and proper ownership semantics against jax's async allocator. The custom force is subscribed into apocharmm'sForceManagerComposite(subscribe<ForceType>(...)), so apocharmm owns the step and the boundary lives in one place. The only remaining question is the cross-stream sync mechanism (§10).
5. Proposed schema — 6 layers¶
┌─ orchestration ─────────────────────────────────────────────┐
│ RunConfig (single dataclass) ← argparse CLI | Snakemake │
│ campaign / sweep / handoff / manifest │
└─────────────────────────────────────────────────────────────┘
│ builds │ selects
▼ ▼
┌─ builders ─────────────┐ ┌─ energy terms (registry) ───────┐
│ SystemBuilder.build() │ │ EnergyTerm.make(system, ctx) │
│ packmol / pyxtal / │ │ ml_intra, ml_pep_water, │
│ peptide_water / tmpl │ │ mm_nonbonded, smd, dihedral, │
└───────────┬────────────┘ │ vdw_core, flat_bottom ... │
│ produces └───────────────┬─────────────────┘
▼ │ composed into
┌─ MolecularSystem (shared, immutable) ──┐ ▼
│ R, Z, box, mol_id, monomer_indices, │ HybridEnergy
│ water_indices, psf, ff_params, excl │ ├─ .as_jax_energy_fn() → jax-md
└────────────────────────────────────────┘ └─ .as_ase_calculator() → ASE
│ consumed by
▼
┌─ drivers (integrators) ──────────────────────────────────────┐
│ Driver.run(system, energy_provider, ensemble) → Trajectory │
│ AseDriver | JaxmdDriver (unify runner + cg loop) | CharmmDrv │
└──────────────────────────────────────────────────────────────┘
Layer responsibilities¶
- Orchestration — lowers CLI args and Snakemake JSON into one
RunConfig; owns campaign/sweep/handoff/manifest. - Builders —
SystemSpec → MolecularSystem; thin wrappers over the existingpycharmmInterfacebuild modules. - MolecularSystem — immutable, backend-agnostic topology artifact that every layer above the builders reads.
- Energy terms — one class per physics term, registry-keyed, each able to emit a jax contribution and/or an ASE contribution.
- HybridEnergy — composes the selected terms and exposes both faces
(
as_jax_energy_fn,as_ase_calculator). - Drivers — one per integrator engine (ASE, jax-md, PyCHARMM, apocharmm);
consume a
HybridEnergyand anEnsembleSpec. - Samplers — a sibling of drivers, selected by
RunConfig. MD is the default sampler; rigid-body sampling is an alternative that moves whole monomers as rigid bodies (translation + rotation), via MC moves or constrained (SHAKE/SETTLE) rigid MD. Reuses the sameHybridEnergyandMolecularSystem; only the propagator differs.
6. Protocol sketches¶
# 0. FF parameters — resolved ONCE by the builder, carried as data (decision A, §10).
# Mirrors NonbondedSystemData field-for-field (see hybrid-mlmm-decomposition §2).
# No energy term re-derives exclusions / e14 / LJ tables at runtime.
@dataclass(frozen=True)
class FFParams:
charges: np.ndarray # (N,) partial charges ← nbdata.charges
epsilon: np.ndarray # (N,) LJ epsilon (kcal/mol) ← nbdata.epsilon
rmin_half: np.ndarray # (N,) LJ Rmin/2 (Å) ← nbdata.rmin
at_codes: np.ndarray # (N,) nonbonded type code ← nbdata.at_codes
exclusions: np.ndarray # (M,2) 1-2/1-3 excluded ← nbdata.excluded_pairs
e14_pairs: np.ndarray # (K,2) 1-4 pairs ← nbdata.e14_pairs
# FFParams.from_nonbonded_system_data(nbdata) does the conversion (duck-typed).
# 1-4 SCALING (e14_scale / vdw14) is a pair-build detail, not per-atom FF data.
# 1. Topology — backend-agnostic, immutable; what builders emit / everyone reads
@dataclass(frozen=True)
class MolecularSystem:
R: np.ndarray # (N,3) positions
Z: np.ndarray # (N,) atomic numbers
box: np.ndarray | None # (3,3), or None for free space
mol_id: np.ndarray # (N,) molecule membership
monomer_indices: list[np.ndarray]
water_indices: list[np.ndarray]
psf_path: Path | None
ff_params: FFParams | None # fully-resolved FF state (decision A)
metadata: dict
# 2. Builders — SystemSpec -> MolecularSystem (also the one place FFParams is built)
class SystemBuilder(Protocol):
def build(self, spec: SystemSpec) -> MolecularSystem: ...
# PackmolLiquidBuilder, PyxtalCrystalBuilder, PeptideWaterBuilder, TemplatePdbBuilder
# 3. Energy terms — composable, registry-keyed, engine-agnostic factory.
# Each term declares the neighbor/pair capacity it needs (decision B): the term
# owns its own padding (peptide-water slots vs. intermolecular pairs differ).
@dataclass(frozen=True)
class NeighborRequest:
cutoff_A: float
kind: str # "intermolecular" | "peptide_water" | ...
capacity_hint: int | None # padded slot count; term sizes its own buffers
class TermFns(NamedTuple):
jax_energy_fn: Callable | None # energy_fn(R, nbrs, box, **kw), jittable
ase_contribution: Callable | None # numpy energy/forces for ASE
neighbor_request: NeighborRequest | None
class EnergyTerm(Protocol):
name: str
def neighbor_request(self, system: MolecularSystem) -> NeighborRequest | None: ...
def make(self, system: MolecularSystem, ctx: EnergyContext) -> TermFns: ...
# registry: {"ml_intra", "ml_pep_water", "mm_nonbonded", "smd",
# "dihedral", "vdw_core", "flat_bottom", ...}
class HybridEnergy:
# Gathers each term's NeighborRequest, allocates the union of neighbor lists,
# and routes the right (padded) list to each term at call time.
def __init__(self, terms: list[EnergyTerm], system, ctx): ...
def as_jax_energy_fn(self) -> Callable: ... # for JaxmdDriver
def as_ase_calculator(self) -> Calculator: ... # sum of contributions
# 4. Drivers — one per integrator engine
class Driver(Protocol):
def run(self, system, energy: HybridEnergy, ensemble: EnsembleSpec) -> Trajectory: ...
# AseDriver, JaxmdDriver, CharmmDriver, ApoCharmmDriver
7. Target directory layout¶
Mostly moves, not new code.
mmml/md/
system.py # MolecularSystem, SystemSpec, FFParams
builders/ # SystemBuilder impls (wrap existing pycharmmInterface builders)
energy/
registry.py # EnergyTerm protocol + term registry
terms/ # ml_intra.py, ml_pep_water.py, mm_nonbonded.py,
# smd.py, dihedral_restraint.py, vdw_core.py ← from cg_jaxmd
hybrid.py # HybridEnergy + as_jax_energy_fn / as_ase_calculator
drivers/
ase.py # ← md_pbc_suite/ase.py
jaxmd.py # ← unify jaxmd_runner.set_up_nhc_sim_routine + cg_jaxmd loop
charmm.py # ← md_pbc_suite/pycharmm_mlpot.py
config.py # RunConfig (target of both argparse and Snakemake JSON)
mmml/cli/run/md_system.py # thin: argv -> RunConfig -> Driver
examples/cg_jaxmd.py # thin: JSON -> RunConfig -> PeptideWaterBuilder + terms + JaxmdDriver
8. How the two converge¶
cg_jaxmdbecomes:PeptideWaterBuilderHybridEnergy([ml_intra | ml_pep_water, mm_nonbonded, smd, dihedral, vdw_core])JaxmdDriver.md-system --backend jaxmdbecomes:PackmolLiquidBuilderHybridEnergy([ml_intra, mm_nonbonded])- the same
JaxmdDriver.
The sweep's two energy-mode toggles (use_ml_intramolecular,
peptide_water_ml) become term-selection flags in the registry — no code fork.
9. Highest-leverage moves (recommended order)¶
- Extract
cg_jaxmd's energy terms (SMD, φ/ψ, vdW-core, peptide–water) intommml/md/energy/terms/behind theEnergyTerminterface. This is what un-forks the science and makes the terms reusable frommd-system. - Merge the two jax-md loops —
jaxmd_runner.set_up_nhc_sim_routineandcg_jaxmd's inline NHC / NVE / FIRE — into oneJaxmdDriver. They are currently two independently-drifting integrators; this is the biggest drift risk. - Introduce
MolecularSystem+SystemBuilderas thin wrappers over the existing pycharmm builders, so both entry points build through one seam. - Lower both configs into
RunConfig; makemd_system.pyandexamples/cg_jaxmd.pythin front-ends.
Suggested migration sequence (non-breaking)¶
- Land the protocols/dataclasses (
system.py,energy/registry.py,config.py) with no behavior change. - Wrap existing builders →
builders/; keep old call sites delegating. - Extract terms one at a time; validate energy parity against
cg_jaxmddiagnostics (diagnose_energy,run_force_and_nl_diagnostics) at each step. - Build
JaxmdDriverfromjaxmd_runner, then portcg_jaxmdonto it behind a flag; compare trajectories before removing the inline loop. - Flip
md_system --backend jaxmdontoJaxmdDriver; retire the duplicate.
10. Decisions & open questions¶
Decided¶
FFParamsboundary → captured as data (option A). CHARMM FF state (charges, LJ tables, exclusions, e14 / vdw14) is resolved once by the builder and carried onMolecularSystem.ff_params(see §6). Energy terms read it; none re-derive it at runtime. This removes the inline recomputationcg_jaxmdcurrently does and makes term evaluation pure w.r.t. FF state.- Neighbor-list ownership → per-term capacities (option B). Each
EnergyTermdeclares aNeighborRequest(cutoff, kind, padded capacity) and owns its padding;HybridEnergyallocates the union and routes the right list to each term. Peptide–water slots and intermolecular pairs keep independent capacities rather than sharing one driver-owned list. - apocharmm ML forces → device-side custom
ForceManager(decided). ML forces enter apocharmm's device force loop; they do not stay host-side. Host-side pybind11 evaluation is only a validation stepping stone (§4). - apocharmm buffer transport → DLPack capsule (decided). jax↔apocharmm
device buffers cross as DLPack capsules (
jax.dlpack.to_dlpack/from_dlpack), not raw CUDA device pointers. Rationale: DLPack carries shape/dtype/device/strides and a capsule deleter, so it is safe against jax's async allocator and buffer donation (a raw pointer pulled from jax can be reused/freed under us, and layout mismatches fail silently). A thin pybind11 shim wraps apocharmm'sCudaContainer/DeviceVectorbuffers asDLManagedTensor. The unavoidablefloat4 xyzq(f32) ↔N×3/doublelayout+precision transform is a small custom kernel on top of the wrap. - apocharmm owns the step (decided). The custom
MLForceManageris subscribed into apocharmm'sForceManagerComposite(subscribe<ForceType>(tag, stream, values, energyVirial)) with its own stream; apocharmm's integrator drives the loop and the DLPack handoff lives in one place (MLForceManager::calcForce), accumulating (+=) into its forcevalues. The jax-owns-loop / XLA-custom-call alternative is rejected — apocharmm is a stateful CUDA engine, not a leaf op. - Rescue path → explicit driver hook. CHARMM repair/minimize is modelled as
a driver hook (
on_overlap), never hidden inside an energy term, so energy terms stay pure and the impurity is confined to the driver. - Rigid-body DOF → quaternions (decided). Rigid-monomer orientation is represented as a unit quaternion (translation = COM), renormalized after each move. Preferred over rotation matrices: 4 numbers vs. 9, no orthogonality drift to project out, and clean composition for MC rotation moves.
- Rigid sampling →
Samplerpeer, not a driver constraint mode (decided). Rigid-body sampling is its ownSampler(peer of the MDDriver, selected byRunConfig), reusing the sameMolecularSystem+HybridEnergy. It is not bolted into each integrator as a constraint mode, so it composes with any energy term and any backend without touching the drivers.
Still open¶
- Cross-stream sync mechanism — how to order the jax ML-force stream against
apocharmm's force-manager stream. Options: (1) host barrier per step
(
block_until_ready+cudaStreamSynchronize) — correct, on-device, but serializes; (2) event-based overlap (cudaEventRecord/cudaStreamWaitEvent) — blocked by jax not exposing its CUDA stream publicly; (3) DLPack stream-aware protocol (__dlpack__(stream=...), needs recent jax + shim support). Plan: ship the barrier (1) first, measure, and only pursue overlap (3) if throughput-bound — the transport (DLPack) and structure (apocharmm-owns-loop) are settled; only the ordering optimization is open.
This is the only genuinely open design question; it is deferred to implementation/measurement rather than blocked. All other seams are decided.
PyCHARMM / MPI runtime contract¶
PyCHARMM is not an ordinary optional Python import in this stack. The vendored
package loads a native libcharmm that may itself be linked to OpenMPI and
Fortran. A native STOP, MPI_Abort, loader mismatch, or segmentation fault
terminates the interpreter; pytest cannot convert it into a Python traceback.
Consequently, successful import of the MMML bootstrap module is not proof that
live CHARMM is available. The authoritative runtime signal is
import_pycharmm.PYCHARMM_AVAILABLE.
There are three supported execution modes:
| Mode | Intended work | Contract |
|---|---|---|
| Offline/unit | restart parsing and text patching, builders mocked at the CHARMM boundary | pytest sets MMML_WARMUP_MLPOT_JAX_ONLY=1; do not enter native PyCHARMM |
| Live, one rank | PSF/parameter reads, restart writes, CHARMM parity tests | launch with MMML_MPI_NP=1 ./scripts/mmml-charmm-mpirun.sh ... when libcharmm is MPI-linked |
| Live, multiple ranks | DOMDEC / spatial ML MPI | launch with the same wrapper and an explicit MMML_MPI_NP; validate Tier 2 first |
The launcher is part of the runtime ABI, not merely a convenience wrapper: it
selects the OpenMPI installation compatible with libcharmm, establishes the
library search path, constrains OpenMP threads, and orders MPI/mpi4py/JAX
initialization. Do not run an MPI-linked libcharmm from a plain pytest or
python process merely because import pycharmm appears to work.
Preflight and test recipes:
# Loader/OpenMPI/mpi4py survey before launching a job
mmml mpi-check --strict
# Live one-rank PyCHARMM smoke in the real runtime
MMML_MPI_NP=1 ./scripts/mmml-charmm-mpirun.sh \
python -m pytest tests/charmm_mpi/test_mpi_live_energy.py -q
# Spatial MPI validation before a multi-rank MD run
mmml mpi-check --tier2 --prelaunch --strict
MMML_MPI_NP=2 MMML_MLPOT_SPATIAL_MPI=1 \
./scripts/mmml-charmm-mpirun.sh mpi-check --tier2 --strict
For normal runs, the implemented mpi-launch front end keeps environment,
rank topology, and JAX policy separate. uv run resolves Python once and the
selected sys.executable is forwarded through the existing ABI-aware wrapper:
uv run mmml mpi-launch --preset single -- md-system --config run.yaml
uv run mmml mpi-launch --preset cpu --jax-cpu-threads 16 -- \
md-system --config run.yaml
uv run mmml mpi-launch --preset spatial --mpi-ranks 4 -- \
md-system --config run.yaml --ml-spatial-mpi
The presets are aliases over independent --mpi-ranks, --jax-mode,
--jax-cpu-threads, and --charmm-omp-threads settings. This also supports
multi-rank CHARMM with rank-0 JAX and generic GPU-per-rank JAX without claiming
that either is spatial decomposition. See mpi-launch.
MMML_NO_CHARMM_MPI=1 is only valid for a genuinely serial libcharmm; it must
not be used to disguise an MPI-linked build. Likewise,
MMML_WARMUP_MLPOT_JAX_ONLY=1 means "bootstrap/import seams only" and must
never be interpreted as permission to call CHARMM state, coordinate, energy,
or restart APIs.
The handoff boundary follows this rule explicitly: with a usable restart
template and PYCHARMM_AVAILABLE=False, save_handoff_to_res performs the
offline, validated text patch. It enters _write_handoff_restart_via_charmm
only when native CHARMM is genuinely loaded. This prevents the former failure
where test_save_handoff_to_res_with_template ended at native level with exit
code 1 and no pytest traceback.
For the full operational matrix and cluster troubleshooting, see
PyCHARMM + MPI, mpi-check, and
health-check.
11. Roadmap checklist¶
Tracking checklist for the unification, plus the two newly-scoped capabilities (rigid sampling, apocharmm interface). Migration order is deliberate: land seams non-breaking, extract terms with parity checks, then add backends.
Core unification (from §9)¶
- [x] Land protocols/dataclasses (
system.py,energy/registry.py,config.py) with no behavior change. (Done:mmml/md/package —system.py(FFParams/MolecularSystem/SystemSpec),config.py(RunConfig/EnsembleSpec),results.py(Trajectory),energy/registry.py(EnergyTerm/TermFns/NeighborRequest/HybridEnergy+ term registry), andbuilders/drivers/samplersProtocols. Imports pull in no jax/ASE/CHARMM; smoke-tested intests/unit/test_md_package_seams.py.) - [~] Wrap existing builders into
builders/(SystemBuilder); old call sites delegate. (In progress.) - [x]
PsfSystemBuilder— builds aMolecularSystemfrom a PSF + params + coords; populatesFFParamsfield-for-field fromNonbondedSystemData(FFParams.from_nonbonded_system_data) and partitions molecules from the PSF bond graph (molecule_ids_from_bonds). CHARMM-integration tested onpept.psf/pept.pdb:tests/unit/test_md_builders.py. - [x] Packmol / PyXtal / peptide-water placement wrappers — legacy builders
now lower PSF-ordered coordinates, molecule partitions, box, residue
identity, and water groups into
MolecularSystem. Packmol/PyXtal can optionally lower a supplied emitted PSF toFFParams; the trialanine-water wrapper always delegates its emitted PSF/parameter files throughPsfSystemBuilder. Backend calls are lazy and mock-backed seam tests require no CHARMM/Packmol/PyXtal:tests/unit/test_md_builders.py. Old front-end call sites still need to delegate before the parent item is complete. - [x] Extract
cg_jaxmdenergy terms intoenergy/terms/one at a time, validating at each step. (Complete — all six terms extracted, registered, and parity-tested.) - [x]
smd—SMDBiasTerm(moving harmonic end-to-end restraint, PBC/free,lambda_tsteering). Parity-tested vs. the cg_jaxmd formula + ASE forces vs. finite difference:tests/unit/test_md_energy_terms.py. - [x]
dihedral—DihedralRestraintTerm(φ/ψ backbone restraints). Parity-tested vs. the cg_jaxmd formula. - [x]
vdw_core—RepulsiveCoreVdwTerm(repulsive LJ wall between a core atom group and molecule groups, smooth cutoff, optional padded slots). Parity-tested vs. the cg_jaxmd formula; takes explicit ε/Rmin arrays now, to be sourced fromFFParamsby the builder. - [x]
mm_nonbonded—MMNonbondedTerm. Faithful (B)-block extraction reusing the shared switching kernels (_pair_vdw_energy/_pair_elec_energy): jittable padded-pair jax face (pair_i/pair_j/pair_mask/e14/vdw14) plus an ASE face wrappingnonbonded_energy_and_forces. ConsumesFFParams. Parity-tested to ~1e-9 vs. the reference (energy + forces, host/padded/ jit/ASE):tests/unit/test_md_mm_nonbonded.py. - [x]
ml_intra—MLIntramolecularTerm(per-monomer ML energy) andml_pep_water—MLCoreGroupTerm(core–group ML dimers over paddedactive_group_slots). Thin wrappers overjaxmdInterface.hybrid_energy; model+params come from the runEnergyContext(model-agnostic). Validated against the factory calls + jit + composition using the last-epoch example checkpoint (examples/sppoky-epoch-0010_params.json):tests/unit/test_md_ml_terms.py. - [x] Build
JaxmdDriverfromjaxmd_runner.set_up_nhc_sim_routine. Shared driver inmmml/md/drivers/jaxmd.py: lazy optional imports, free/PBC NVE, NVT-NHC, FIRE, NPT (Nosé–Hoover barostat), block-boundary per-term neighbor refresh, overlap-repair hook, partial-final-block recording, and optional NPZ output. NPT keeps the real-space terms via a fractional↔real bridge and box-aware terms (mm_nonbonded/vdw_core/smdaccept aboxkwarg; box threaded each step). Tests (FIRE/NVE/NVT, fixed-box PBC, box-evolving NPT):tests/unit/test_md_jaxmd_driver.py. - [~] Wire the front-ends onto the shared stack. (Assembly glue landed:
mmml/md/assemble.py— builder registry (get_builder/available_builders/build_system),build_hybrid_energy(term registry + per-term kwargs), andassemble_and_run(RunConfig → builder → HybridEnergy → JaxmdDriver). End-to-end tested from aRunConfig:tests/unit/test_md_assemble.py.) - [x] Lowering adapters (
mmml/md/lowering.py):runconfig_from_md_system_args(argparseNamespace→RunConfig) andrunconfig_from_cg_config(cg_jaxmd JSON + phase →RunConfig), withterms_from_cg_configimplementing the sweep-toggle → term-selection mapping (doc §8). Pure / numpy-only; tested intests/unit/test_md_lowering.py. - [x]
mmml/cli/run/md_system_unified.py— an opt-in--jaxmd-unifiedflag onmmml md-system --backend jaxmdthat routes throughrunconfig_from_md_system_args→assemble_and_runinstead of the legacymd_pbc_suite/jaxmd.pyinline loop. Opt-in (not the default) so the existing path is untouched until this covers more of md-system's feature surface — only the packmol composition builder is wired;--builder pyxtal,--template-pdb, and--continue-fromraiseNotImplementedErrorrather than silently diverging.build_packmol_system_with_ffparamscloses a real gap: the packmol composition builder (build_packmol_composition_cluster) never writes a PSF file, so it writes the live CHARMM PSF to a scratch file to resolveFFParams(mirrors_lower_optional_psf). Also explicitly callsensure_pycharmm_loaded()before building — unlikePeptideWaterSystemBuilder's underlying builder, the packmol composition builder does not self-bootstrap CHARMM. Validated end-to-end via the real CLI (mmml md-system --backend jaxmd --jaxmd-unified ...) for bothpbc_nveandpbc_nvtwith a real checkpoint; existing (non-flagged)--backend jaxmdbehavior is unaffected (66 pre-existingmd_systemCLI tests still pass). Tests intests/unit/test_md_system_unified.py(9 tests: 6 fast config validation + 3 real end-to-end CHARMM+checkpoint runs, includingFFParamsresolution and bothpbc_nve/pbc_nvt). The full legacy dispatch remains the default; retiring the duplicate inline loop is deferred until the unified path covers pyxtal/template builders and handoff. - [x]
examples/cg_jaxmd_unified.py— a thin, validated front-end overassemble_and_runfor cg_jaxmd-style peptide-water runs: reads a cg_jaxmd JSON config, chainsfire → nvt → nvephases (carrying positions forward between phases viadataclasses.replace), and wiresml_intra/mm_nonbonded/ml_pep_water/vdw_core/smdterm kwargs from the config toggles. Deliberately raisesNotImplementedErrorforconstrain_phi_psi(not yet wired) rather than silently diverging from the legacy script. This is a new parallel script, not an in-place rewrite of the 2624-lineexamples/cg_jaxmd.py— safer to validate, andcg_jaxmd.pyis untouched. Validated end-to-end against real CHARMM builds (PeptideWaterSystemBuilder) + the example checkpoint: single-phase runs, multi-phase position chaining (nvt's start energy matchesfire's end energy exactly), and thepeptide_water_ml+peptide_water_ml_core_vdwcombination. Tests intests/unit/test_cg_jaxmd_unified.py(11 tests: config validation, pureterm_kwargsderivation, and 3 real end-to-end CHARMM+checkpoint integration tests).
Rigid sampling¶
- [x] Define a
Samplerprotocol (peer ofDriver) selected byRunConfig; MD is the default sampler.assemble_and_rundispatchesconfig.sampler == "rigid"to the sampler. - [x] Rigid-body state: per-monomer COM + unit-quaternion orientation, groups
from
MolecularSystem.monomer_indices(quat_from_axis_angle/quat_to_matrixinmmml/md/samplers/rigid.py). - [x] Rigid-move propagator: Metropolis MC translation + quaternion rotation,
reusing the jitted
HybridEnergy. Rigid moves preserve every intramolecular distance (tested to 1e-8):tests/unit/test_md_samplers.py. - [x] Acceptance uses the composed energy, so existing bias terms (SMD, etc.)
participate automatically (exercised via
assemble_and_runintests/unit/test_md_assemble.py). - [x]
neighbor_fnsupport (found while building the cross-backend sweep, §11 "Cross-backend validation workflow"):RigidBodySampleroriginally had no way to supplymm_nonbonded's padded pair list, so any intermolecular-pair term raisedTracerArrayConversionErrorunder the sampler'sjax.jit.RigidBodySamplernow acceptsneighbor_fn/neighbor_refresh_every(rebuilt once per N sweeps, not per proposed move — rebuilding on the host for every trial move would dominate runtime) mirroringJaxmdDriver;assemble_and_runauto-wires it for the rigid path exactly as it already did for the MD driver (refactored into a shared_auto_neighbor_fnhelper). Regression-tested intests/unit/test_md_samplers.py(fails without a neighbor_fn, passes with one, and the auto-wiring is tested throughassemble_and_run). - [ ] Validate rigid sampling reproduces liquid structure (RDF) vs. flexible MD on a small box. (Needs a real forcefield run — deferred.)
- [ ] Optional: constrained (SHAKE/SETTLE) rigid MD variant. (MC path landed first; not required for the sampler to be usable.)
Cross-backend validation workflow¶
- [x]
workflows/unified_backend_sweep/— a Snakemake workflow that runs every driver/sampler + ensemble combination reachable throughassemble_and_run(jaxmd_min,jaxmd_nve,jaxmd_nvt,jaxmd_npt,rigid_mc) against the same small real system (4-water TIP3 box via the packmol composition builder,ml_intra+mm_nonbonded, the example checkpoint), across several seeds, and aggregates a pass/fail + energy-drift summary. This is what surfaced theRigidBodySamplerneighbor_fn gap above, and confirmed NPT's much longer JIT-compile time (~4 min vs. ~30s for the fixed-box ensembles on this tiny system — reflected in the workflow's per-backendruntime_min). Seeworkflows/unified_backend_sweep/README.md. - [x] Submitted the sweep on the real cluster (pc-studix, SLURM
gpupartition). Found and fixed three more real deployment gaps this surfaced (none were bugs in the sweep itself): - The vendored Packmol binary isn't committed to git (platform-specific);
a fresh checkout needs
bash scripts/rebuild_packmol.shonce. Documented in the workflow README. - Two CPU-JIT-heavy jobs (
jaxmd_npt) landing on the same node concurrently hit a transient XLA CPU-backend compile race (JaxRuntimeError: ... Failed to materialize symbols). Mitigated withretries: 2inprofiles/slurm(-cpu)(clears on a fresh node allocation) rather than chasing the underlying XLA race. - This cluster's mmml
.venvhas no CUDA jaxlib, sogpu-partition jobs were already silently falling back to CPU. Addedconfig.cpu.yaml+profiles/slurm-cpu/(Snakemake deep-merges--configfile config.yaml config.cpu.yaml) to submit to the CPU partition (short) directly instead — identical behavior, better use of the shared GPU allocation. - [x]
lr_solversupport inmm_nonbonded(decided): the term now acceptslr_solver(micdefault, plusjax_pme/nvalchemiops_pme/scafacosinherited fromnonbonded_energy_and_forces), but only through the ASE face —jax_pme's evaluator builds an ASEAtomsand its own neighbor list on the host and returns plain numpy, which cannot be traced inside thejax.jitgraphJaxmdDriver/RigidBodySamplerboth require. The jax/jit face raisesNotImplementedErrorimmediately for any non-micsolver rather than silently usingmicunder the hood. Tested intests/unit/test_md_mm_nonbonded.py(jax face rejects cleanly; ASE face withjax_pmegives a finite energy that differs frommic, proving the kwarg actually reaches the reference call). Not wired into anyunified_backend_sweepbackend — doing so would just raise, which is correct but not a useful smoke-test row. Real jax_pme/etc. support needs either a non-jit ASE-based driver (not built) or ajax.pure_callbackbridge around the evaluator — tracked here, not attempted this session.
Mixed-system (ML core + MM/ML solvent) support checklist¶
The project goal behind all of §11 is a validated mixed system: one
ML-scored "core" region (e.g. a peptide) embedded in a solvent that is
partly ML (near the core) and partly classical MM (far from it), evaluated
through the same assemble_and_run seam as every pure-MM or pure-ML setting.
This tracks what's actually wired and tested for that goal today, vs. still
open. See workflows/mixed_calculator_sweep/ for the 10000-step validation
sweep that exercises the checked items below on a real trialanine+TIP3-water
system with a real ML checkpoint.
Builder / topology
- [x] PeptideWaterSystemBuilder (mmml/md/builders/placement.py) — one ML
"core" molecule (e.g. trialanine) + N TIP3 waters, all built through
live CHARMM (PSF + params), lowered into one MolecularSystem with
monomer_indices[0] = core atoms and water_indices = per-water atom
groups.
- [x] SystemSpec(builder="peptide_water", box_size=..., n_molecules=<n_waters>,
seed=...) — box sizing and water count are builder params, not
hardcoded.
- [ ] Builder support for a non-peptide ML core (e.g. a small-molecule
solute) — PeptideWaterSystemBuilder is specifically trialanine-shaped
today (build_trialanine_water_box_in_charmm); a general "ML solute +
MM solvent" builder is not yet factored out.
Energy terms (the ML/MM boundary itself)
- [x] ml_pep_water (MLCoreGroupTerm) — ML dimer energy between the core
and each water, gated by interaction_cutoff_A /
interaction_switch_width_A (unset = every core-water pair is ML,
correct for small systems, not capacity-padded for large ones).
- [x] vdw_core (RepulsiveCoreVdwTerm) — repulsive-only LJ wall from the
core to all solvent, filling the region outside ml_pep_water's
cutoff so waters that fall outside the ML shell don't collapse into the
core unscored. Composes with ml_pep_water (mixed_core_vdw setting
in the sweep) rather than replacing it.
- [x] mm_nonbonded scores water-water (MM/MM) interactions in the same
composed HybridEnergy, so a mixed run is ml_intra (core
intramolecular) + ml_pep_water (core↔near-water, ML) + vdw_core
(core↔far-water, repulsive wall) + mm_nonbonded (water↔water, MM) —
four terms, one HybridEnergy, one driver.
- [ ] A smooth ML→MM handoff for individual core-water pairs as a water
crosses ml_pep_water's cutoff during dynamics (today: ml_pep_water
is evaluated over a fixed, build-time-selected group; there is no
dynamic re-assignment of a water from "ML-scored" to "MM-scored" mid
trajectory). This is the actual "complementary handoff" problem
referenced in long_range_backend.py's jax-pme hybrid path
(complementary_handoff kwarg) but not yet threaded through
assemble_and_run's term composition.
Calculator / force-field knobs (validated via mixed_calculator_sweep's
checkpoint / calculator / mm_nonbonded setting axes)
- [x] Swapping the ML checkpoint (create_calculator_from_checkpoint's
checkpoint_path) — different training epochs, same term wiring.
- [x] electrostatics_damping_sigma — Gaussian damping of the ML model's
point-Coulomb electrostatics; baked into the extracted model/params
at calculator-construction time (checkpoint_loading.py), so it
affects ml_intra/ml_pep_water identically to how it'd affect a bare
ASE calculator.
- [x] disable_physnet_point_coulomb — toggles the ML model's own
point-charge Coulomb term off (e.g. when mm_nonbonded's classical
Coulomb, or an external long-range solver, should be the only source of
electrostatics for MM-scored atoms).
- [x] mm_nonbonded cutoffs (cutnb/ctonnb/ctofnb) via EnergyContext.options
— independent of the ML side, only affects the MM/MM water-water term.
- [ ] lr_solver other than mic (jax_pme/nvalchemiops_pme/scafacos) in
a mixed system's mm_nonbonded — same limitation as the pure-MM case
(§ "Cross-backend validation workflow" above): ASE-face only, not
wireable into JaxmdDriver's jit loop. Not exercised by
mixed_calculator_sweep for the same reason.
Validation depth
- [x] 10000-step NVE energy trace (100 recorded samples via JaxmdDriver's
default record_every=100) per setting — enough to see real
conservation behavior, not just a 2-3 point endpoint delta. Fully
achieved for all 6 water_box settings on the real cluster
(pc-studix, 2026-07-12). For the 6 mixed peptide_water settings,
achieved at a reduced scale (n_waters=8, n_steps=200, 3 recorded
samples) after two rounds of measured-not-guessed downsizing — see
below; 12/12 settings completed, confirmed in
results/summary.md on the cluster.
- [x] Mixed-system (peptide_water) settings: the first attempt
(n_waters=20, n_steps=10000) ran clean for ~3 real hours with no
crash (confirmed via ps: steady CPU usage, stable memory) but was
intentionally cancelled — water_baseline finished all 10000 steps in
~324s in the same window, while the mixed settings hadn't finished
after ~10800s, at least ~33x slower. Root cause, verified against
the term's source (not just inferred from timing):
ml_pep_water's interaction_cutoff_A does not reduce its
per-step cost — checked directly in
mmml/interfaces/jaxmdInterface/hybrid_energy.py::make_peptide_water_ml_energy_fn,
every core-water dimer is vmapped through the ML model every step
regardless of the cutoff; the cutoff only applies a post-hoc energy
switching weight (correct physics, zero runtime effect). Unlike
mm_nonbonded, no neighbor-list-style pruning exists for this term
yet (no active_group_slots/active_group_mask auto-wiring in
assemble.py for ml_core_group-kind NeighborRequests — tracked as
an open item below). Fix, in two measured rounds: (1) n_waters
20→8 + n_steps 10000→2000 — still hit the 240 min Slurm walltime
(TIMEOUT); a local 20-step check at n_waters=8 measured ~9.85s/step
even past JIT compile (this term's cost scales ~linearly with steps,
unlike the compile-dominated water_box case), so 2000 steps ≈ 5.5h,
consistent with the timeout. (2) n_steps cut again to 200 (~33 min
at the measured rate) — this completed successfully, 12/12.
interaction_cutoff_A is kept (still physically correct) but is not a
speed fix. See workflows/mixed_calculator_sweep/README.md "Real-run
status" for the full incident log (packmol binary clobbered with a
macOS build — recurred 3 times across this session with no identified
cause, AVX2-vs-non-AVX2 node split, OOM at the original memory budget,
Snakemake driver dying on SSH disconnect, now run in tmux — all
fixed or mitigated).
- [x] Found via this same real run: mixed_core_vdw showed a genuine
initial-configuration energy blow-up (energy_initial_ev ≈ 1.7M,
vs. the other mixed settings' ~-500 to -560 eV range) at step 0,
before any dynamics, relaxing to a sane value within ~100 steps as NVE
forces fling the clashing atoms apart. Root cause: this workflow runs
raw NVE directly from the packmol-built geometry with no minimization
step first; vdw_core's repulsive wall is steep at short range, and
packmol's placement tolerance (2.0 Å) can leave a close initial
core-water contact that only vdw_core (not the other mixed settings,
which lack that term) punishes severely. Fix for a real production run
using vdw_core: minimize first (e.g. the fire → nvt → nve chain
examples/cg_jaxmd_unified.py already implements) rather than starting
dynamics straight from the raw placement. Not yet fixed in
mixed_calculator_sweep itself (documented as a known caveat in its
README instead, since fixing it means adding a minimization phase to
scripts/run_setting.py, out of scope for this validation pass).
- [ ] Real neighbor-list support for ml_core_group-kind terms (wiring
active_group_slots/active_group_mask through assemble.py's
auto-wiring, mirroring the intermolecular kind) so
peptide_water_ml_cutoff_A becomes a genuine per-step compute
reduction rather than only an energy-switching weight. Needed before
larger n_waters/n_steps mixed-system runs are tractable.
- [ ] NVT/NPT for mixed systems — mixed_calculator_sweep only runs NVE;
jaxmd_npt's deterministic cluster-specific failure (documented in
workflows/unified_backend_sweep/README.md) would need root-causing
before extending the mixed sweep to NPT.
- [ ] RDF / structural validation of the ML-water shell against a reference
(flagged as open in "Rigid sampling" above for the pure-MM case; doubly
untested for the mixed ML/MM boundary).
- [ ] Build apocharmm and confirm the pybind11 module imports in the MMML env (CUDA 11.1.1+, GCC 10.1+, NetCDF4).
- [ ]
MolecularSystem→CharmmContextlowering (CharmmPSF,CharmmParameters,CharmmCrd). - [ ]
ApoCharmmDriver(Python/pybind11): mapEnsembleSpectoCudaLangevinPistonIntegrator(NPT) /CudaLangevinThermostatIntegrator/CudaNoseHooverThermostatIntegrator/ velocity-Verlet. - [ ] Wire
SubscriberI/O (DcdSubscriber/NetCDFSubscriber/RestartSubscriber) into the sharedTrajectoryoutput. - [ ] Host-side ML (validation stepping stone, not shipping): evaluate MMML jax forces between steps and add them via a pybind11 force hook; use only to pin energy/force parity before touching CUDA.
- [ ] Device-side ML — the target (committed, §10): expose MMML ML forces
as a custom apocharmm
ForceManagerso forces never leave the GPU. - [ ] pybind11 shim wrapping apocharmm
CudaContainer/DeviceVectorbuffers as DLPackDLManagedTensor(both directions); confirm zero-copyjax.dlpack.from_dlpack/to_dlpackround-trip on a toy buffer. - [ ] Layout+precision kernel:
float4 xyzq(f32) →N×3for the model, and model forces → accumulate (+=) into thedoubleforce buffer. - [ ] Implement
MLForceManagerand subscribe it into theForceManagerComposite(subscribe<ForceType>(tag, stream, values, energyVirial)) with its own stream; DLPack handoff + jax call live incalcForce. - [ ] Sync: host barrier first (
block_until_ready+cudaStreamSynchronize) — correct and on-device; validate forces match the host-side reference bitwise-close. - [ ] Benchmark step throughput vs. host-side path; only if throughput-bound,
move to DLPack stream-aware sync (
__dlpack__(stream=...)) for overlap. - [ ] Bridge free-energy managers (
FEPEIForceManager,MBARForceManager,EDSForceManager,DualTopologyForceManager) to the existinglambda_dynamics/lambda_mbar/lambda_tipaths. - [ ] Parity check: apocharmm vs. PyCHARMM vs. jax-md energies/forces on a shared PBC box.