MCPcopy Create free account
hub / github.com/NVIDIA/DALI / run_multiprocess_workflow

Function run_multiprocess_workflow

dali/test/python/jax_plugin/jax_server.py:110–121  ·  view source on GitHub ↗
(process_id=0)

Source from the content-addressed store, hash-verified

108
109
110def run_multiprocess_workflow(process_id=0):
111 jax.distributed.initialize(
112 coordinator_address="localhost:12321", num_processes=2, process_id=process_id
113 )
114
115 log.basicConfig(level=log.INFO, format=f"PID {process_id}: %(message)s")
116
117 print_devices(process_id=process_id)
118
119 test_lax_workflow(process_id=process_id)
120 test_positional_sharding_workflow(process_id=process_id)
121 test_named_sharding_workflow(process_id=process_id)
122
123
124if __name__ == "__main__":

Callers 2

jax_client.pyFile · 0.90
jax_server.pyFile · 0.70

Calls 5

test_lax_workflowFunction · 0.85
print_devicesFunction · 0.70
initializeMethod · 0.45

Tested by

no test coverage detected