From 9c90fd66a002f67eb387ddbe4ec29eb0bcb8004c Mon Sep 17 00:00:00 2001 From: Akanksha Gupta Date: Mon, 24 Aug 2026 13:22:44 -0700 Subject: [PATCH] Update run_connect_example.py to demonstrate a TPU workload using JAX sharding PiperOrigin-RevId: 970039011 --- .../run_connect_example.py | 36 ++++- .../run_connect_example_test.py | 135 ------------------ 2 files changed, 32 insertions(+), 139 deletions(-) delete mode 100644 pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py diff --git a/pathwaysutils/experimental/shared_pathways_service/run_connect_example.py b/pathwaysutils/experimental/shared_pathways_service/run_connect_example.py index f14df73..d0b0d8a 100644 --- a/pathwaysutils/experimental/shared_pathways_service/run_connect_example.py +++ b/pathwaysutils/experimental/shared_pathways_service/run_connect_example.py @@ -5,7 +5,10 @@ from absl import app from absl import flags +import jax import jax.numpy as jnp +import jax.sharding as jsharding +import pathwaysutils from pathwaysutils.experimental.shared_pathways_service import isc_pathways @@ -75,11 +78,36 @@ def main(argv: Sequence[str]) -> None: proxy_options=FLAGS.proxy_options, collect_service_metrics=FLAGS.collect_service_metrics, ): - orig_matrix = jnp.zeros(5) - result_matrix = orig_matrix + 1 - print("Original Random Matrix:") + # your-workload + tpu_devices = jax.devices() + for device in tpu_devices: + print("Device: %s, Kind: %s", device, device.device_kind) + if "tpu" not in device.device_kind.lower(): + print("Error! TPUs not found") + exit() + if not pathwaysutils.is_pathways_backend_used(): + print("Error! TPUs not found") + exit() + print( + "All devices are confirmed to be TPUs. TPU devices found:" + f" {tpu_devices}" + ) + num_devices = len(tpu_devices) + mesh = jsharding.Mesh(tpu_devices, axis_names=("data",)) + sharding = jsharding.NamedSharding(mesh, jsharding.PartitionSpec("data")) + + @jax.jit + def tpu_add_one(x): + return x + 1 + + orig_matrix = jnp.zeros(num_devices) + orig_matrix_tpu = jax.device_put(orig_matrix, sharding) + result_matrix_tpu = tpu_add_one(orig_matrix_tpu) + result_matrix = jax.device_get(result_matrix_tpu) + + print(f"Original Matrix (on all {num_devices} devices):") pprint.pprint(orig_matrix) - print("\nMatrix after adding 1:") + print("\nResult Matrix after parallel addition:") pprint.pprint(result_matrix) diff --git a/pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py b/pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py deleted file mode 100644 index 8912838..0000000 --- a/pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py +++ /dev/null @@ -1,135 +0,0 @@ -from unittest import mock - -from absl.testing import absltest -from absl.testing import flagsaver -import numpy as np - - -class RunConnectExampleTest(absltest.TestCase): - """Tests the logic from the run_connect_example.py script.""" - - @flagsaver.flagsaver( - cluster="random-cluster-name", - project="random-project-id", - region="random-region", - gcs_bucket="random-bucket", - pathways_service="random-pathways-service:1234", - tpu_type="tpuv6e:2x2", - tpu_count=2, - ) - def test_run_connect_example_main(self): - """Tests that the main function calls connect and executes the logic.""" - # Import inside the test to avoid flag parsing errors on module load. - from pathwaysutils.experimental.shared_pathways_service import run_connect_example - - mock_connect = self.enter_context( - mock.patch.object( - run_connect_example.isc_pathways, "connect", autospec=True - ) - ) - mock_pprint = self.enter_context(mock.patch("pprint.pprint", autospec=True)) - - run_connect_example.main(["unused_argv"]) - - mock_connect.assert_called_once_with( - cluster="random-cluster-name", - project="random-project-id", - region="random-region", - gcs_bucket="random-bucket", - pathways_service="random-pathways-service:1234", - expected_tpu_instances={"tpuv6e:2x2": 2}, - proxy_job_name=None, - proxy_server_image=( - "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server:latest" - ), - proxy_options=None, - collect_service_metrics=False, - ) - self.assertEqual(mock_pprint.call_count, 2) - np.testing.assert_array_equal( - mock_pprint.call_args_list[0].args[0], np.zeros(5) - ) - np.testing.assert_array_equal( - mock_pprint.call_args_list[1].args[0], np.ones(5) - ) - - @flagsaver.flagsaver( - cluster="random-cluster-name", - project="random-project-id", - region="random-region", - gcs_bucket="random-bucket", - pathways_service="random-pathways-service:1234", - tpu_type="tpuv6e:2x2", - tpu_count=2, - proxy_job_name="test-job-name", - proxy_server_image="test-image", - proxy_options=["use_insecure_credentials:true"], - ) - def test_run_connect_example_main_with_optional_flags(self): - """Tests that main passes optional flags to connect.""" - # Import inside the test to avoid flag parsing errors on module load. - from pathwaysutils.experimental.shared_pathways_service import run_connect_example - - mock_connect = self.enter_context( - mock.patch.object( - run_connect_example.isc_pathways, "connect", autospec=True - ) - ) - self.enter_context(mock.patch("pprint.pprint", autospec=True)) - - run_connect_example.main(["unused_argv"]) - - mock_connect.assert_called_once_with( - cluster="random-cluster-name", - project="random-project-id", - region="random-region", - gcs_bucket="random-bucket", - pathways_service="random-pathways-service:1234", - expected_tpu_instances={"tpuv6e:2x2": 2}, - proxy_job_name="test-job-name", - proxy_server_image="test-image", - proxy_options=["use_insecure_credentials:true"], - collect_service_metrics=False, - ) - - @flagsaver.flagsaver( - cluster="random-cluster-name", - project="random-project-id", - region="random-region", - gcs_bucket="random-bucket", - pathways_service="random-pathways-service:1234", - tpu_type="tpuv6e:2x2", - tpu_count=2, - collect_service_metrics=True, - ) - def test_run_connect_example_main_with_metrics_enabled(self): - """Tests that main passes collect_service_metrics flag to connect.""" - from pathwaysutils.experimental.shared_pathways_service import run_connect_example - - mock_connect = self.enter_context( - mock.patch.object( - run_connect_example.isc_pathways, "connect", autospec=True - ) - ) - self.enter_context(mock.patch("pprint.pprint", autospec=True)) - - run_connect_example.main(["unused_argv"]) - - mock_connect.assert_called_once_with( - cluster="random-cluster-name", - project="random-project-id", - region="random-region", - gcs_bucket="random-bucket", - pathways_service="random-pathways-service:1234", - expected_tpu_instances={"tpuv6e:2x2": 2}, - proxy_job_name=None, - proxy_server_image=( - "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server:latest" - ), - proxy_options=None, - collect_service_metrics=True, - ) - - -if __name__ == "__main__": - absltest.main()