From ba1985f0cacdab84cb16bad95d39bd8976dff569 Mon Sep 17 00:00:00 2001 From: Googler Date: Sun, 30 Aug 2026 22:03:40 -0700 Subject: [PATCH] Expand PyTorch WeightSynchronizer test suite to reach parity with JAX: - Added `WeightSynchronizer::BindWeights` in `framework/` and nanobind module, enabling live weight memory re-binding without daemon teardown. - Add `test_bind_weights` in `weight_synchronizer_test.py` - Add test cases for variable layer dimensions. PiperOrigin-RevId: 973655927 --- tpu_sync/api/torch/BUILD | 2 - tpu_sync/api/torch/kv_cache_manager_test.py | 326 ++++++++++------ .../torch/kv_cache_manager_transfer_test.py | 350 ++++++++++++------ tpu_sync/api/torch/weight_synchronizer.py | 4 + .../api/torch/weight_synchronizer_test.py | 265 ++++++++++--- .../torch/tpu_raiden_torch_module.cc | 12 + .../frameworks/torch/weight_synchronizer.cc | 12 +- .../frameworks/torch/weight_synchronizer.h | 4 + 8 files changed, 676 insertions(+), 299 deletions(-) diff --git a/tpu_sync/api/torch/BUILD b/tpu_sync/api/torch/BUILD index 9fba8527..916eddf2 100644 --- a/tpu_sync/api/torch/BUILD +++ b/tpu_sync/api/torch/BUILD @@ -159,7 +159,6 @@ py_test( ":weight_synchronizer_torch_py", "@com_google_absl_py//absl/testing:absltest", "@com_google_absl_py//absl/testing:parameterized", - "@pypi//numpy", "@torch_tpu//shims/torch:pytorch", "@torch_tpu//torch_tpu", ], @@ -201,7 +200,6 @@ py_test( ":kv_cache_manager_torch_py", "@com_google_absl_py//absl/testing:absltest", "@com_google_absl_py//absl/testing:parameterized", - "@pypi//numpy", "@torch_tpu//shims/torch:pytorch", "@torch_tpu//torch_tpu", ], diff --git a/tpu_sync/api/torch/kv_cache_manager_test.py b/tpu_sync/api/torch/kv_cache_manager_test.py index aa259816..026038d5 100644 --- a/tpu_sync/api/torch/kv_cache_manager_test.py +++ b/tpu_sync/api/torch/kv_cache_manager_test.py @@ -15,10 +15,10 @@ """E2E physical unit tests for KVCacheManager on XLA TPUs.""" import time +from typing import Sequence from absl.testing import absltest from absl.testing import parameterized -import numpy as np import torch from tpu_sync.api.torch import kv_cache_manager @@ -32,65 +32,119 @@ def setUp(self): super().setUp() # Initialize PyTorch XLA accelerator device E2E self.device = torch.device("tpu") - self.num_layers = 1 - self.block_size = 1 - - def test_initialization(self): - device = self.device - shape = (4, 128, 8) - kv_caches = [torch.zeros(shape, device=device)] - - manager = KVCacheManager( - kv_caches=kv_caches, - node_id=0, - local_control_port=0, - max_blocks=4, - num_slots=2, + self.num_layers = 2 + self.skip_lock = True + + def _generate_random_cache( + self, + shape: Sequence[int], + dtype: torch.dtype = torch.float32, + seed: int = 123, + ): + torch.manual_seed(seed) + if dtype == torch.bfloat16: + host_data = torch.randn(shape, dtype=torch.bfloat16) + else: + host_data = torch.randn(shape, dtype=torch.float32) + dev_arr = host_data.to(self.device) + return dev_arr, host_data.cpu() + + def _wait_for_transfer( + self, + manager: kv_cache_manager.KVCacheManager, + req_id: str, + is_receiver: bool = True, + timeout_s: float = 10.0, + sleep_sec: float = 0.05, + ): + """Polls KVCacheManager until transfer finishes or times out.""" + deadline = time.time() + timeout_s + while time.time() < deadline: + done_sending, done_recving, failed_recving = manager.poll_stats() + if is_receiver: + if req_id in failed_recving: + self.fail(f"Receiver transfer failed for request {req_id}") + if req_id in done_recving: + return + else: + if req_id in done_sending: + return + time.sleep(sleep_sec) + role = "Receiver" if is_receiver else "Producer" + self.fail( + f"{role} did not finish transfer within {timeout_s}s for request" + f" {req_id}" ) - self.assertIsNotNone(manager) - def test_e2e_transfer_polling(self): - num_blocks = 2 - shape = (num_blocks, 128, 8) - - src_caches = [] - for _ in range(self.num_layers): - src_caches.append( - torch.full( - shape, fill_value=1.0, dtype=torch.float32, device=self.device - ) + def _setup_test_pair( + self, + num_blocks: int, + dtype: torch.dtype, + seed: int = 123, + ) -> tuple[ + kv_cache_manager.KVCacheManager, + kv_cache_manager.KVCacheManager, + list[torch.Tensor], + list[torch.Tensor], + ]: + """Creates a Producer-Consumer pair with randomized source and zeroed destination caches.""" + shape = (num_blocks, 128, 8, 8, 128) + src_caches, src_refs = [], [] + for i in range(self.num_layers): + dev_arr, host_ref = self._generate_random_cache( + shape, dtype=dtype, seed=seed + i ) + src_caches.append(dev_arr) + src_refs.append(host_ref) - dst_caches = [] - for _ in range(self.num_layers): - dst_caches.append( - torch.zeros(shape, dtype=torch.float32, device=self.device) - ) + dst_caches = [ + torch.zeros(shape, dtype=dtype, device=self.device) + for _ in range(self.num_layers) + ] producer = KVCacheManager( kv_caches=src_caches, node_id=0, local_control_port=0, - max_blocks=2, + max_blocks=num_blocks, num_slots=2, + unsafe_skip_buffer_lock=self.skip_lock, ) - consumer = KVCacheManager( kv_caches=dst_caches, node_id=0, local_control_port=0, - max_blocks=2, + max_blocks=num_blocks, + num_slots=2, + unsafe_skip_buffer_lock=self.skip_lock, + ) + self.assertGreater(producer.local_control_port, 0) + return producer, consumer, src_refs, dst_caches + + def test_initialization(self): + shape = (4, 128, 8, 8, 128) + kv_caches = [torch.zeros(shape, device=self.device)] + + manager = KVCacheManager( + kv_caches=kv_caches, + node_id=0, + local_control_port=0, + max_blocks=4, num_slots=2, ) + self.assertIsNotNone(manager) - port = producer.local_control_port - self.assertGreater(port, 0) + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_e2e_transfer_polling(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=2, dtype=dtype, seed=100 + ) req_id = "test_req_poll" uuid = 12345 producer.register_read(req_id, uuid, [0, 1]) - remote_endpoint = f"127.0.0.1:{port}" + remote_endpoint = f"127.0.0.1:{producer.local_control_port}" consumer.start_read( req_id=req_id, uuid=uuid, @@ -99,76 +153,132 @@ def test_e2e_transfer_polling(self): local_block_ids=[0, 1], ) - # Poll until consumer is done receiving - done = False - for _ in range(50): - _, done_recving, failed_recving = consumer.poll_stats() - if req_id in failed_recving: - self.fail("Transfer failed") - if req_id in done_recving: - done = True - break - time.sleep(0.1) + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + # Check that consumer correctly loaded all layer values + for idx, t in enumerate(dst_caches): + self.assertTrue(torch.equal(t.cpu(), src_refs[idx])) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) + + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_non_contiguous_blocks(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=3, dtype=dtype, seed=200 + ) + + req_id = "test_req_non_contig" + uuid = 54321 + producer.register_read(req_id, uuid, [0, 2]) + + remote_endpoint = f"127.0.0.1:{producer.local_control_port}" + consumer.start_read( + req_id=req_id, + uuid=uuid, + remote_endpoint=remote_endpoint, + remote_block_ids=[0, 2], + local_block_ids=[0, 1], + ) - self.assertTrue(done, "Consumer did not finish transfer in time") + self._wait_for_transfer(consumer, req_id, is_receiver=True) - # Check that consumer correctly loaded the values - for t in dst_caches: - np.testing.assert_allclose(t.cpu().numpy(), 1.0, atol=1e-5) + for idx, t in enumerate(dst_caches): + # local block 0 <- remote block 0 + self.assertTrue(torch.equal(t[0].cpu(), src_refs[idx][0])) + # local block 1 <- remote block 2 + self.assertTrue(torch.equal(t[1].cpu(), src_refs[idx][2])) + # local block 2 was not copied, should remain 0 + self.assertTrue(torch.equal(t[2].cpu(), torch.zeros_like(t[2].cpu()))) # Poll producer until it's done sending - done_prod = False - for _ in range(50): - done_sending, _, _ = producer.poll_stats() - if req_id in done_sending: - done_prod = True - break - time.sleep(0.1) - - self.assertTrue(done_prod, "Producer did not finish sending in time") - - def test_parallel_pull(self): - num_blocks = 2 - shape = (num_blocks, 128, 8) - - src_caches = [] - for _ in range(self.num_layers): - src_caches.append( - torch.full( - shape, fill_value=2.0, dtype=torch.float32, device=self.device - ) - ) + self._wait_for_transfer(producer, req_id, is_receiver=False) - dst_caches = [] - for _ in range(self.num_layers): - dst_caches.append( - torch.zeros(shape, dtype=torch.float32, device=self.device) - ) + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_host_reordering(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=2, dtype=dtype, seed=300 + ) - producer = KVCacheManager( - kv_caches=src_caches, - node_id=0, - local_control_port=0, - max_blocks=2, - num_slots=2, + req_id = "test_req_reorder" + uuid = 98765 + producer.register_read(req_id, uuid, [0, 1]) + + remote_endpoint = f"127.0.0.1:{producer.local_control_port}" + consumer.start_read( + req_id=req_id, + uuid=uuid, + remote_endpoint=remote_endpoint, + remote_block_ids=[1, 0], + local_block_ids=[0, 1], ) - consumer = KVCacheManager( - kv_caches=dst_caches, - node_id=0, - local_control_port=0, - max_blocks=2, - num_slots=2, + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + for idx, t in enumerate(dst_caches): + # local block 0 <- remote block 1 + self.assertTrue(torch.equal(t[0].cpu(), src_refs[idx][1])) + # local block 1 <- remote block 0 + self.assertTrue(torch.equal(t[1].cpu(), src_refs[idx][0])) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) + + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_large_complex_non_contiguous_and_reorder(self, dtype): + num_blocks = 16 + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=num_blocks, dtype=dtype, seed=400 + ) + + req_id = "test_req_large_complex" + uuid = 13579 + + remote_blocks = [0, 2, 3, 5, 6, 7, 9, 11, 12, 14] + requested_remote = list(reversed(remote_blocks)) + local_blocks = list(range(len(remote_blocks))) + + producer.register_read(req_id, uuid, remote_blocks) + + remote_endpoint = f"127.0.0.1:{producer.local_control_port}" + consumer.start_read( + req_id=req_id, + uuid=uuid, + remote_endpoint=remote_endpoint, + remote_block_ids=requested_remote, + local_block_ids=local_blocks, ) - port = producer.local_control_port - self.assertGreater(port, 0) + self._wait_for_transfer(consumer, req_id, is_receiver=True, timeout_s=15.0) + + for idx, t in enumerate(dst_caches): + for local_idx, local_block in enumerate(local_blocks): + remote_block = requested_remote[local_idx] + self.assertTrue( + torch.equal(t[local_block].cpu(), src_refs[idx][remote_block]) + ) + + for local_block in range(len(local_blocks), num_blocks): + self.assertTrue( + torch.equal( + t[local_block].cpu(), torch.zeros_like(t[local_block].cpu()) + ) + ) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) + + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_parallel_pull(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=2, dtype=dtype, seed=500 + ) req_id = "test_req_parallel" uuid = 99999 producer.register_read(req_id, uuid, [0, 1]) - remote_endpoint = f"127.0.0.1:{port}" + remote_endpoint = f"127.0.0.1:{producer.local_control_port}" consumer.start_read( req_id=req_id, uuid=uuid, @@ -178,31 +288,13 @@ def test_parallel_pull(self): parallelism=2, ) - done = False - for _ in range(50): - _, done_recving, failed_recving = consumer.poll_stats() - if req_id in failed_recving: - self.fail("Transfer failed") - if req_id in done_recving: - done = True - break - time.sleep(0.1) - - self.assertTrue(done, "Consumer did not finish transfer in time") - time.sleep(0.5) - - for t in dst_caches: - np.testing.assert_allclose(t.cpu().numpy(), 2.0, atol=1e-5) - - done_prod = False - for _ in range(50): - done_sending, _, _ = producer.poll_stats() - if req_id in done_sending: - done_prod = True - break - time.sleep(0.1) - - self.assertTrue(done_prod, "Producer did not finish sending in time") + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + for idx, t in enumerate(dst_caches): + self.assertTrue(torch.equal(t.cpu(), src_refs[idx])) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) if __name__ == "__main__": diff --git a/tpu_sync/api/torch/kv_cache_manager_transfer_test.py b/tpu_sync/api/torch/kv_cache_manager_transfer_test.py index 651927df..c10ef7c0 100644 --- a/tpu_sync/api/torch/kv_cache_manager_transfer_test.py +++ b/tpu_sync/api/torch/kv_cache_manager_transfer_test.py @@ -14,16 +14,16 @@ """E2E physical unit tests for KVCacheManager transfer on XLA TPUs.""" -import threading import time +from typing import Sequence from absl.testing import absltest from absl.testing import parameterized -import numpy as np import torch -from tpu_sync.api.torch.kv_cache_manager import KVCacheManager +from tpu_sync.api.torch import kv_cache_manager +KVCacheManager = kv_cache_manager.KVCacheManager class KVCacheManagerTransferTest(parameterized.TestCase): @@ -31,176 +31,284 @@ def setUp(self): super().setUp() # Initialize PyTorch XLA accelerator device E2E self.device = torch.device("tpu") - self.num_layers = 1 - self.block_size = 1 - - def test_initialization(self): - shape = (4, 128, 8) - kv_caches = [torch.zeros(shape, device=self.device)] - - engine = KVCacheManager( - kv_caches=kv_caches, - node_id=0, - local_control_port=0, - max_blocks=4, - num_slots=2, + self.num_layers = 2 + self.skip_lock = True + + def _generate_random_cache( + self, + shape: Sequence[int], + dtype: torch.dtype = torch.float32, + seed: int = 123, + ): + torch.manual_seed(seed) + if dtype == torch.bfloat16: + host_data = torch.randn(shape, dtype=torch.bfloat16) + else: + host_data = torch.randn(shape, dtype=torch.float32) + dev_arr = host_data.to(self.device) + return dev_arr, host_data.cpu() + + def _wait_for_transfer( + self, + manager: kv_cache_manager.KVCacheManager, + req_id: str, + is_receiver: bool = True, + timeout_s: float = 10.0, + sleep_sec: float = 0.05, + ): + """Polls KVCacheManager until transfer finishes or times out.""" + deadline = time.time() + timeout_s + while time.time() < deadline: + done_sending, done_recving, failed_recving = manager.poll_stats() + if is_receiver: + if req_id in failed_recving: + self.fail(f"Receiver transfer failed for request {req_id}") + if req_id in done_recving: + return + else: + if req_id in done_sending: + return + time.sleep(sleep_sec) + role = "Receiver" if is_receiver else "Producer" + self.fail( + f"{role} did not finish transfer within {timeout_s}s for request" + f" {req_id}" ) - self.assertIsNotNone(engine) - - def test_e2e_transfer_polling(self): - num_blocks = 2 - shape = (num_blocks, 128, 8) - - src_caches = [] - for _ in range(self.num_layers): - src_caches.append( - torch.full( - shape, fill_value=1.0, dtype=torch.float32, device=self.device - ) - ) - dst_caches = [] - for _ in range(self.num_layers): - dst_caches.append( - torch.zeros(shape, dtype=torch.float32, device=self.device) + def _to_loopback_endpoints(self, endpoints): + res = [] + for ep in endpoints: + d = dict(ep) + endpoint_str = d["endpoint"] + port = endpoint_str.split(":")[-1] + d["endpoint"] = f"127.0.0.1:{port}" + res.append(d) + return res + + def _setup_test_pair( + self, + num_blocks: int, + dtype: torch.dtype, + seed: int = 123, + ) -> tuple[ + kv_cache_manager.KVCacheManager, + kv_cache_manager.KVCacheManager, + list[torch.Tensor], + list[torch.Tensor], + ]: + """Creates a Producer-Consumer pair with randomized source and zeroed destination caches.""" + shape = (num_blocks, 128, 8, 8, 128) + src_caches, src_refs = [], [] + for i in range(self.num_layers): + dev_arr, host_ref = self._generate_random_cache( + shape, dtype=dtype, seed=seed + i ) + src_caches.append(dev_arr) + src_refs.append(host_ref) + + dst_caches = [ + torch.zeros(shape, dtype=dtype, device=self.device) + for _ in range(self.num_layers) + ] producer = KVCacheManager( kv_caches=src_caches, node_id=0, local_control_port=0, - max_blocks=2, + max_blocks=num_blocks, num_slots=2, + unsafe_skip_buffer_lock=self.skip_lock, ) - consumer = KVCacheManager( kv_caches=dst_caches, node_id=0, local_control_port=0, - max_blocks=2, + max_blocks=num_blocks, + num_slots=2, + unsafe_skip_buffer_lock=self.skip_lock, + ) + self.assertGreater(producer.local_control_port, 0) + return producer, consumer, src_refs, dst_caches + + def test_initialization(self): + shape = (4, 128, 8, 8, 128) + kv_caches = [torch.zeros(shape, device=self.device)] + + manager = KVCacheManager( + kv_caches=kv_caches, + node_id=0, + local_control_port=0, + max_blocks=4, num_slots=2, ) + self.assertIsNotNone(manager) - # Use getattr just in case local_control_port was completely hidden - port = getattr(producer, "local_control_port", 0) - self.assertGreater(port, 0) + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_e2e_transfer_polling(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=2, dtype=dtype, seed=100 + ) req_id = "test_req_poll" uuid = 12345 producer.register_read(req_id, uuid, [0, 1]) - remote_endpoint = f"127.0.0.1:{port}" + endpoints = self._to_loopback_endpoints(producer.get_local_endpoints()) + self.assertNotEmpty(endpoints) consumer.start_read( req_id=req_id, uuid=uuid, - remote_endpoint=remote_endpoint, + remote_endpoint=endpoints, remote_block_ids=[0, 1], local_block_ids=[0, 1], ) - # Poll until consumer is done receiving - done = False - for _ in range(50): - done_sending, done_recving, failed_recving = consumer.poll_stats() - if req_id in failed_recving: - self.fail("Transfer failed") - if req_id in done_recving: - done = True - break - time.sleep(0.1) + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + # Check that consumer correctly loaded all layer values + for idx, t in enumerate(dst_caches): + self.assertTrue(torch.equal(t.cpu(), src_refs[idx])) - self.assertTrue(done, "Consumer did not finish transfer in time") + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) + + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_non_contiguous_blocks(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=3, dtype=dtype, seed=200 + ) - # Check that consumer correctly loaded the values - for t in dst_caches: - np.testing.assert_allclose(t.cpu().numpy(), 1.0, atol=1e-5) + req_id = "test_req_non_contig" + uuid = 54321 + producer.register_read(req_id, uuid, [0, 2]) + + endpoints = self._to_loopback_endpoints(producer.get_local_endpoints()) + self.assertNotEmpty(endpoints) + consumer.start_read( + req_id=req_id, + uuid=uuid, + remote_endpoint=endpoints, + remote_block_ids=[0, 2], + local_block_ids=[0, 1], + ) + + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + for idx, t in enumerate(dst_caches): + # local block 0 <- remote block 0 + self.assertTrue(torch.equal(t[0].cpu(), src_refs[idx][0])) + # local block 1 <- remote block 2 + self.assertTrue(torch.equal(t[1].cpu(), src_refs[idx][2])) + # local block 2 was not copied, should remain 0 + self.assertTrue(torch.equal(t[2].cpu(), torch.zeros_like(t[2].cpu()))) # Poll producer until it's done sending - done_prod = False - for _ in range(50): - done_sending, done_recving, failed_recving = producer.poll_stats() - if req_id in done_sending: - done_prod = True - break - time.sleep(0.1) - - self.assertTrue(done_prod, "Producer did not finish sending in time") - - def test_parallel_pull(self): - num_blocks = 2 - shape = (num_blocks, 128, 8) - - src_caches = [] - for _ in range(self.num_layers): - src_caches.append( - torch.full( - shape, fill_value=2.0, dtype=torch.float32, device=self.device - ) - ) + self._wait_for_transfer(producer, req_id, is_receiver=False) - dst_caches = [] - for _ in range(self.num_layers): - dst_caches.append( - torch.zeros(shape, dtype=torch.float32, device=self.device) - ) + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_host_reordering(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=2, dtype=dtype, seed=300 + ) - producer = KVCacheManager( - kv_caches=src_caches, - node_id=0, - local_control_port=0, - max_blocks=2, - num_slots=2, + req_id = "test_req_reorder" + uuid = 98765 + producer.register_read(req_id, uuid, [0, 1]) + + endpoints = self._to_loopback_endpoints(producer.get_local_endpoints()) + self.assertNotEmpty(endpoints) + consumer.start_read( + req_id=req_id, + uuid=uuid, + remote_endpoint=endpoints, + remote_block_ids=[1, 0], + local_block_ids=[0, 1], ) - consumer = KVCacheManager( - kv_caches=dst_caches, - node_id=0, - local_control_port=0, - max_blocks=2, - num_slots=2, + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + for idx, t in enumerate(dst_caches): + # local block 0 <- remote block 1 + self.assertTrue(torch.equal(t[0].cpu(), src_refs[idx][1])) + # local block 1 <- remote block 0 + self.assertTrue(torch.equal(t[1].cpu(), src_refs[idx][0])) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) + + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_large_complex_non_contiguous_and_reorder(self, dtype): + num_blocks = 16 + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=num_blocks, dtype=dtype, seed=400 + ) + + req_id = "test_req_large_complex" + uuid = 13579 + + remote_blocks = [0, 2, 3, 5, 6, 7, 9, 11, 12, 14] + requested_remote = list(reversed(remote_blocks)) + local_blocks = list(range(len(remote_blocks))) + + producer.register_read(req_id, uuid, remote_blocks) + + endpoints = self._to_loopback_endpoints(producer.get_local_endpoints()) + self.assertNotEmpty(endpoints) + consumer.start_read( + req_id=req_id, + uuid=uuid, + remote_endpoint=endpoints, + remote_block_ids=requested_remote, + local_block_ids=local_blocks, ) - port = getattr(producer, "local_control_port", 0) - self.assertGreater(port, 0) + self._wait_for_transfer(consumer, req_id, is_receiver=True, timeout_s=15.0) + + for idx, t in enumerate(dst_caches): + for local_idx, local_block in enumerate(local_blocks): + remote_block = requested_remote[local_idx] + self.assertTrue( + torch.equal(t[local_block].cpu(), src_refs[idx][remote_block]) + ) + + for local_block in range(len(local_blocks), num_blocks): + self.assertTrue( + torch.equal( + t[local_block].cpu(), torch.zeros_like(t[local_block].cpu()) + ) + ) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) + + @parameterized.parameters(torch.float32, torch.bfloat16) + def test_parallel_pull(self, dtype): + producer, consumer, src_refs, dst_caches = self._setup_test_pair( + num_blocks=2, dtype=dtype, seed=500 + ) req_id = "test_req_parallel" uuid = 99999 producer.register_read(req_id, uuid, [0, 1]) - remote_endpoint = f"127.0.0.1:{port}" + endpoints = self._to_loopback_endpoints(producer.get_local_endpoints()) + self.assertNotEmpty(endpoints) consumer.start_read( req_id=req_id, uuid=uuid, - remote_endpoint=remote_endpoint, + remote_endpoint=endpoints, remote_block_ids=[0, 1], local_block_ids=[0, 1], parallelism=2, ) - done = False - for _ in range(50): - done_sending, done_recving, failed_recving = consumer.poll_stats() - if req_id in failed_recving: - self.fail("Transfer failed") - if req_id in done_recving: - done = True - break - time.sleep(0.1) - - self.assertTrue(done, "Consumer did not finish transfer in time") - - for t in dst_caches: - np.testing.assert_allclose(t.cpu().numpy(), 2.0, atol=1e-5) - - done_prod = False - for _ in range(50): - done_sending, done_recving, failed_recving = producer.poll_stats() - if req_id in done_sending: - done_prod = True - break - time.sleep(0.1) - - self.assertTrue(done_prod, "Producer did not finish sending in time") + self._wait_for_transfer(consumer, req_id, is_receiver=True) + + for idx, t in enumerate(dst_caches): + self.assertTrue(torch.equal(t.cpu(), src_refs[idx])) + + # Poll producer until it's done sending + self._wait_for_transfer(producer, req_id, is_receiver=False) if __name__ == "__main__": diff --git a/tpu_sync/api/torch/weight_synchronizer.py b/tpu_sync/api/torch/weight_synchronizer.py index 4cf9481a..f11ec9dd 100644 --- a/tpu_sync/api/torch/weight_synchronizer.py +++ b/tpu_sync/api/torch/weight_synchronizer.py @@ -73,6 +73,10 @@ def push_weights(self, peers: List[str]) -> None: """Trainer pushes model weights to peer inference server coordinates.""" self._impl.PushWeights(peers) + def bind_weights(self, device_tensors: List[List[torch.Tensor]]) -> None: + """Dynamically re-binds new device weights in-place without daemon restart.""" + self._impl.bind_weights(device_tensors) + def d2h(self) -> None: """Triggers asynchronous D2H copy of current weights to Host buffer.""" self._impl.D2h() diff --git a/tpu_sync/api/torch/weight_synchronizer_test.py b/tpu_sync/api/torch/weight_synchronizer_test.py index fca4cd75..28cd3beb 100644 --- a/tpu_sync/api/torch/weight_synchronizer_test.py +++ b/tpu_sync/api/torch/weight_synchronizer_test.py @@ -14,12 +14,8 @@ """E2E physical integration tests for PyTorch WeightSynchronizer on XLA TPUs.""" -import os -import time - from absl.testing import absltest from absl.testing import parameterized -import numpy as np import torch import torch_tpu @@ -35,87 +31,240 @@ def setUp(self): self.num_layers = 2 self.num_shards = 1 self.block_size = 2 - self.slice_byte_size = 16384 // 4 # float32 capacity + + def _run_push_sync( + self, + src_tensors: list[list[torch.Tensor]], + dst_tensors: list[list[torch.Tensor]], + ): + ws_source = WeightSynchronizer( + src_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) + ws_dest = WeightSynchronizer( + dst_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) + self.assertIsNotNone(ws_source.local_port) + self.assertIsNotNone(ws_dest.local_port) + + peer_dest = f"127.0.0.1:{ws_dest.local_port}" + ws_source.push_weights([peer_dest]) + ws_dest.h2d() + + for l in range(len(src_tensors)): + for sh in range(len(src_tensors[l])): + self.assertTrue( + torch.equal(dst_tensors[l][sh].cpu(), src_tensors[l][sh].cpu()) + ) @parameterized.named_parameters( ("fp32", torch.float32), + ("bf16", torch.bfloat16), ("int32", torch.int32), ) def test_e2e_3node_distributed_weight_push(self, dtype): - shape = (self.block_size, 128, 8) # 16384 bytes capacity per layer shard - - # 1. Allocate source (Trainer) weights on Device TPU - src_tensors = [] - for l in range(self.num_layers): - shards = [] - for sh in range(self.num_shards): - t = torch.zeros(shape, dtype=dtype, device=self.device) - shards.append(t) - src_tensors.append(shards) - - # Allocate destination 1 (Inference Peer 1) weights - dst1_tensors = [] - for l in range(self.num_layers): - shards = [] - for sh in range(self.num_shards): - t = torch.zeros(shape, dtype=dtype, device=self.device) - shards.append(t) - dst1_tensors.append(shards) + shape = (self.block_size, 128, 8) - # Allocate destination 2 (Inference Peer 2) weights - dst2_tensors = [] - for l in range(self.num_layers): - shards = [] - for sh in range(self.num_shards): - t = torch.zeros(shape, dtype=dtype, device=self.device) - shards.append(t) - dst2_tensors.append(shards) - - # 2. Instantiate destination WeightSynchronizers on ephemeral ports! - ws_dest1 = WeightSynchronizer(dst1_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1") - ws_dest2 = WeightSynchronizer(dst2_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1") + src_tensors = [ + [ + torch.zeros(shape, dtype=dtype, device=self.device) + for _ in range(self.num_shards) + ] + for _ in range(self.num_layers) + ] + dst1_tensors = [ + [ + torch.zeros(shape, dtype=dtype, device=self.device) + for _ in range(self.num_shards) + ] + for _ in range(self.num_layers) + ] + dst2_tensors = [ + [ + torch.zeros(shape, dtype=dtype, device=self.device) + for _ in range(self.num_shards) + ] + for _ in range(self.num_layers) + ] + ws_dest1 = WeightSynchronizer( + dst1_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) + ws_dest2 = WeightSynchronizer( + dst2_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) self.assertIsNotNone(ws_dest1.local_port) self.assertIsNotNone(ws_dest2.local_port) - peer_dest1 = f"localhost:{ws_dest1.local_port}" - peer_dest2 = f"localhost:{ws_dest2.local_port}" + peer_dest1 = f"127.0.0.1:{ws_dest1.local_port}" + peer_dest2 = f"127.0.0.1:{ws_dest2.local_port}" - # ========================================================================== - # Scenario A: Test the Push API E2E (1 Source pushes to 2 Destinations!) - # ========================================================================== - # Trainer fills source weights with distinct values per layer for l in range(self.num_layers): for sh in range(self.num_shards): - val = float(l + 10.0) # Layer 0=10.0, Layer 1=11.0 + val = int(l + 10) if dtype == torch.int32 else float(l + 10.0) src_tensors[l][sh].fill_(val) - # Force execution of fill_ on source tensors to ensure TPU memory is updated - for l in range(self.num_layers): - for sh in range(self.num_shards): - _ = src_tensors[l][sh].cpu() - - # Recreate/Instantiate ws_source to capture filled buffers! - ws_source = WeightSynchronizer(src_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1") + ws_source = WeightSynchronizer( + src_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) self.assertIsNotNone(ws_source.local_port) - peer_source = f"localhost:{ws_source.local_port}" - # Source pushes weights to both dest1 and dest2 socket servers E2E! ws_source.push_weights([peer_dest1, peer_dest2]) ws_dest1.h2d() ws_dest2.h2d() - # Assert both destinations have received the trainer's weights on their TPU HBM! for l in range(self.num_layers): for sh in range(self.num_shards): - expected_val = float(l + 10.0) - np.testing.assert_allclose( - dst1_tensors[l][sh].cpu().numpy(), expected_val, atol=1e-5 + self.assertTrue( + torch.equal(dst1_tensors[l][sh].cpu(), src_tensors[l][sh].cpu()) ) - np.testing.assert_allclose( - dst2_tensors[l][sh].cpu().numpy(), expected_val, atol=1e-5 + self.assertTrue( + torch.equal(dst2_tensors[l][sh].cpu(), src_tensors[l][sh].cpu()) ) + @parameterized.named_parameters( + ("fp32", torch.float32), + ("bf16", torch.bfloat16), + ) + def test_bind_weights(self, dtype): + shape = (self.block_size, 128, 8) + + src_tensors = [ + [torch.full(shape, fill_value=5.0, dtype=dtype, device=self.device)] + for _ in range(self.num_layers) + ] + dst_tensors = [ + [torch.zeros(shape, dtype=dtype, device=self.device)] + for _ in range(self.num_layers) + ] + + ws_source = WeightSynchronizer( + src_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) + ws_dest = WeightSynchronizer( + dst_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) + peer_dest = f"127.0.0.1:{ws_dest.local_port}" + + # --- Sync 1 (V1: 5.0 -> 0.0) --- + ws_source.push_weights([peer_dest]) + ws_dest.h2d() + + for l in range(self.num_layers): + self.assertTrue( + torch.equal(dst_tensors[l][0].cpu(), src_tensors[l][0].cpu()) + ) + + # --- Bind weights to V2 --- + new_src_tensors = [ + [torch.full(shape, fill_value=10.0, dtype=dtype, device=self.device)] + for _ in range(self.num_layers) + ] + ws_source.bind_weights(new_src_tensors) + ws_source.d2h() + + new_dst_tensors = [ + [torch.full(shape, fill_value=-1.0, dtype=dtype, device=self.device)] + for _ in range(self.num_layers) + ] + ws_dest.bind_weights(new_dst_tensors) + + # --- Sync 2 (V2: 10.0 -> -1.0) --- + ws_source.push_weights([peer_dest]) + ws_dest.h2d() + + # Verify Sync 2 updated new_dst_tensors to 10.0 + for l in range(self.num_layers): + self.assertTrue( + torch.equal(new_dst_tensors[l][0].cpu(), new_src_tensors[l][0].cpu()) + ) + + # Verify original V1 dst_tensors were NOT overwritten (still 5.0) + for l in range(self.num_layers): + self.assertTrue( + torch.equal(dst_tensors[l][0].cpu(), src_tensors[l][0].cpu()) + ) + + @parameterized.named_parameters( + ("fp32", torch.float32), + ("bf16", torch.bfloat16), + ) + def test_heterogeneous_layers_small_first(self, dtype): + shapes = [(1024,), (1024, 3072), (2048, 2048)] + src_tensors = [ + [ + torch.full( + shape, + fill_value=float(i + 1.0), + dtype=dtype, + device=self.device, + ) + ] + for i, shape in enumerate(shapes) + ] + dst_tensors = [ + [torch.zeros(shape, dtype=dtype, device=self.device)] + for shape in shapes + ] + self._run_push_sync(src_tensors, dst_tensors) + + @parameterized.named_parameters( + ("fp32", torch.float32), + ("bf16", torch.bfloat16), + ) + def test_heterogeneous_layers_large_first(self, dtype): + shapes = [(1024, 3072), (1024,), (128,)] + src_tensors = [ + [ + torch.full( + shape, + fill_value=float(i + 1.0), + dtype=dtype, + device=self.device, + ) + ] + for i, shape in enumerate(shapes) + ] + dst_tensors = [ + [torch.zeros(shape, dtype=dtype, device=self.device)] + for shape in shapes + ] + self._run_push_sync(src_tensors, dst_tensors) + + @parameterized.named_parameters( + ("fp32", torch.float32), + ("bf16", torch.bfloat16), + ) + def test_heterogeneous_layers_local_roundtrip(self, dtype): + shapes = [(1024,), (1024, 3072), (2048, 2048)] + src_tensors = [ + [ + torch.full( + shape, + fill_value=float(i + 10.0), + dtype=dtype, + device=self.device, + ) + ] + for i, shape in enumerate(shapes) + ] + ws = WeightSynchronizer( + src_tensors, local_port=0, parallelism=1, bind_ip="127.0.0.1" + ) + ws.d2h() + + # Zero out new destination tensors and bind them + zero_tensors = [ + [torch.zeros(shape, dtype=dtype, device=self.device)] + for shape in shapes + ] + ws.bind_weights(zero_tensors) + ws.h2d() + + for i in range(len(shapes)): + self.assertTrue( + torch.equal(zero_tensors[i][0].cpu(), src_tensors[i][0].cpu()) + ) + if __name__ == "__main__": absltest.main() diff --git a/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc b/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc index a6a41744..00d767ab 100644 --- a/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc +++ b/tpu_sync/frameworks/torch/tpu_raiden_torch_module.cc @@ -535,6 +535,18 @@ NB_MODULE(_tpu_raiden_torch, m) { } }, nb::arg("peers"), nb::call_guard()) + .def( + "bind_weights", + [](WeightSynchronizer& self, + const std::vector>& device_tensors) { + absl::Status s = self.BindWeights(device_tensors); + if (!s.ok()) { + throw std::runtime_error( + "WeightSynchronizer bind_weights failed: " + + std::string(s.message())); + } + }, + nb::arg("device_tensors"), nb::call_guard()) .def( "D2h", diff --git a/tpu_sync/frameworks/torch/weight_synchronizer.cc b/tpu_sync/frameworks/torch/weight_synchronizer.cc index d4fa1775..33f30a2f 100644 --- a/tpu_sync/frameworks/torch/weight_synchronizer.cc +++ b/tpu_sync/frameworks/torch/weight_synchronizer.cc @@ -45,7 +45,17 @@ WeightSynchronizer::WeightSynchronizer( /*external_host_ptrs=*/std::nullopt, unsafe_skip_buffer_lock, parallelism, listener_port, bind_ip, /*layer_names=*/{}, auto_h2d), - buffer_refs_(std::move(unpacked.refs)) {} + buffer_refs_(std::move(unpacked.refs)), + unsafe_skip_buffer_lock_(unsafe_skip_buffer_lock) {} + +absl::Status WeightSynchronizer::BindWeights( + const std::vector>& device_tensors) { + UnpackedTensors unpacked = + UnpackTorchTensors(device_tensors, unsafe_skip_buffer_lock_); + buffer_refs_ = std::move(unpacked.refs); + return weight_sync::WeightSynchronizerBase::BindWeights( + std::move(unpacked.buffers)); +} WeightSynchronizer::~WeightSynchronizer() = default; diff --git a/tpu_sync/frameworks/torch/weight_synchronizer.h b/tpu_sync/frameworks/torch/weight_synchronizer.h index bd402112..11aea0db 100644 --- a/tpu_sync/frameworks/torch/weight_synchronizer.h +++ b/tpu_sync/frameworks/torch/weight_synchronizer.h @@ -38,6 +38,9 @@ class WeightSynchronizer : public weight_sync::WeightSynchronizerBase { bool unsafe_skip_buffer_lock = true, bool auto_h2d = false); + absl::Status BindWeights( + const std::vector>& device_tensors); + ~WeightSynchronizer() override; private: @@ -49,6 +52,7 @@ class WeightSynchronizer : public weight_sync::WeightSynchronizerBase { bool unsafe_skip_buffer_lock, bool auto_h2d); std::vector buffer_refs_; + bool unsafe_skip_buffer_lock_ = true; }; } // namespace torch