MCPcopy Create free account
hub / github.com/NVIDIA/cuda-quantum / testMPI

Function testMPI

python/tests/parallel/test_mpi_api.py:18–55  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

16
17@skipIfUnsupported
18def 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

Callers

nothing calls this directly

Calls 10

enumerateFunction · 0.85
rankMethod · 0.80
num_ranksMethod · 0.80
all_gatherMethod · 0.80
broadcastMethod · 0.80
printFunction · 0.50
initializeMethod · 0.45
is_initializedMethod · 0.45
getMethod · 0.45
finalizeMethod · 0.45

Tested by

no test coverage detected