|
16 | 16 | from collections import Counter |
17 | 17 | from unittest import mock |
18 | 18 |
|
| 19 | +import numpy as np |
19 | 20 | import pyarrow as pa |
20 | 21 | from parameterized import parameterized |
21 | 22 | from pyarrow import csv, parquet |
|
25 | 26 | CollisionResolutionRunner, |
26 | 27 | ResolveSidCollisionsConfig, |
27 | 28 | ) |
| 29 | +from tzrec.utils.sid.collision import stable_order_hash |
28 | 30 | from tzrec.utils.test_util import make_test_dir, parameterized_name_func |
29 | 31 |
|
30 | 32 |
|
31 | | -def _parquet(path, item_ids, codes, candidate_codes=None): |
| 33 | +def _parquet( |
| 34 | + path, |
| 35 | + item_ids, |
| 36 | + codes, |
| 37 | + candidate_codes=None, |
| 38 | + item_id_type=None, |
| 39 | +): |
| 40 | + if item_id_type is None: |
| 41 | + item_id_type = pa.int64() |
32 | 42 | cols = { |
33 | | - "item_id": pa.array(item_ids, type=pa.int64()), |
| 43 | + "item_id": pa.array(item_ids, type=item_id_type), |
34 | 44 | "codes": pa.array(codes, type=pa.list_(pa.int64())), |
35 | 45 | } |
36 | 46 | if candidate_codes is not None: |
@@ -824,12 +834,91 @@ def test_empty_input_raises(self) -> None: |
824 | 834 | with self.assertRaisesRegex(ValueError, "SID input is empty"): |
825 | 835 | self._run(inp, out, max_items_per_codebook=2) |
826 | 836 |
|
827 | | - def test_duplicate_item_id_raises(self) -> None: |
| 837 | + @parameterized.expand( |
| 838 | + [("same_batch", 100000), ("across_batches", 1)], |
| 839 | + name_func=parameterized_name_func, |
| 840 | + ) |
| 841 | + def test_candidate_tolerates_duplicate_item_ids( |
| 842 | + self, _case_name, batch_size |
| 843 | + ) -> None: |
828 | 844 | inp = os.path.join(self.test_dir, "in.parquet") |
829 | 845 | out = os.path.join(self.test_dir, "out") |
830 | | - _parquet(inp, [0, 1, 1], [[0, 0], [0, 1], [0, 2]]) |
831 | | - with self.assertRaisesRegex(ValueError, "item IDs must be unique"): |
832 | | - self._run(inp, out, max_items_per_codebook=2) |
| 846 | + distinct_ids = np.asarray(["a", "b"], dtype=object) |
| 847 | + hash_order = np.argsort(stable_order_hash(distinct_ids)) |
| 848 | + keeper = distinct_ids[hash_order[0]] |
| 849 | + duplicate = distinct_ids[hash_order[1]] |
| 850 | + _parquet( |
| 851 | + inp, |
| 852 | + [keeper, duplicate, duplicate], |
| 853 | + [[0, 0]] * 3, |
| 854 | + [ |
| 855 | + [[0, 1], [0, 2]], |
| 856 | + [[0, 1], [0, 2]], |
| 857 | + [[0, 2], [0, 3]], |
| 858 | + ], |
| 859 | + item_id_type=pa.string(), |
| 860 | + ) |
| 861 | + |
| 862 | + stats = self._run( |
| 863 | + inp, |
| 864 | + out, |
| 865 | + batch_size=batch_size, |
| 866 | + include_original=True, |
| 867 | + max_items_per_codebook=1, |
| 868 | + ) |
| 869 | + |
| 870 | + self.assertEqual(stats.total_items, 3) |
| 871 | + self.assertEqual(stats.relocated_count, 2) |
| 872 | + result = self._read_parquet(out) |
| 873 | + self.assertEqual(result["item_id"], [keeper, duplicate, duplicate]) |
| 874 | + self.assertEqual(result["item_id"].count(duplicate), 2) |
| 875 | + self._assert_map_matches_resolved_groups(out) |
| 876 | + original_path, _ = self._group_paths(out) |
| 877 | + original_groups = self._read_parquet(original_path) |
| 878 | + self.assertCountEqual( |
| 879 | + [item_id for group in original_groups["itemids"] for item_id in group], |
| 880 | + [keeper, duplicate, duplicate], |
| 881 | + ) |
| 882 | + |
| 883 | + def test_item_id_lookup_broadcasts_all_duplicate_targets(self) -> None: |
| 884 | + lookup = resolve_sid_collisions._ItemIdLookup( |
| 885 | + np.asarray(["b", "a", "b", "b"], dtype=object) |
| 886 | + ) |
| 887 | + source_rows, target_rows = lookup.match( |
| 888 | + np.asarray(["b", "missing", "a"], dtype=object) |
| 889 | + ) |
| 890 | + np.testing.assert_array_equal(source_rows, [0, 2]) |
| 891 | + np.testing.assert_array_equal(target_rows, [0, 1]) |
| 892 | + |
| 893 | + values = np.asarray([[10, 11], [20, 21], [-1, -1], [-1, -1]]) |
| 894 | + lookup.broadcast_duplicate_targets(values) |
| 895 | + np.testing.assert_array_equal( |
| 896 | + values, |
| 897 | + [[10, 11], [20, 21], [10, 11], [10, 11]], |
| 898 | + ) |
| 899 | + |
| 900 | + def test_random_tolerates_duplicate_item_ids(self) -> None: |
| 901 | + inp = os.path.join(self.test_dir, "in.parquet") |
| 902 | + out = os.path.join(self.test_dir, "out") |
| 903 | + _parquet( |
| 904 | + inp, |
| 905 | + ["b", "a", "a"], |
| 906 | + [[0, 0]] * 3, |
| 907 | + item_id_type=pa.string(), |
| 908 | + ) |
| 909 | + |
| 910 | + stats = self._run( |
| 911 | + inp, |
| 912 | + out, |
| 913 | + strategy="random", |
| 914 | + random_num_candidates=8, |
| 915 | + max_items_per_codebook=1, |
| 916 | + ) |
| 917 | + |
| 918 | + self.assertEqual(stats.total_items, 3) |
| 919 | + result = self._read_parquet(out) |
| 920 | + self.assertEqual(result["item_id"], ["b", "a", "a"]) |
| 921 | + self._assert_map_matches_resolved_groups(out) |
833 | 922 |
|
834 | 923 | def test_empty_codebook_token_raises(self) -> None: |
835 | 924 | args = resolve_sid_collisions.build_parser().parse_args( |
|
0 commit comments