(process_id=0)
| 108 | |
| 109 | |
| 110 | def 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 | |
| 124 | if __name__ == "__main__": |
no test coverage detected