Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions nvalchemi/distributed/_core/particle_halo.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,13 @@ def _identify_ghosts_split(
# fractional coordinates that miss PBC halo atoms at cell boundaries and
# under-count neighbors on per-rank neighbor lists.
frac_pos = positions @ inv_cell

# Use canonical fractional coordinates only when deciding which receiver
# needs each atom. Keep the original Cartesian positions for transmission;
# the receiver's neighbor list applies the required integer cell image.
pbc = partitioner.pbc.to(device=positions.device)
frac_pos = torch.where(pbc.unsqueeze(0), frac_pos - torch.floor(frac_pos), frac_pos)

gw_frac = _ghost_width_fractional(partitioner, config.ghost_width).to(
device=positions.device
)
Expand Down
23 changes: 20 additions & 3 deletions test/distributed/_core/test_particle_halo.py
Original file line number Diff line number Diff line change
Expand Up @@ -192,25 +192,42 @@ def test_sender_owned_atom_inside_receiver_core_is_ghost(self):
assert direct_mask.tolist() == [True]
assert pbc_list == []

def test_sender_owned_atom_wrapped_into_receiver_core_is_pbc_ghost(self):
def test_sender_owned_atom_wrapped_into_receiver_core_is_direct_ghost(self):
config = _make_halo_config(
box_length=30.0,
ghost_width=3.0,
world_size=2,
rank=1,
)
# Rank 1 retains an atom just beyond the periodic z boundary. Rank 0
# needs the -cell image at z=0.1 as a ghost until ownership migrates.
# needs it as a ghost until ownership migrates. Halo routing uses the
# canonical z=0.1 position to select rank 0 but transmits raw z=30.1;
# the receiver's neighbor list supplies the -cell image.
positions = torch.tensor([[15.0, 15.0, 30.1]])
direct_mask, pbc_list = _identify_ghosts_split(positions, 0, config)

assert direct_mask.tolist() == [True]
assert pbc_list == []

def test_primary_cell_atom_uses_pbc_image_for_periodic_neighbor(self):
config = _make_halo_config(
box_length=30.0,
ghost_width=3.0,
world_size=2,
rank=0,
)
# This atom is in rank 0's primary-cell core but lies in rank 1's
# periodic halo. Rank 1 therefore receives the +cell image at z=30.1.
positions = torch.tensor([[15.0, 15.0, 0.1]])
direct_mask, pbc_list = _identify_ghosts_split(positions, 1, config)

assert direct_mask.tolist() == [False]
assert len(pbc_list) == 1
pbc_mask, shift = pbc_list[0]
assert pbc_mask.tolist() == [True]
torch.testing.assert_close(
positions[pbc_mask] + shift,
torch.tensor([[15.0, 15.0, 0.1]]),
torch.tensor([[15.0, 15.0, 30.1]]),
)

def test_center_atom_not_ghost(self):
Expand Down
166 changes: 166 additions & 0 deletions test/distributed/test_deferred_migration_halo.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,11 @@

from typing import Any

import pytest
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from _dd_harness import nccl_worker
from _gloo_harness import run_gloo
from torch.distributed import DeviceMesh

Expand Down Expand Up @@ -247,6 +250,139 @@ def _batch(positions: torch.Tensor) -> Batch:
)


def _unwrapped_periodic_migrant_force_matches_reference(
rank: int,
world_size: int,
_queue: Any = None,
device_type: str = "cpu",
) -> None:
"""An unwrapped migrant must remain visible across an internal boundary."""
assert world_size == 4
dtype = torch.float64
device = (
torch.device(f"cuda:{rank}") if device_type == "cuda" else torch.device("cpu")
)
box = 20.0
cell = torch.eye(3, dtype=dtype, device=device) * box
pbc = torch.ones(3, dtype=torch.bool, device=device)
mesh = DeviceMesh(
device_type,
list(range(world_size)),
mesh_dim_names=("domain",),
)
domain_config = DomainConfig(
cutoff=1.0,
skin=0.5,
mesh=mesh,
grid_dims=(1, 1, 4),
)
partitioner = SpatialPartitioner(
config=domain_config,
cell_matrix=cell.unsqueeze(0),
pbc=pbc.unsqueeze(0),
)
assert partitioner.rank_grid == (1, 1, 4)

