PMIx based wire-up for PyTorch
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354import 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_value_0", "1") if rank == 1: store.set("test_value_1", "2") store.wait(["test_value_0", "test_value_1"])
val1 = store.get("test_value_0") val2 = store.get("test_value_1") assert val1.decode("utf-8") == "1" assert val2.decode("utf-8") == "2"
print("RIXA_CPU_WORKER_CLEAN_EXIT")
sys.exit(0)
@pytest.mark.cpudef test_pmix_set_get_wait(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)