Distributed Multi-GPU Pipeline: Parallel FIRE → Langevin#

This example orchestrates two independent FIRE → NVTLangevin pipelines running in parallel across 4 GPUs using DistributedPipeline.

Topology

digraph topology { rankdir=LR fontname="Helvetica" node [fontname="Helvetica" fontsize=11 shape=box style="rounded,filled" fillcolor="#dce6f1" fontcolor="#111111"] edge [fontname="Helvetica" fontsize=10] r0 [label="Rank 0\nFIRE + sampler_a"] r1 [label="Rank 1\nNVTLangevin + sink_a" fillcolor="#f9e2ae" fontcolor="#111111"] r2 [label="Rank 2\nFIRE + sampler_b"] r3 [label="Rank 3\nNVTLangevin + sink_b" fillcolor="#f9e2ae" fontcolor="#111111"] r0 -> r1 [style=bold color="#c0392b" penwidth=2] r2 -> r3 [style=bold color="#c0392b" penwidth=2] }

Two independent FIRE → Langevin pipelines across 4 GPUs.#

Each FIRE rank draws molecules from a dataset, optimises them until convergence or a 50-step limit, and sends them to the paired Langevin rank for 20 steps of short MD production. A SnapshotHook on the Langevin ranks writes the trajectories to a HostMemory sink.

Note

This example requires 4 GPUs. Run with:

torchrun --nproc_per_node=4 examples/distributed/01_distributed_pipeline.py

For CPU-only testing, change backend="nccl" to backend="gloo".

from __future__ import annotations

import logging
import os

import torch
import torch.distributed as dist
from ase.build import molecule
from loguru import logger

from nvalchemi.data import AtomicData
from nvalchemi.dynamics import (
    FIRE,
    ConvergenceHook,
    DistributedPipeline,
    FusedStage,
    HostMemory,
    NVTLangevin,
    SizeAwareSampler,
)
from nvalchemi.dynamics.base import BufferConfig
from nvalchemi.dynamics.hooks import SnapshotHook
from nvalchemi.models.demo import DemoModel, DemoModelWrapper

logging.basicConfig(level=logging.INFO)

# Distributed examples are launcher-only. Sphinx sets this flag during docs
# builds, and torchrun sets rank/world-size variables during real launches.
_DOCS_BUILD = os.environ.get("NVALCHEMI_SPHINX_BUILD") == "1"
_DISTRIBUTED_ENV = "RANK" in os.environ and "WORLD_SIZE" in os.environ
_RUN_DISTRIBUTED_EXAMPLE = _DISTRIBUTED_ENV and not _DOCS_BUILD

# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------


def atoms_to_data(atoms) -> AtomicData:
    """Convert an ASE Atoms object to AtomicData with dynamics fields."""
    data = AtomicData.from_atoms(atoms)
    n = data.num_nodes
    data.forces = torch.zeros(n, 3)
    data.energy = torch.zeros(1, 1)
    data.add_node_property("velocities", torch.zeros(n, 3))
    return data


class InMemoryDataset:
    """Minimal dataset wrapper for ``SizeAwareSampler``."""

    def __init__(self, data_list: list[AtomicData]) -> None:
        self._data = data_list

    def __len__(self) -> int:
        return len(self._data)

    def __getitem__(self, idx: int) -> tuple[AtomicData, dict]:
        d = self._data[idx]
        return d, {"num_atoms": d.num_nodes, "num_edges": d.num_edges}

    def get_metadata(self, idx: int) -> tuple[int, int]:
        """Get the metadata associated with idx"""
        d = self._data[idx]
        return d.num_nodes, d.num_edges


# ---------------------------------------------------------------------------
# Build molecules
# ---------------------------------------------------------------------------


def build_dataset() -> list[AtomicData]:
    """Create a handful of small rattled molecules."""
    names = ["H2O", "CH4", "NH3", "H2O", "CH4", "NH3", "H2O", "CH4"]
    data_list = []
    for name in names:
        atoms = molecule(name)
        atoms.rattle(stdev=0.15)
        data_list.append(atoms_to_data(atoms))
    return data_list

DistributedPipeline topology#

DistributedPipeline maps integer GPU ranks to dynamics stage instances. When pipeline.run() is called, each process (launched by torchrun) looks up its own rank and executes only the corresponding stage. Inter-rank communication is handled transparently via NCCL isend/irecv calls.

Here we create two independent sub-pipelines: ranks 0→1 and 2→3. Each sub-pipeline is a FIRE optimiser feeding into a Langevin MD stage. The stages dict is built on every rank but only the local stage is ever executed; constructing all stages on every rank keeps the code identical across processes, which simplifies debugging.

Dataset and samplers#

SizeAwareSampler draws molecules from a dataset and packs them into variable-size batches that fit within atom and edge count budgets. In a distributed setting, each upstream rank owns its own sampler and dataset partition so that work is distributed evenly. Downstream ranks (the Langevin stages) do not need a sampler — they receive systems directly from the paired FIRE rank via NCCL.

BufferConfig: fixed-size communication#

NCCL requires that every isend/irecv pair transfers an identical number of bytes. BufferConfig specifies the fixed sizes (num_systems, num_nodes, num_edges) used to pre-allocate communication buffers on both sender and receiver. Choose values that are at least as large as the largest batch you expect to send in a single step; excess capacity is padded with zeros and stripped on receipt.