# The migrant previously crossed the global periodic boundary. Its raw
# coordinate therefore carries a +20 Å box offset, while its physical
# position is obtained modulo the cell.
positions_before_crossing = torch.tensor(
[
[10.0, 10.0, 24.9], # physical z=4.9, owned by rank 0
[10.0, 10.0, 5.6], # rank-1 interaction partner
[10.0, 10.0, 12.5], # rank-2 non-interacting atom
[10.0, 10.0, 17.5], # rank-3 non-interacting atom
],
dtype=dtype,
device=device,
)
positions_at_compute = positions_before_crossing.clone()
positions_at_compute[0, 2] = 25.1 # physical z=5.1
reference_positions = positions_at_compute.clone()
reference_positions[0, 2] = 5.1

assert torch.equal(
partitioner.assign_atoms_to_ranks(positions_before_crossing),
torch.tensor([0, 1, 2, 3], device=device),
)

def _batch(positions: torch.Tensor) -> Batch:
n_atoms = positions.shape[0]
return Batch.from_data_list(
[
AtomicData(
positions=positions.clone(),
atomic_numbers=torch.full(
(n_atoms,), 18, dtype=torch.long, device=device
),
atomic_masses=torch.full(
(n_atoms,), 39.948, dtype=dtype, device=device
),
cell=cell.unsqueeze(0),
pbc=pbc.unsqueeze(0),
)
]
)

sharded = ShardedBatch.from_batch(
_batch(positions_before_crossing) if rank == 0 else None,
mesh=mesh,
config=domain_config,
src=0,
)
assert sharded.n_owned == 1
if rank == 0:
owned_positions = sharded.positions.to_local()
owned_positions[0].copy_(positions_at_compute[0])
assert partitioner.assign_atoms_to_ranks(owned_positions).item() == 1
assert partitioner.keeps_owner(
owned_positions,
owner_rank=0,
hysteresis=domain_config.effective_migration_hysteresis(),
).item()

# Compare against the same physical geometry with the migrant represented
# by its canonical z=5.1 image.
reference_forces = torch.zeros(4, 3, dtype=dtype, device=device)
if rank == 0:
reference_model = LennardJonesModelWrapper(
epsilon=1.0,
sigma=0.5,
cutoff=1.0,
).to(device)
reference_batch = _batch(reference_positions)
compute_neighbors(
reference_batch,
config=reference_model.model_config.neighbor_config,
)
reference_forces.copy_(reference_model(reference_batch)["forces"].detach())
dist.broadcast(reference_forces, src=0)
assert torch.linalg.vector_norm(reference_forces[1]).item() > 1.0

distributed_model = LennardJonesModelWrapper(
epsilon=1.0,
sigma=0.5,
cutoff=1.0,
).to(device)
with DistributedModel(distributed_model, domain_config) as model:
distributed_forces = model(sharded)["forces"]

expected_force = reference_forces[rank]
torch.testing.assert_close(
distributed_forces[0],
expected_force,
rtol=1e-12,
atol=1e-12,
msg=(
f"rank {rank} force is wrong at the periodic migrant geometry; "
"rank 1 must receive the rank-0-owned raw z=25.1 atom as its "
"physical z=5.1 ghost; "
f"distributed={distributed_forces[0].tolist()}, "
f"same_geometry_reference={expected_force.tolist()}"
),
)


def test_crossed_atom_reaches_receiver_before_deferred_migration() -> None:
"""The current owner must ghost an atom that entered a neighbor's core."""
run_gloo(world_size=2, fn=_crossed_atom_reaches_receiver_before_migration)
Expand All @@ -255,3 +391,33 @@ def test_crossed_atom_reaches_receiver_before_deferred_migration() -> None:
def test_crossed_atom_force_matches_same_geometry_reference() -> None:
"""Distributed force must match a reference at the identical geometry."""
run_gloo(world_size=2, fn=_crossed_atom_force_matches_same_geometry_reference)


def test_unwrapped_periodic_migrant_force_matches_same_geometry_reference() -> None:
"""A multi-box raw coordinate must not disappear from an internal halo."""
run_gloo(
world_size=4,
fn=_unwrapped_periodic_migrant_force_matches_reference,
)


@pytest.mark.multigpu
@pytest.mark.skipif(
torch.cuda.device_count() < 4,
reason="requires >=4 CUDA GPUs",
)
def test_unwrapped_periodic_migrant_force_matches_reference_nccl(
unused_tcp_port: int,
) -> None:
"""The periodic-migrant force regression also fails over a real NCCL halo."""
mp.spawn(
nccl_worker,
args=(
4,
str(unused_tcp_port),
_unwrapped_periodic_migrant_force_matches_reference,
None,
"cuda",
),
nprocs=4,
)
Loading