diff --git a/test/registered/unit/managers/test_hicache_controller_cp.py b/test/registered/unit/managers/test_hicache_controller_cp.py index e898462f0..49d25d9a8 100644 --- a/test/registered/unit/managers/test_hicache_controller_cp.py +++ b/test/registered/unit/managers/test_hicache_controller_cp.py @@ -285,7 +285,9 @@ class TestHiCacheControllerCPWrite(CustomTestCase): controller = self.make_controller(host_pool, cp_rank=1) logical_locs = torch.tensor([8, 9, 10], dtype=torch.int64) - with self.assertRaisesRegex(ValueError, "host_indices.*whole pages"): + with self.assertRaisesRegex( + ValueError, "(host_indices|physical_device_indices).*whole pages" + ): controller.write(logical_locs, node_id=21) def test_cp_write_rejects_non_contiguous_owned_physical_page(self): diff --git a/test/registered/unit/mem_cache/test_cp_hicache_metadata.py b/test/registered/unit/mem_cache/test_cp_hicache_metadata.py index 117c2acf6..dc8d16fa9 100644 --- a/test/registered/unit/mem_cache/test_cp_hicache_metadata.py +++ b/test/registered/unit/mem_cache/test_cp_hicache_metadata.py @@ -668,6 +668,7 @@ class TestHiRadixCacheCPSplitEvict(CustomTestCase): cache._uses_cp_hicache = True cache.tp_world_size = 2 cache.tp_group = object() + cache._tp_group_rank = 0 cache.root_node = TreeNode() cache.root_node.key = RadixKey([]) cache.evictable_host_leaves = set() @@ -700,7 +701,7 @@ class TestHiRadixCacheCPSplitEvict(CustomTestCase): cache.evictable_host_leaves.add(node) all_done_states = iter([False, True]) - cache._cp_all_ranks_true = lambda done: next(all_done_states) + cache._cp_all_ranks_true = lambda done: next(all_done_states, True) cache._cp_broadcast_node_ids = lambda node_ids, max_ids: node_ids[:max_ids] cache._cp_filter_all_ranks_safe_node_ids = ( lambda node_ids, is_safe, **_kwargs: [