PMIx based wire-up for PyTorch
Something went wrong. Try again.
12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849import osimport subprocessimport sys
import pytest
from rixa.pytorch import get_pmix_store
def is_worker(): return os.environ.get("RIXA_WORKER_MODE") == "1"
if is_worker(): store = get_pmix_store()() rank = store.rank() if rank == 0: store.set("test_values", "1") store.wait(["test_values"])
assert store.check(["test_values"]) assert not store.check(["test_values_not"])
print("RIXA_CPU_WORKER_CLEAN_EXIT") sys.exit(0)
@pytest.mark.cpudef test_pmix_check(nprocs): env = os.environ.copy() env["RIXA_WORKER_MODE"] = "1"
cmd = [ "mpirun", "-n", nprocs, "-x", "RIXA_WORKER_MODE", "-x", "PYTHONPATH", sys.executable, __file__, ]
result = subprocess.run(cmd, capture_output=True, text=True, env=env)
assert result.returncode == 0 assert result.stdout.count("RIXA_CPU_WORKER_CLEAN_EXIT") == int(nprocs)