Stage construction#

Each stage is wired to its neighbours via prior_rank and next_rank. An upstream stage has prior_rank=None and next_rank=<downstream>; a downstream stage has prior_rank=<upstream> and next_rank=None. The buffer_config must be identical for all stages that communicate with each other.

def make_fire(model: DemoModelWrapper, rank: int, **kwargs) -> FusedStage:
    """Create a convergence- or step-limited FIRE optimiser stage."""
    dynamics = FIRE(
        model=model,
        dt=1.0,
        n_steps=50,
        convergence_hook=ConvergenceHook(
            criteria=[
                {
                    "key": "forces",
                    "threshold": 0.05,
                    "reduce_op": "norm",
                    "reduce_dims": -1,
                }
            ],
        ),
    )
    return FusedStage(sub_stages=[(0, dynamics)], **kwargs)


def make_langevin(
    model: DemoModelWrapper,
    sink: HostMemory,
    rank: int,
    **kwargs,
) -> FusedStage:
    """Create a fixed-duration NVTLangevin stage with trajectory snapshots."""
    dynamics = NVTLangevin(
        model=model,
        dt=0.5,
        temperature=300.0,
        friction=0.01,
        n_steps=20,
        hooks=[SnapshotHook(sink=sink, frequency=1)],
    )
    return FusedStage(sub_stages=[(0, dynamics)], **kwargs)

Running the pipeline#

DistributedPipeline is used as a context manager. On __enter__ it initialises the PyTorch process group (dist.init_process_group) and assigns each rank its GPU device. On __exit__ it tears down the process group gracefully.

pipeline.run() blocks until every stage signals completion. Upstream ranks finish when their sampler is exhausted. Downstream ranks finish after the upstream rank is done and all fixed-duration MD work has drained.

def main() -> None:
    """Launch two parallel FIRE -> Langevin pipelines on 4 GPUs."""
    model = DemoModelWrapper(DemoModel())

    # Sinks (only used by ranks 1 and 3, but created on all for simplicity)
    sink_a = HostMemory(capacity=100)
    sink_b = HostMemory(capacity=100)

    # Dataset (only used by ranks 0 and 2)
    all_data = build_dataset()
    mid = len(all_data) // 2
    dataset_a = InMemoryDataset(all_data[:mid])
    dataset_b = InMemoryDataset(all_data[mid:])

    sampler_a = SizeAwareSampler(
        dataset=dataset_a,
        max_atoms=50,
        max_edges=0,
        max_batch_size=4,
    )
    sampler_b = SizeAwareSampler(
        dataset=dataset_b,
        max_atoms=50,
        max_edges=0,
        max_batch_size=4,
    )

    # Buffer config — matches sampler capacities and ensures fixed-size comm
    # buffers for NCCL compatibility (identical message count every step).
    buffer_cfg = BufferConfig(num_systems=4, num_nodes=50, num_edges=0)

    # Stages — one per rank.
    # By default prior_rank / next_rank are -1 (unset) and
    # DistributedPipeline.setup() would auto-wire a linear chain
    # 0 -> 1 -> 2 -> 3.  Setting them explicitly here creates two
    # independent sub-pipelines: 0 -> 1 and 2 -> 3.
    stages = {
        0: make_fire(
            model,
            rank=0,
            sampler=sampler_a,
            refill_frequency=1,
            prior_rank=None,
            next_rank=1,
            buffer_config=buffer_cfg,
        ),
        1: make_langevin(
            model,
            sink=sink_a,
            rank=1,
            prior_rank=0,
            next_rank=None,
            buffer_config=buffer_cfg,
        ),
        2: make_fire(
            model,
            rank=2,
            sampler=sampler_b,
            refill_frequency=1,
            prior_rank=None,
            next_rank=3,
            buffer_config=buffer_cfg,
        ),
        3: make_langevin(
            model,
            sink=sink_b,
            rank=3,
            prior_rank=2,
            next_rank=None,
            buffer_config=buffer_cfg,
        ),
    }

    if not _RUN_DISTRIBUTED_EXAMPLE:
        logger.info(
            "Not running under torchrun — skipping pipeline launch. "
            "Run with: torchrun --nproc_per_node=4 "
            "examples/distributed/01_distributed_pipeline.py",
        )
        return

    backend = "nccl"  # use gloo when testing with CPUs
    # debug mode will provide insight into what rank is doing what
    pipeline = DistributedPipeline(stages=stages, backend=backend, debug_mode=True)
    with pipeline:
        pipeline.run()
        rank = dist.get_rank()
        if rank == 1:
            expected_frames = len(dataset_a) * 20
            assert len(sink_a) == expected_frames, (
                f"rank 1 collected {len(sink_a)} of {expected_frames} frames"
            )
            logger.info(f"Rank 1 sink collected {len(sink_a)} trajectory frames")
        elif rank == 3:
            expected_frames = len(dataset_b) * 20
            assert len(sink_b) == expected_frames, (
                f"rank 3 collected {len(sink_b)} of {expected_frames} frames"
            )
            logger.info(f"Rank 3 sink collected {len(sink_b)} trajectory frames")


if __name__ == "__main__":
    main()

Total running time of the script: (0 minutes 0.005 seconds)

Gallery generated by Sphinx-Gallery