From fd322a0e0b89aa0029b9489d99dbf49be03f9790 Mon Sep 17 00:00:00 2001 From: The tunix Authors Date: Tue, 1 Sep 2026 17:43:06 -0700 Subject: [PATCH] Introduce page pool manager for experimental v2 sampler. PiperOrigin-RevId: 974802788 --- .../generate/tiered_page_pool_test.py | 656 ++++++++++++++++++ .../experimental/generate/tiered_page_pool.py | 442 ++++++++++++ 2 files changed, 1098 insertions(+) create mode 100644 tests/experimental/generate/tiered_page_pool_test.py create mode 100644 tunix/experimental/generate/tiered_page_pool.py diff --git a/tests/experimental/generate/tiered_page_pool_test.py b/tests/experimental/generate/tiered_page_pool_test.py new file mode 100644 index 000000000..4f435fe56 --- /dev/null +++ b/tests/experimental/generate/tiered_page_pool_test.py @@ -0,0 +1,656 @@ +import os + +os.environ["XLA_FLAGS"] = "--xla_force_host_platform_device_count=4" + +from absl.testing import absltest +from absl.testing import parameterized +import jax +import jax.numpy as jnp +from jax.sharding import Mesh +import numpy as np + +from tunix.experimental.generate import tiered_page_pool + + +class PagePoolTest(parameterized.TestCase): + + def test_init_state(self): + total_pages = 10 + pages_dict: dict[str, jax.Array | np.ndarray] = { + "layer1": jnp.zeros((total_pages, 8)) + } + pool = tiered_page_pool._PagePool(partition_pages=pages_dict) + self.assertEqual(pool.num_free_pages, total_pages) + self.assertEqual(pool._available_page_indices, list(range(total_pages))) + self.assertEqual(pool._in_use, set()) + + @parameterized.parameters((0,), (1,), (5,)) + def test_allocate(self, num_pages: int): + total_pages = 10 + pages_dict: dict[str, jax.Array | np.ndarray] = { + "layer1": jnp.zeros((total_pages, 8)) + } + pool = tiered_page_pool._PagePool(partition_pages=pages_dict) + prev_unallocated = set(range(total_pages)) + prev_len = total_pages + + allocated = pool.allocate(num_pages) + + # Check set(available page indices) does not have allocated pages + avail_set = set(pool._available_page_indices) + for idx in allocated: + self.assertNotIn(idx, avail_set) + + # Check available page indices contains all unallocated pages + expected_unallocated = prev_unallocated - set(allocated) + self.assertEqual(avail_set, expected_unallocated) + + # Check that returned indices were previously unallocated + for idx in allocated: + self.assertIn(idx, prev_unallocated) + + # Check len + self.assertLen(pool._available_page_indices, prev_len - num_pages) + self.assertEqual(pool.num_free_pages, prev_len - num_pages) + + @parameterized.parameters((0,), (1,), (5,)) + def test_free(self, num_pages: int): + total_pages = 10 + pages_dict: dict[str, jax.Array | np.ndarray] = { + "layer1": jnp.zeros((total_pages, 8)) + } + pool = tiered_page_pool._PagePool(partition_pages=pages_dict) + allocated = pool.allocate(num_pages) + prev_avail = list(pool._available_page_indices) + prev_len = len(prev_avail) + + pool.free(allocated) + + avail_set = set(pool._available_page_indices) + + # Check available page indices contains all previous pages + for idx in prev_avail: + self.assertIn(idx, avail_set) + + # Check available page indices contains new freed pages + for idx in allocated: + self.assertIn(idx, avail_set) + + # Check len + self.assertLen(pool._available_page_indices, prev_len + num_pages) + self.assertEqual(pool.num_free_pages, prev_len + num_pages) + + def test_validations(self): + with self.assertRaisesRegex( + ValueError, r"Partition pages cannot be empty\." + ): + tiered_page_pool._PagePool(partition_pages={}) + + pages_dict: dict[str, jax.Array | np.ndarray] = { + "layer1": jnp.zeros((5, 8)) + } + mismatched: dict[str, jax.Array | np.ndarray] = { + "layer1": jnp.zeros((5, 8)), + "layer2": jnp.zeros((6, 8)), + } + with self.assertRaisesRegex( + ValueError, r"All partitions must have the same number of pages\." + ): + tiered_page_pool._PagePool(partition_pages=mismatched) + + pool = tiered_page_pool._PagePool(partition_pages=pages_dict) + + with self.assertRaisesRegex( + ValueError, r"Cannot allocate a negative number of pages: -1\." + ): + pool.allocate(-1) + + with self.assertRaisesRegex( + ValueError, r"Cannot allocate 10 pages, only 5 available\." + ): + pool.allocate(10) + + allocated = pool.allocate(2) + + with self.assertRaisesRegex( + ValueError, r"Cannot free duplicate page indices\." + ): + pool.free([allocated[0], allocated[0]]) + + with self.assertRaisesRegex( + ValueError, r"Cannot free pages \{0\}\. These pages are not in use\." + ): + pool.free([0]) + + +class TieredPagePoolConfigTest(parameterized.TestCase): + + def test_config_validations(self): + with self.assertRaisesRegex( + ValueError, + r"All dimensions of page_shape must be positive, got 0 in \(10, 0\)\.", + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=0, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + ) + with self.assertRaisesRegex( + ValueError, r"partition_keys cannot be empty\." + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=(), + num_tpu_pages=10, + ) + with self.assertRaisesRegex( + ValueError, r"num_tpu_pages must be positive, got -1\." + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=-1, + ) + with self.assertRaisesRegex( + ValueError, r"num_tpu_pages must be positive, got 0\." + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=0, + ) + with self.assertRaisesRegex( + ValueError, r"num_cpu_pages cannot be negative, got -1\." + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + num_cpu_pages=-1, + ) + with self.assertRaisesRegex( + ValueError, + r"All dimensions of page_shape must be positive, got 0 in \(10, 16, 2," + r" 0\)\.", + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=16, + element_shape=(2, 0), + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + ) + with self.assertRaisesRegex( + ValueError, + r"All dimensions of page_shape must be positive, got -1 in \(10, 16," + r" -1\)\.", + ): + tiered_page_pool.TieredPagePoolConfig( + page_size=16, + element_shape=(-1,), + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + ) + # Valid configs with PartitionSpec succeed. + config_1d = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + page_sharding=jax.sharding.PartitionSpec("dp"), + ) + self.assertIsNotNone(config_1d) + + config_full = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + element_shape=(2, 3), + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + page_sharding=jax.sharding.PartitionSpec("dp", None, None, None), + ) + self.assertIsNotNone(config_full) + + def test_page_shape(self): + config = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + element_shape=(2, 3), + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + ) + self.assertEqual(config.page_shape(5), (5, 16, 2, 3)) + + config_no_subshape = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + ) + self.assertEqual(config_no_subshape.page_shape(5), (5, 16)) + + def test_cpu_sharding_error(self): + config = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + num_cpu_pages=5, + ) + with self.assertRaisesRegex(ValueError, r"Cannot shard pages on CPU\."): + config._make_pool( + num_pages=5, + sharding=jax.sharding.PartitionSpec("dp"), + is_cpu=True, + ) + + def test_init(self): + config = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0", "layer_1"), + num_tpu_pages=10, + num_cpu_pages=5, + ) + tpu_pool, cpu_pool = config.init() + self.assertIsNotNone(cpu_pool) + self.assertEqual(tpu_pool.num_free_pages, 10) + self.assertEqual(cpu_pool.num_free_pages, 5) + self.assertIsInstance(cpu_pool.partition_pages["layer_0"], np.ndarray) + self.assertIsInstance(tpu_pool.partition_pages["layer_0"], jax.Array) + + config_no_cpu = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + num_cpu_pages=0, + ) + _, cpu_pool_2 = config_no_cpu.init() + self.assertIsNone(cpu_pool_2) + + +class InternalHelpersTest(parameterized.TestCase): + + def test_scatter_tpu_pages(self): + tpu_pages = { + "layer_0": jnp.zeros((4, 8), dtype=jnp.float32), + "layer_1": jnp.zeros((4, 8), dtype=jnp.float32), + } + indices = jnp.array([1, 3], dtype=jnp.int32) + slices = { + "layer_0": jnp.ones((2, 8), dtype=jnp.float32), + "layer_1": jnp.full((2, 8), 2.0, dtype=jnp.float32), + } + updated = tiered_page_pool._scatter_tpu_pages(tpu_pages, indices, slices) + np.testing.assert_allclose(updated["layer_0"][1], np.ones(8)) + np.testing.assert_allclose(updated["layer_0"][3], np.ones(8)) + np.testing.assert_allclose(updated["layer_0"][0], np.zeros(8)) + np.testing.assert_allclose(updated["layer_0"][2], np.zeros(8)) + np.testing.assert_allclose(updated["layer_1"][1], np.full(8, 2.0)) + np.testing.assert_allclose(updated["layer_1"][3], np.full(8, 2.0)) + + def test_get_tpu_slices(self): + layer_0 = jnp.arange(32, dtype=jnp.float32).reshape((4, 8)) + tpu_pages = {"layer_0": layer_0} + indices = jnp.array([0, 2], dtype=jnp.int32) + slices = tiered_page_pool._get_tpu_slices(tpu_pages, indices) + np.testing.assert_allclose(slices["layer_0"], layer_0[indices]) + + +class TieredPagePoolManagerTest(parameterized.TestCase): + + def setUp(self): + super().setUp() + if len(jax.devices()) < 4: + self.skipTest("Requires at least 4 devices") + mesh_shape = (2, 2) + self.devices = np.array(jax.devices()[:4]).reshape(mesh_shape) + self.mesh = Mesh(self.devices, axis_names=("dp", "tp")) + + def get_config(self, sharding_type: str, has_subshape: bool = True): + page_size = 16 + element_shape = (2, 1, 5) if has_subshape else () + page_sharding = None + + if has_subshape: + if sharding_type == "dp TPU sharding": + page_sharding = jax.sharding.PartitionSpec("dp", None, None, None, None) + elif sharding_type == "tp TPU sharding": + page_sharding = jax.sharding.PartitionSpec(None, None, "tp", None, None) + elif sharding_type == "dp + tp TPU sharding": + page_sharding = jax.sharding.PartitionSpec("dp", None, "tp", None, None) + else: + if sharding_type == "dp TPU sharding": + page_sharding = jax.sharding.PartitionSpec("dp", None) + elif sharding_type == "tp TPU sharding": + page_sharding = jax.sharding.PartitionSpec(None, "tp") + elif sharding_type == "dp + tp TPU sharding": + page_sharding = jax.sharding.PartitionSpec("dp", "tp") + + return tiered_page_pool.TieredPagePoolConfig( + page_size=page_size, + element_shape=element_shape, + dtype=jnp.float32, + partition_keys=("layer_0", "layer_1"), + num_tpu_pages=10, + num_cpu_pages=10, + page_sharding=page_sharding, + ) + + @parameterized.parameters((0,), (1,), (5,)) + def test_allocate_tpu_pages(self, num_pages: int): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, cpu_pool = config.init() + manager = tiered_page_pool.TieredPagePoolManager(config, tpu_pool, cpu_pool) + + allocated = manager.allocate_tpu_pages(num_pages) + + self.assertLen(set(allocated), num_pages) + for pid in allocated: + self.assertEqual(manager.get_page_location(pid), "tpu") + phys_idx = manager.get_page_idx(pid) + self.assertNotIn(phys_idx, manager.tpu_pool._available_page_indices) + + def test_allocate_tpu_pages_errors(self): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, _ = config.init() + manager = tiered_page_pool.TieredPagePoolManager(config, tpu_pool, None) + + with self.assertRaisesRegex( + ValueError, r"Cannot allocate a negative number of pages\." + ): + manager.allocate_tpu_pages(-1) + + with self.assertRaisesRegex( + ValueError, r"Cannot allocate 100 TPU pages, only 10 available\." + ): + manager.allocate_tpu_pages(100) + + def test_num_free_cpu_pages_when_cpu_pool_is_none(self): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, _ = config.init() + manager = tiered_page_pool.TieredPagePoolManager( + config, tpu_pool, cpu_pool=None + ) + self.assertEqual(manager.num_free_cpu_pages, 0) + + config_no_cpu = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + num_cpu_pages=0, + ) + tpu_pool_no_cpu, cpu_pool_no_cpu = config_no_cpu.init() + self.assertIsNone(cpu_pool_no_cpu) + manager_no_cpu = tiered_page_pool.TieredPagePoolManager( + config_no_cpu, tpu_pool_no_cpu, cpu_pool_no_cpu + ) + self.assertEqual(manager_no_cpu.num_free_cpu_pages, 0) + + @parameterized.product( + [ + dict(sharding_type="no TPU sharding", has_subshape=False), + dict(sharding_type="no TPU sharding", has_subshape=True), + dict(sharding_type="dp TPU sharding", has_subshape=True), + dict(sharding_type="tp TPU sharding", has_subshape=True), + dict(sharding_type="dp + tp TPU sharding", has_subshape=True), + ], + num_pages=[1, 2, 5], + ) + def test_load_offload( + self, sharding_type: str, has_subshape: bool, num_pages: int + ): + with jax.set_mesh(self.mesh): + config = self.get_config(sharding_type, has_subshape=has_subshape) + tpu_pool, cpu_pool = config.init() + manager = tiered_page_pool.TieredPagePoolManager( + config, tpu_pool, cpu_pool + ) + + n_layers = len(config.partition_keys) + page_vals = np.zeros((n_layers, num_pages), dtype=np.float32) + for l in range(n_layers): + for p in range(num_pages): + page_vals[l, p] = (l + 1) * 100.0 + (p + 1) + + tpu_pids = manager.allocate_tpu_pages(num_pages) + orig_tpu_idxs: list[int] = [] + for pid in tpu_pids: + idx = manager.get_page_idx(pid) + self.assertIsNotNone(idx) + assert idx is not None + orig_tpu_idxs.append(idx) + + # Populate allocated TPU pages with distinct values from page_vals. + new_tpu_pages = dict(manager.tpu_pool.partition_pages) + for l_idx, layer in enumerate(config.partition_keys): + pages = new_tpu_pages[layer] + assert isinstance(pages, jax.Array) + for p_idx, phys_idx in enumerate(orig_tpu_idxs): + pages = pages.at[phys_idx].set(page_vals[l_idx, p_idx]) + new_tpu_pages[layer] = pages + manager.update_tpu_pool(new_tpu_pages) + + prev_cpu_free = manager.num_free_cpu_pages + prev_tpu_free = manager.num_free_tpu_pages + + manager.offload(tpu_pids) + + for pid in tpu_pids: + self.assertEqual(manager.get_page_location(pid), "cpu") + + self.assertEqual(manager.num_free_cpu_pages, prev_cpu_free - num_pages) + self.assertEqual(manager.num_free_tpu_pages, prev_tpu_free + num_pages) + + # Verify pages on CPU have their distinct values per page and per layer. + self.assertIsNotNone(manager.cpu_pool) + for l_idx, layer in enumerate(config.partition_keys): + cpu_pages = manager.cpu_pool.partition_pages[layer] + for p_idx, pid in enumerate(tpu_pids): + cpu_idx = manager.get_page_idx(pid) + assert cpu_idx is not None + np.testing.assert_allclose( + cpu_pages[cpu_idx], page_vals[l_idx, p_idx] + ) + + # Overwrite the freed TPU slots with sentinel values before loading to + # guarantee that load() actively transfers data rather than reusing stale + # TPU buffers. + dirty_tpu_pages = {} + for layer, pages in manager.tpu_pool.partition_pages.items(): + assert isinstance(pages, jax.Array) + for phys_idx in orig_tpu_idxs: + pages = pages.at[phys_idx].set(-999.0) + dirty_tpu_pages[layer] = pages + manager.update_tpu_pool(dirty_tpu_pages) + + for pages in manager.tpu_pool.partition_pages.values(): + for phys_idx in orig_tpu_idxs: + np.testing.assert_allclose(pages[phys_idx], -999.0) + + manager.load(tpu_pids) + for pid in tpu_pids: + self.assertEqual(manager.get_page_location(pid), "tpu") + + self.assertEqual(manager.num_free_cpu_pages, prev_cpu_free) + self.assertEqual(manager.num_free_tpu_pages, prev_tpu_free) + + # Verify TPU pages have restored their distinct per-page and per-layer values. + for l_idx, layer in enumerate(config.partition_keys): + hbm_pages = manager.tpu_pool.partition_pages[layer] + for p_idx, pid in enumerate(tpu_pids): + tpu_idx = manager.get_page_idx(pid) + assert tpu_idx is not None + np.testing.assert_allclose( + hbm_pages[tpu_idx], page_vals[l_idx, p_idx] + ) + if config.page_sharding is not None and hasattr(hbm_pages, "sharding"): + self.assertEqual( + hbm_pages.sharding, + jax.sharding.NamedSharding(self.mesh, config.page_sharding), + ) + + def test_empty_load_offload(self): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, cpu_pool = config.init() + manager = tiered_page_pool.TieredPagePoolManager(config, tpu_pool, cpu_pool) + # Empty operations should be no-ops + manager.load([]) + manager.offload([]) + + def test_load_offload_errors(self): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, _ = config.init() + manager_no_cpu = tiered_page_pool.TieredPagePoolManager( + config, tpu_pool, None + ) + pids = manager_no_cpu.allocate_tpu_pages(2) + with self.assertRaisesRegex( + ValueError, + r"Cannot offload pages to CPU, CPU pool is not initialized\.", + ): + manager_no_cpu.offload(pids) + + with self.assertRaisesRegex( + ValueError, + r"Cannot load pages from CPU to TPU, CPU pool is not initialized\.", + ): + manager_no_cpu.load(pids) + + tpu_pool, cpu_pool = config.init() + manager = tiered_page_pool.TieredPagePoolManager(config, tpu_pool, cpu_pool) + tpu_pids = manager.allocate_tpu_pages(2) + + with self.assertRaisesRegex( + ValueError, r"Cannot offload duplicate pages\." + ): + manager.offload([tpu_pids[0], tpu_pids[0]]) + + with self.assertRaisesRegex( + ValueError, r"Page ID 999 is not on TPU \(location: None\)\." + ): + manager.offload([999]) + + config_small_cpu = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=10, + num_cpu_pages=1, + ) + tpu_p, cpu_p = config_small_cpu.init() + mgr_small = tiered_page_pool.TieredPagePoolManager( + config_small_cpu, tpu_p, cpu_p + ) + more_pids = mgr_small.allocate_tpu_pages(2) + with self.assertRaisesRegex( + ValueError, r"Cannot offload 2 pages, only 1 available\." + ): + mgr_small.offload(more_pids) + + # Attempting to load a page that is already on TPU + with self.assertRaisesRegex( + ValueError, r"Page ID \d+ is not on CPU \(location: tpu\)\." + ): + manager.load(tpu_pids) + + manager.offload(tpu_pids) + + # Attempting to offload a page that is already on CPU + with self.assertRaisesRegex( + ValueError, r"Page ID \d+ is not on TPU \(location: cpu\)\." + ): + manager.offload(tpu_pids) + + with self.assertRaisesRegex(ValueError, r"Cannot load duplicate pages\."): + manager.load([tpu_pids[0], tpu_pids[0]]) + + with self.assertRaisesRegex( + ValueError, r"Page ID 999 is not on CPU \(location: None\)\." + ): + manager.load([999]) + + # Test load when TPU pool is full / has insufficient free pages. + config_small_tpu = tiered_page_pool.TieredPagePoolConfig( + page_size=16, + dtype=jnp.float32, + partition_keys=("layer_0",), + num_tpu_pages=2, + num_cpu_pages=2, + ) + tpu_p2, cpu_p2 = config_small_tpu.init() + mgr_small_tpu = tiered_page_pool.TieredPagePoolManager( + config_small_tpu, tpu_p2, cpu_p2 + ) + pids_tpu = mgr_small_tpu.allocate_tpu_pages(2) + mgr_small_tpu.offload(pids_tpu) + # Re-allocate TPU pool to capacity so 0 free TPU pages remain + _ = mgr_small_tpu.allocate_tpu_pages(2) + with self.assertRaisesRegex( + ValueError, r"Cannot load 2 pages, only 0 available\." + ): + mgr_small_tpu.load(pids_tpu) + + def test_free(self): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, cpu_pool = config.init() + manager = tiered_page_pool.TieredPagePoolManager(config, tpu_pool, cpu_pool) + + tpu_pids = manager.allocate_tpu_pages(4) + # Offload 2 pages to CPU + manager.offload(tpu_pids[:2]) + + self.assertEqual(manager.get_page_location(tpu_pids[0]), "cpu") + self.assertEqual(manager.get_page_location(tpu_pids[2]), "tpu") + + prev_tpu_free = manager.num_free_tpu_pages + prev_cpu_free = manager.num_free_cpu_pages + + manager.free(tpu_pids) + + self.assertEqual(manager.num_free_tpu_pages, prev_tpu_free + 2) + self.assertEqual(manager.num_free_cpu_pages, prev_cpu_free + 2) + + for pid in tpu_pids: + self.assertIsNone(manager.get_page_location(pid)) + self.assertIsNone(manager.get_page_idx(pid)) + + # Freeing an empty list is a safe no-op. + manager.free([]) + + # Test freeing only CPU pages + tpu_pids_cpu_only = manager.allocate_tpu_pages(2) + manager.offload(tpu_pids_cpu_only) + manager.free(tpu_pids_cpu_only) + self.assertIsNone(manager.get_page_location(tpu_pids_cpu_only[0])) + + # Test freeing only TPU pages + tpu_pids_tpu_only = manager.allocate_tpu_pages(2) + manager.free(tpu_pids_tpu_only) + self.assertIsNone(manager.get_page_location(tpu_pids_tpu_only[0])) + + def test_free_errors(self): + config = self.get_config("no TPU sharding", has_subshape=False) + tpu_pool, cpu_pool = config.init() + manager = tiered_page_pool.TieredPagePoolManager(config, tpu_pool, cpu_pool) + pids = manager.allocate_tpu_pages(2) + + with self.assertRaisesRegex( + ValueError, r"Attempting to free page 999 which is not in use\." + ): + manager.free([999]) + + with self.assertRaisesRegex(ValueError, r"Cannot free duplicate pages\."): + manager.free([pids[0], pids[0]]) + + +if __name__ == "__main__": + absltest.main() diff --git a/tunix/experimental/generate/tiered_page_pool.py b/tunix/experimental/generate/tiered_page_pool.py new file mode 100644 index 000000000..de1da058a --- /dev/null +++ b/tunix/experimental/generate/tiered_page_pool.py @@ -0,0 +1,442 @@ +"""A tiered memory cache manager for TPU and CPU memory. + +This module acts as the interface for logical page management, abstracting away +the allocation, tracking, and transfer of physical pages across devices. + +It is used by the KV cache coordinator to track and manage logical pages for +requests. + +Key components: + TieredPagePoolManager: Provides an interface for page pools, exposing logical + page IDs instead of physical indices. It handles the allocation, tracking, + and transfer of low-level physical page indices across devices. + TieredPagePoolConfig: Defines the layout, sharding, and capacity of the + underlying physical page pools. It is used to initialize the page pools. +""" + +from collections.abc import Sequence +import dataclasses +import functools +import jax +import jax.numpy as jnp +import numpy as np + + +@dataclasses.dataclass(kw_only=True) +class _PagePool: + """A pool of pages.""" + + # A mapping of partition names to the pages for that partition. + partition_pages: dict[str, jax.Array | np.ndarray] + # A list of available page indices across all partitions. + _available_page_indices: list[int] = dataclasses.field( + default_factory=list, init=False + ) + # A set of allocated pages. This is used to validate and prevent double-free + # operations. + _in_use: set[int] = dataclasses.field(default_factory=set, init=False) + + def __post_init__(self): + if not self.partition_pages: + raise ValueError("Partition pages cannot be empty.") + + n_pages = [self.partition_pages[k].shape[0] for k in self.partition_pages] + for n in n_pages: + if n != n_pages[0]: + raise ValueError("All partitions must have the same number of pages.") + + self._available_page_indices = list(range(n_pages[0])) + self._in_use = set() + + def allocate(self, num_pages: int) -> list[int]: + """Allocates pages in the pool.""" + if num_pages < 0: + raise ValueError( + f"Cannot allocate a negative number of pages: {num_pages}." + ) + + if num_pages > self.num_free_pages: + raise ValueError( + f"Cannot allocate {num_pages} pages, " + f"only {self.num_free_pages} available." + ) + + if num_pages == 0: + return [] + + indices = self._available_page_indices[-num_pages:] + del self._available_page_indices[-num_pages:] + + self._in_use.update(indices) + + return indices + + def free(self, indices: Sequence[int]) -> None: + """Frees pages in the pool.""" + indices_set = set(indices) + if len(indices_set) != len(indices): + raise ValueError("Cannot free duplicate page indices.") + + if len(indices_set - self._in_use) > 0: + raise ValueError( + f"Cannot free pages {indices_set - self._in_use}. " + "These pages are not in use." + ) + + for idx in indices: + self._in_use.remove(idx) + self._available_page_indices.extend(indices) + + @property + def num_free_pages(self) -> int: + return len(self._available_page_indices) + + +@dataclasses.dataclass(frozen=True, kw_only=True) +class TieredPagePoolConfig: + """Configuration for tiered page pool.""" + + # The number of elements in a page. + page_size: int + # The shape of an individual element in a page. + element_shape: tuple[int, ...] = () + # The data type of the elements in a page. + dtype: jnp.dtype + # The names of the pool partitions (e.g. layer1, layer2). + partition_keys: tuple[str, ...] + # The number of TPU pages to allocate. + num_tpu_pages: int + # The number of CPU pages to allocate. + num_cpu_pages: int = 0 + # The TPU sharding of the page pool tensor. + # (Page dim, elements dim, element shape dim0, element shape dim1, ...) + page_sharding: jax.sharding.PartitionSpec | None = None + + def __post_init__(self): + if self.num_tpu_pages <= 0: + raise ValueError( + f"num_tpu_pages must be positive, got {self.num_tpu_pages}." + ) + if self.num_cpu_pages < 0: + raise ValueError( + f"num_cpu_pages cannot be negative, got {self.num_cpu_pages}." + ) + if not self.partition_keys: + raise ValueError("partition_keys cannot be empty.") + + page_shape = self.page_shape() + if not page_shape: + raise ValueError("page_shape cannot be empty.") + + for dim in page_shape: + if dim <= 0: + raise ValueError( + f"All dimensions of page_shape must be positive, got {dim} in" + f" {page_shape}." + ) + + def page_shape(self, num_pages: int | None = None) -> tuple[int, ...]: + if num_pages is None: + num_pages = self.num_tpu_pages + return (num_pages, self.page_size, *self.element_shape) + + def _make_pool( + self, + num_pages: int, + sharding: jax.sharding.PartitionSpec | None = None, + is_cpu: bool = False, + ) -> _PagePool: + """Creates a page pool.""" + if is_cpu and sharding is not None: + raise ValueError("Cannot shard pages on CPU.") + + pages_dict = {} + page_shape = self.page_shape(num_pages) + + if sharding is not None: + init_sharded_fn = jax.jit( + lambda: jnp.zeros(page_shape, dtype=self.dtype), + out_shardings=sharding, + ) + for k in self.partition_keys: + pages_dict[k] = init_sharded_fn() + elif is_cpu: + for k in self.partition_keys: + pages_dict[k] = np.zeros(page_shape, dtype=self.dtype) + else: + for k in self.partition_keys: + pages_dict[k] = jnp.zeros(page_shape, dtype=self.dtype) + + return _PagePool( + partition_pages=pages_dict, + ) + + def init(self) -> tuple[_PagePool, _PagePool | None]: + """Initializes physical page tensors for TPU and CPU.""" + + tpu_pool = self._make_pool( + num_pages=self.num_tpu_pages, sharding=self.page_sharding + ) + + cpu_pool = ( + self._make_pool(num_pages=self.num_cpu_pages, is_cpu=True) + if self.num_cpu_pages > 0 + else None + ) + + return (tpu_pool, cpu_pool) + + +@functools.partial(jax.jit, donate_argnames=("tpu_pages", "slices")) +def _scatter_tpu_pages( + tpu_pages: dict[str, jax.Array], + indices: jax.Array, + slices: dict[str, jax.Array], +) -> dict[str, jax.Array]: + """Scatters pages into the TPU page pool. + + To optimize memory and prevent XLA from allocating redundant copies, the + `tpu_pages` and `slices` buffers are donated. The function is jitted so that + these scattered updates are executed concurrently across all TPU partitions. + + Args: + tpu_pages: The current TPU page pool. + indices: The indices of the pages to scatter. + slices: The slices of the pages to scatter. + + Returns: + The updated TPU page pool. + """ + return {k: tpu_pages[k].at[indices].set(slices[k]) for k in tpu_pages} + + +@jax.jit +def _get_tpu_slices( + tpu_pages: dict[str, jax.Array], + indices: jax.Array, +) -> dict[str, jax.Array]: + """Returns the slices of the TPU pages for the given indices. + + The function is jitted to ensure these slices are executed concurrently across + all TPU partitions. + + Args: + tpu_pages: The current TPU page pool. + indices: The indices of the pages to get slices for. + + Returns: + The slices of the TPU pages for the given indices. + """ + return {layer: tpu_pages[layer][indices] for layer in tpu_pages} + + +class TieredPagePoolManager: + """Manager for tiered TPU/CPU memory.""" + + def __init__( + self, + tiered_config: TieredPagePoolConfig, + tpu_pool: _PagePool, + cpu_pool: _PagePool | None, # CPU pool is optional. + ): + self.config = tiered_config + self.tpu_pool = tpu_pool + self.cpu_pool = cpu_pool + self.page_size = tiered_config.page_size + + self._next_page_id: int = 0 + self._page_id_to_idx: dict[int, int] = {} + self._page_location: dict[int, str] = {} + + @property + def num_free_tpu_pages(self) -> int: + return self.tpu_pool.num_free_pages + + @property + def num_free_cpu_pages(self) -> int: + if self.cpu_pool: + return self.cpu_pool.num_free_pages + return 0 + + def get_page_location(self, page_id: int) -> str | None: + return self._page_location.get(page_id) + + def get_page_idx(self, page_id: int) -> int | None: + return self._page_id_to_idx.get(page_id) + + def allocate_tpu_pages(self, num_pages: int) -> list[int]: + """Allocate logical TPU pages.""" + if num_pages < 0: + raise ValueError("Cannot allocate a negative number of pages.") + + if num_pages == 0: + return [] + + if num_pages > self.num_free_tpu_pages: + raise ValueError( + f"Cannot allocate {num_pages} TPU pages, " + f"only {self.num_free_tpu_pages} available." + ) + + allocated_ids = [] + phys_indices = self.tpu_pool.allocate(num_pages) + + for phys_idx in phys_indices: + pid = self._next_page_id + self._next_page_id += 1 + self._page_id_to_idx[pid] = phys_idx + self._page_location[pid] = "tpu" + + allocated_ids.append(pid) + + return allocated_ids + + def update_tpu_pool( + self, new_pages: dict[str, jax.Array | np.ndarray] + ) -> None: + """Updates the underlying TPU pool partition pages with new pages.""" + self.tpu_pool.partition_pages = new_pages + + @property + def _tpu_sharding(self) -> jax.sharding.Sharding | None: + first_layer_pages = next(iter(self.tpu_pool.partition_pages.values())) + return getattr(first_layer_pages, "sharding", None) + + @property + def _transferred_page_sharding(self) -> jax.sharding.Sharding | None: + """Returns the target sharding for transferring to TPU.""" + tpu_sharding = self._tpu_sharding + + # Replicate the page dimension across all devices, so that pages are + # scattered without cross-device communication. + if isinstance(tpu_sharding, jax.sharding.NamedSharding): + slice_spec = jax.sharding.PartitionSpec(None, *tpu_sharding.spec[1:]) + return jax.sharding.NamedSharding(tpu_sharding.mesh, slice_spec) + + return tpu_sharding + + def load(self, page_ids: Sequence[int]) -> None: + """Transfers logical pages from CPU to TPU.""" + if not page_ids: + return + + if self.cpu_pool is None: + raise ValueError( + "Cannot load pages from CPU to TPU, CPU pool is not initialized." + ) + + if len(page_ids) > self.num_free_tpu_pages: + raise ValueError( + f"Cannot load {len(page_ids)} pages, " + f"only {self.num_free_tpu_pages} available." + ) + + if len(set(page_ids)) != len(page_ids): + raise ValueError("Cannot load duplicate pages.") + + for pid in page_ids: + if self._page_location.get(pid) != "cpu": + raise ValueError( + f"Page ID {pid} is not on CPU " + f"(location: {self._page_location.get(pid)})." + ) + + cpu_idxs = [self._page_id_to_idx[pid] for pid in page_ids] + tpu_idxs = self.tpu_pool.allocate(len(page_ids)) + + # Gather all the pages that need to be transferred to TPU. + cpu_slices = { + k: self.cpu_pool.partition_pages[k][cpu_idxs] + for k in self.tpu_pool.partition_pages + } + + # Transfer pages to the TPU + tpu_slices = jax.device_put(cpu_slices, self._transferred_page_sharding) + + # Scatter pages to the TPU partitions. + tpu_indices_arr = jnp.array(tpu_idxs, dtype=jnp.int32) + tpu_partitions = self.tpu_pool.partition_pages + + # Use a jit-compiled function to avoid replicating page pools when + # scattering pages across all TPU partitions. + self.tpu_pool.partition_pages = _scatter_tpu_pages( + tpu_partitions, tpu_indices_arr, tpu_slices + ) + + # Update page state + self.cpu_pool.free(cpu_idxs) + for pid, p_idx in zip(page_ids, tpu_idxs): + self._page_id_to_idx[pid] = p_idx + self._page_location[pid] = "tpu" + + def offload(self, page_ids: Sequence[int]) -> None: + """Moves logical pages from TPU to CPU transferring only active pages.""" + if not page_ids: + return + + if self.cpu_pool is None: + raise ValueError( + "Cannot offload pages to CPU, CPU pool is not initialized." + ) + + if len(page_ids) > self.num_free_cpu_pages: + raise ValueError( + f"Cannot offload {len(page_ids)} pages, " + f"only {self.num_free_cpu_pages} available." + ) + + if len(set(page_ids)) != len(page_ids): + raise ValueError("Cannot offload duplicate pages.") + + for pid in page_ids: + if self._page_location.get(pid) != "tpu": + raise ValueError( + f"Page ID {pid} is not on TPU " + f"(location: {self._page_location.get(pid)})." + ) + + physical_tpu_idxs = [self._page_id_to_idx[pid] for pid in page_ids] + physical_cpu_idxs = self.cpu_pool.allocate(len(page_ids)) + tpu_indices_arr = jnp.array(physical_tpu_idxs, dtype=jnp.int32) + + # Use a jit-compiled function here to concurrently gather slices from all + # TPU partitions. + tpu_slices = _get_tpu_slices(self.tpu_pool.partition_pages, tpu_indices_arr) + host_slices = jax.device_get(tpu_slices) + for layer, host_slice in host_slices.items(): + self.cpu_pool.partition_pages[layer][physical_cpu_idxs] = host_slice + + self.tpu_pool.free(physical_tpu_idxs) + for pid, p_idx in zip(page_ids, physical_cpu_idxs): + self._page_id_to_idx[pid] = p_idx + self._page_location[pid] = "cpu" + + def free(self, page_ids: Sequence[int]) -> None: + """Releases physical allocations in tpu_pool or cpu_pool and removes logical IDs.""" + if not page_ids: + return + + if len(set(page_ids)) != len(page_ids): + raise ValueError("Cannot free duplicate pages.") + + for pid in page_ids: + if pid not in self._page_location or pid not in self._page_id_to_idx: + raise ValueError(f"Attempting to free page {pid} which is not in use.") + + cpu_idxs_to_free = [] + tpu_idxs_to_free = [] + + for pid in page_ids: + loc = self._page_location[pid] + if loc == "cpu": + cpu_idxs_to_free.append(self._page_id_to_idx[pid]) + elif loc == "tpu": + tpu_idxs_to_free.append(self._page_id_to_idx[pid]) + + del self._page_location[pid] + del self._page_id_to_idx[pid] + + if cpu_idxs_to_free and self.cpu_pool: + self.cpu_pool.free(cpu_idxs_to_free) + if tpu_idxs_to_free and self.tpu_pool: + self.tpu_pool.free(tpu_idxs_to_free)