()
| 16 | |
| 17 | @skipIfUnsupported |
| 18 | def testMPI(): |
| 19 | cudaq.mpi.initialize() |
| 20 | assert cudaq.mpi.is_initialized() == True |
| 21 | # Check rank API |
| 22 | if os.environ.get('OMPI_COMM_WORLD_RANK') is not None: |
| 23 | print("Rank:", os.environ.get('OMPI_COMM_WORLD_RANK')) |
| 24 | assert cudaq.mpi.rank() == int(os.environ.get('OMPI_COMM_WORLD_RANK')) |
| 25 | |
| 26 | if os.environ.get('OMPI_COMM_WORLD_SIZE') is not None: |
| 27 | assert cudaq.mpi.num_ranks() == int( |
| 28 | os.environ.get('OMPI_COMM_WORLD_SIZE')) |
| 29 | |
| 30 | # all_gather integers |
| 31 | localData = [cudaq.mpi.rank()] |
| 32 | gatherData = cudaq.mpi.all_gather(cudaq.mpi.num_ranks(), localData) |
| 33 | assert len(gatherData) == cudaq.mpi.num_ranks() |
| 34 | for idx, x in enumerate(gatherData): |
| 35 | assert x == idx |
| 36 | |
| 37 | # all_gather floats |
| 38 | localData = [float(cudaq.mpi.rank())] |
| 39 | gatherData = cudaq.mpi.all_gather(cudaq.mpi.num_ranks(), localData) |
| 40 | assert len(gatherData) == cudaq.mpi.num_ranks() |
| 41 | for idx, x in enumerate(gatherData): |
| 42 | assert abs(gatherData[idx] - float(idx)) < 1e-12 |
| 43 | |
| 44 | # Broadcast |
| 45 | ref_data = [1.0, 2.0, 3.0] |
| 46 | if cudaq.mpi.rank() == 0: |
| 47 | data = ref_data |
| 48 | else: |
| 49 | data = [] |
| 50 | |
| 51 | data = cudaq.mpi.broadcast(data, len(ref_data), 0) |
| 52 | for idx, x in enumerate(data): |
| 53 | assert abs(x - ref_data[idx]) < 1e-12 |
| 54 | |
| 55 | cudaq.mpi.finalize() |
| 56 | |
| 57 | |
| 58 | # leave for gdb debugging |
nothing calls this directly
no test coverage detected