diff --git a/data_juicer/ops/deduplicator/ray_bts_minhash_deduplicator.py b/data_juicer/ops/deduplicator/ray_bts_minhash_deduplicator.py index acdd87df905..a50bd977a3e 100644 --- a/data_juicer/ops/deduplicator/ray_bts_minhash_deduplicator.py +++ b/data_juicer/ops/deduplicator/ray_bts_minhash_deduplicator.py @@ -26,6 +26,89 @@ ray = LazyLoader("ray") BATCH_SIZE = 1000 +UID_DTYPE = np.dtype(np.int64) +EDGE_DTYPE = np.dtype([("u", UID_DTYPE), ("v", UID_DTYPE)]) +_OBJECT_EDGE_DTYPE = np.dtype([("u", object), ("v", object)]) + + +def _empty_edge_array(): + return np.empty(0, dtype=EDGE_DTYPE) + + +def _parent_to_arrays(parent): + """Convert a parent mapping to compact, order-aligned UID arrays.""" + size = len(parent) + try: + uids = np.fromiter(parent.keys(), dtype=UID_DTYPE, count=size) + parents = np.fromiter(parent.values(), dtype=UID_DTYPE, count=size) + except OverflowError: + uids = np.fromiter(parent.keys(), dtype=object, count=size) + parents = np.fromiter(parent.values(), dtype=object, count=size) + return uids, parents + + +def _group_edges_by_destination(uids, parents, parallel_num): + """Pack edges into one array grouped by their destination actor. + + Every edge is sent to the actor owning its source UID and, when the two + endpoints have different owners, to the actor owning its parent UID too. + The returned offsets delimit the slice for each destination. + """ + if parallel_num <= 0: + raise ValueError("parallel_num must be positive") + if len(uids) != len(parents): + raise ValueError("uids and parents must have the same length") + if len(uids) == 0: + return _empty_edge_array(), np.zeros(parallel_num + 1, dtype=np.int64) + + use_object_edges = uids.dtype.hasobject or parents.dtype.hasobject + if use_object_edges: + hash_u = np.fromiter((int(uid) // BATCH_SIZE % parallel_num for uid in uids), dtype=np.intp, count=len(uids)) + hash_v = np.fromiter( + (int(parent) // BATCH_SIZE % parallel_num for parent in parents), dtype=np.intp, count=len(parents) + ) + else: + hash_u = (uids // BATCH_SIZE) % parallel_num + hash_v = (parents // BATCH_SIZE) % parallel_num + cross_partition = hash_u != hash_v + cross_count = int(np.count_nonzero(cross_partition)) + + size = len(uids) + cross_count + destination_dtype = np.min_scalar_type(parallel_num - 1) + destinations = np.empty(size, dtype=destination_dtype) + destinations[: len(uids)] = hash_u + edge_dtype = _OBJECT_EDGE_DTYPE if use_object_edges else EDGE_DTYPE + unsorted_edges = np.empty(size, dtype=edge_dtype) + unsorted_edges["u"][: len(uids)] = uids + unsorted_edges["v"][: len(uids)] = parents + if cross_count: + destinations[len(uids) :] = hash_v[cross_partition] + unsorted_edges["u"][len(uids) :] = uids[cross_partition] + unsorted_edges["v"][len(uids) :] = parents[cross_partition] + del hash_u, hash_v, cross_partition + + # Union always chooses the smaller root, so the order within a destination + # does not affect the resulting connected components. + counts = np.bincount(destinations, minlength=parallel_num) + offsets = np.empty(parallel_num + 1, dtype=np.int64) + offsets[0] = 0 + np.cumsum(counts, out=offsets[1:]) + del counts + + order = np.argsort(destinations) + del destinations + edges = np.empty(size, dtype=edge_dtype) + np.take(unsorted_edges["u"], order, out=edges["u"]) + np.take(unsorted_edges["v"], order, out=edges["v"]) + return edges, offsets + + +def _union_edges(union_find, edges, chunk_size=1 << 18): + """Union compact edges without materializing all Python integers at once.""" + for start in range(0, len(edges), chunk_size): + end = min(start + chunk_size, len(edges)) + for u, v in zip(edges["u"][start:end].tolist(), edges["v"][start:end].tolist()): + union_find.union(u, v) class IdGenerator: @@ -42,16 +125,31 @@ def get_next_id(self, count): class EdgeBuffer: def __init__(self): - self.edge_dict = {} + self.clear() def clear(self): - self.edge_dict = {} + self.edges = _empty_edge_array() + self.offsets = np.zeros(1, dtype=np.int64) + self.consumed = np.empty(0, dtype=np.bool_) + self.remaining = 0 - def set_edges(self, edge_dict): - self.edge_dict = edge_dict + def set_edges(self, edges, offsets): + self.edges = edges + self.offsets = offsets + self.consumed = np.zeros(len(offsets) - 1, dtype=np.bool_) + self.remaining = len(self.consumed) def get_edges(self, key): - return self.edge_dict.pop(key, []) + if key < 0 or key >= len(self.consumed) or self.consumed[key]: + return _empty_edge_array() + + self.consumed[key] = True + self.remaining -= 1 + start, end = self.offsets[key : key + 2] + result = self.edges[start:end] + if self.remaining == 0: + self.clear() + return result class BTSUnionFind: @@ -78,8 +176,6 @@ def __init__( self.parent = {} self.old_parent = {} self.remote_edge_buffers = remote_edge_buffers - self.edge_buffer = [] - self.edge_list_dict = {} self.max_pending_edge_buffer_task = max_pending_edge_buffer_task self.num_edge_buffer_task_returns = num_edge_buffer_task_returns @@ -98,69 +194,57 @@ def flush_key_value_pairs(self): self.hash_table = {} def balanced_union_find(self): - for x, y in self.edge_buffer: - self.union(x, y) - self.edge_buffer = [] result_refs = [] for remote_edge_buffer in self.remote_edge_buffers: if len(result_refs) > self.max_pending_edge_buffer_task: ready_refs, result_refs = ray.wait(result_refs, num_returns=self.num_edge_buffer_task_returns) - edge_list = ray.get(ready_refs) - for edges in edge_list: - for x, y in edges: - self.union(x, y) + for edges in ray.get(ready_refs): + _union_edges(self, edges) del ready_refs result_refs.append(remote_edge_buffer.get_edges.remote(self.parallel_id)) - edge_list = ray.get(result_refs) - for edges in edge_list: - for x, y in edges: - self.union(x, y) - del edge_list, result_refs + for edges in ray.get(result_refs): + _union_edges(self, edges) + del result_refs self.rebalancing() return self.old_parent != self.parent - def distribute_edge(self, u, v): - hash_u = u // BATCH_SIZE % self.parallel_num - hash_v = v // BATCH_SIZE % self.parallel_num - if hash_u not in self.edge_list_dict: - self.edge_list_dict[hash_u] = [] - self.edge_list_dict[hash_u].append((u, v)) - if hash_u != hash_v: - if hash_v not in self.edge_list_dict: - self.edge_list_dict[hash_v] = [] - self.edge_list_dict[hash_v].append((u, v)) - - def set_edge_buffer(self): - if self.parallel_id in self.edge_list_dict: - self.edge_buffer = self.edge_list_dict[self.parallel_id] - del self.edge_list_dict[self.parallel_id] - else: - self.edge_buffer = [] - ray.get(self.remote_edge_buffers[self.parallel_id].set_edges.remote(self.edge_list_dict)) - self.edge_list_dict = {} + def set_edge_buffer(self, uids, parents): + edges, offsets = _group_edges_by_destination(uids, parents, self.parallel_num) + ray.get(self.remote_edge_buffers[self.parallel_id].set_edges.remote(edges, offsets)) def edge_redistribution(self): self.flush_key_value_pairs() self.rebalancing() - self.edge_list_dict = {} - for u, v in self.parent.items(): - self.distribute_edge(u, v) + uids, parents = _parent_to_arrays(self.parent) self.parent = {} - self.set_edge_buffer() + self.set_edge_buffer(uids, parents) def communication(self): - self.edge_list_dict = {} - del_list = [] + try: + edge_uids, edge_parents, deleted_uids = self._collect_communication_edges(UID_DTYPE) + except OverflowError: + edge_uids, edge_parents, deleted_uids = self._collect_communication_edges(object) + self.old_parent = self.parent.copy() + for u in deleted_uids: + del self.parent[int(u)] + self.set_edge_buffer(edge_uids, edge_parents) + + def _collect_communication_edges(self, dtype): + edge_uids = np.empty(len(self.parent), dtype=dtype) + edge_parents = np.empty(len(self.parent), dtype=dtype) + deleted_uids = np.empty(len(self.parent), dtype=dtype) + edge_count = 0 + deleted_count = 0 for u, v in self.parent.items(): hash_u = u // BATCH_SIZE % self.parallel_num - if self.parent[u] != self.old_parent.get(u, u) or (hash_u != self.parallel_id and v not in self.parent): - self.distribute_edge(u, v) + if v != self.old_parent.get(u, u) or (hash_u != self.parallel_id and v not in self.parent): + edge_uids[edge_count] = u + edge_parents[edge_count] = v + edge_count += 1 if hash_u != self.parallel_id: - del_list.append(u) - self.old_parent = self.parent.copy() - for u in del_list: - del self.parent[u] - self.set_edge_buffer() + deleted_uids[deleted_count] = u + deleted_count += 1 + return edge_uids[:edge_count], edge_parents[:edge_count], deleted_uids[:deleted_count] def find(self, x): if x not in self.parent: @@ -217,7 +301,6 @@ def squeeze(self): dup_keys = {x for x in self.parent if x // BATCH_SIZE % self.parallel_num == self.parallel_id} self.parent = dup_keys self.old_parent = {} - self.edge_buffer = [] ray.get(self.remote_edge_buffers[self.parallel_id].clear.remote()) def dup_idx(self, queries): diff --git a/tests/ops/deduplicator/test_ray_bts_minhash_deduplicator.py b/tests/ops/deduplicator/test_ray_bts_minhash_deduplicator.py index 7be6e393cf3..58d5c0b8056 100644 --- a/tests/ops/deduplicator/test_ray_bts_minhash_deduplicator.py +++ b/tests/ops/deduplicator/test_ray_bts_minhash_deduplicator.py @@ -1,12 +1,22 @@ -import unittest import os import shutil +import unittest +from unittest.mock import patch + +import numpy as np from data_juicer.core.data import NestedDataset as Dataset from data_juicer.ops.deduplicator.ray_bts_minhash_deduplicator import ( + EDGE_DTYPE, + BTSUnionFind, + EdgeBuffer, RayBTSMinhashDeduplicator, RayBTSMinhashDeduplicatorWithUid, + _group_edges_by_destination, + _parent_to_arrays, + _union_edges, + get_remote_classes, ) from data_juicer.ops.deduplicator.ray_bts_minhash_cpp_deduplicator import ( RayBTSMinhashCppDeduplicator, @@ -15,6 +25,142 @@ from data_juicer.utils.unittest_utils import DataJuicerTestCaseBase, TEST_TAG +class RayBTSMinhashStorageTest(unittest.TestCase): + + def test_compact_edge_layout_and_partitioning(self): + self.assertEqual(EDGE_DTYPE.itemsize, 16) + uids = np.array([-2000, 10, 1500, 3100], dtype=np.int64) + parents = np.array([10, 2500, 1500, -100], dtype=np.int64) + + edges, offsets = _group_edges_by_destination(uids, parents, parallel_num=3) + actual = {} + for destination in range(3): + start, end = offsets[destination : destination + 2] + actual[destination] = sorted( + zip(edges["u"][start:end].tolist(), edges["v"][start:end].tolist()) + ) + + expected = {destination: [] for destination in range(3)} + for u, v in zip(uids.tolist(), parents.tolist()): + hash_u = u // 1000 % 3 + hash_v = v // 1000 % 3 + expected[hash_u].append((u, v)) + if hash_u != hash_v: + expected[hash_v].append((u, v)) + expected = {destination: sorted(values) for destination, values in expected.items()} + + self.assertEqual(actual, expected) + + def test_edge_buffer_returns_each_partition_once(self): + uids = np.array([1, 1001, 2001], dtype=np.int64) + parents = np.array([2001, 1, 1001], dtype=np.int64) + edges, offsets = _group_edges_by_destination(uids, parents, parallel_num=3) + buffer = EdgeBuffer() + buffer.set_edges(edges, offsets) + + for destination in (2, 0, 1): + expected_size = int(offsets[destination + 1] - offsets[destination]) + self.assertEqual(len(buffer.get_edges(destination)), expected_size) + self.assertEqual(len(buffer.get_edges(destination)), 0) + + self.assertEqual(buffer.remaining, 0) + self.assertEqual(len(buffer.edges), 0) + + def test_compact_edges_preserve_connected_components(self): + union_find = BTSUnionFind(256, 3, 0, [], 20, 10) + edges = np.array([(5, 2), (2, -3), (9, 10)], dtype=EDGE_DTYPE) + _union_edges(union_find, edges, chunk_size=2) + + self.assertEqual(union_find.find(5), -3) + self.assertEqual(union_find.find(2), -3) + self.assertEqual(union_find.find(9), 9) + self.assertEqual(union_find.find(10), 9) + + def test_parent_arrays_are_signed_and_order_aligned(self): + uids, parents = _parent_to_arrays({-2: -3, 7: -2}) + + self.assertEqual(uids.dtype, np.dtype(np.int64)) + self.assertEqual(parents.dtype, np.dtype(np.int64)) + self.assertEqual(uids.tolist(), [-2, 7]) + self.assertEqual(parents.tolist(), [-3, -2]) + + def test_large_uids_fall_back_without_changing_values(self): + high_uid = 1 << 63 + low_uid = -(1 << 63) - 1001 + uids, parents = _parent_to_arrays({high_uid: low_uid}) + + self.assertTrue(uids.dtype.hasobject) + self.assertTrue(parents.dtype.hasobject) + edges, offsets = _group_edges_by_destination(uids, parents, parallel_num=3) + self.assertTrue(edges.dtype.hasobject) + self.assertEqual(offsets[-1], 2) + self.assertEqual(set(zip(edges["u"].tolist(), edges["v"].tolist())), {(high_uid, low_uid)}) + + def test_large_uids_fall_back_during_bts_phases(self): + high_uid = 1 << 63 + union_find = BTSUnionFind(256, 2, 0, [], 20, 10) + union_find.add_key_value_pairs([(b"same", high_uid), (b"same", high_uid + 1000)]) + with patch.object(union_find, "set_edge_buffer") as set_edge_buffer: + union_find.edge_redistribution() + uids, parents = set_edge_buffer.call_args.args + self.assertTrue(uids.dtype.hasobject) + self.assertEqual(list(zip(uids.tolist(), parents.tolist())), [(high_uid + 1000, high_uid)]) + + union_find.parent = {high_uid + 1000: high_uid} + union_find.old_parent = {} + with patch.object(union_find, "set_edge_buffer") as set_edge_buffer: + union_find.communication() + uids, parents = set_edge_buffer.call_args.args + self.assertTrue(uids.dtype.hasobject) + self.assertEqual(list(zip(uids.tolist(), parents.tolist())), [(high_uid + 1000, high_uid)]) + + @TEST_TAG("ray") + def test_compact_edges_merge_across_ray_actors(self): + import ray + + if not ray.is_initialized(): + ray.init("auto", ignore_reinit_error=True) + + remote_classes = get_remote_classes() + for uid_base in (0, 1 << 63): + edge_buffers = [remote_classes["EdgeBuffer"].remote() for _ in range(2)] + union_finds = [ + remote_classes["BTSUnionFind"].remote(256, 2, actor_id, edge_buffers, 20, 10) + for actor_id in range(2) + ] + try: + ray.get( + [ + union_finds[0].add_key_value_pairs.remote( + [(b"a", uid_base), (b"a", uid_base + 1000)] + ) + ] + + [ + union_finds[1].add_key_value_pairs.remote( + [(b"b", uid_base + 1000), (b"b", uid_base + 2000)] + ) + ] + ) + ray.get([union_find.edge_redistribution.remote() for union_find in union_finds]) + while any(ray.get([union_find.balanced_union_find.remote() for union_find in union_finds])): + ray.get([union_find.communication.remote() for union_find in union_finds]) + ray.get([union_find.squeeze.remote() for union_find in union_finds]) + + queries = [(uid_base + offset, index) for index, offset in enumerate((0, 1000, 2000))] + duplicate_indices = ray.get( + [ + union_find.dup_idx.remote( + [(uid, index) for uid, index in queries if uid // 1000 % 2 == actor_id] + ) + for actor_id, union_find in enumerate(union_finds) + ] + ) + self.assertEqual(sorted(duplicate_indices[0] + duplicate_indices[1]), [1, 2]) + finally: + for actor in union_finds + edge_buffers: + ray.kill(actor) + + class RayBTSMinhashDeduplicatorTest(DataJuicerTestCaseBase): def _run_minhash_dedup(self, dataset: Dataset, target_list, op):