From 81f2ff7d754d4d83d7aa9844b166e0ac97f58030 Mon Sep 17 00:00:00 2001 From: ji11dube Date: Tue, 22 Sep 2026 21:18:37 +0200 Subject: [PATCH] Fix default workload for two-column datasets --- dpsynth/discrete_mechanisms/common.py | 9 ++++++--- tests/discrete_mechanisms/common_test.py | 12 ++++++++++++ 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/dpsynth/discrete_mechanisms/common.py b/dpsynth/discrete_mechanisms/common.py index fc7d4ad..998b604 100644 --- a/dpsynth/discrete_mechanisms/common.py +++ b/dpsynth/discrete_mechanisms/common.py @@ -413,14 +413,16 @@ def supporting_cliques( Args: domain: The domain of the dataset. - workload: A workload specification. Defaults to all three-way marginals. + workload: A workload specification. Defaults to all three-way marginals, + or the maximum possible degree for fewer than three columns. max_marginal_size: The maximum domain size of a clique to include. Returns: A list of cliques from the workload whose domain size is within the limit. """ if workload is None: - cliques = list(itertools.combinations(domain.attributes, 3)) + degree = min(3, len(domain.attributes)) + cliques = list(itertools.combinations(domain.attributes, degree)) elif isinstance(workload, Mapping): cliques = [tuple(cl) for cl in workload.keys()] else: @@ -458,7 +460,8 @@ def compiled_workload( """ if workload is None: - workload = list(itertools.combinations(domain.attributes, 3)) + degree = min(3, len(domain.attributes)) + workload = list(itertools.combinations(domain.attributes, degree)) if not isinstance(workload, Mapping): workload = {tuple(cl): 1.0 for cl in workload} diff --git a/tests/discrete_mechanisms/common_test.py b/tests/discrete_mechanisms/common_test.py index 5e86697..433cd01 100644 --- a/tests/discrete_mechanisms/common_test.py +++ b/tests/discrete_mechanisms/common_test.py @@ -80,6 +80,10 @@ def test_supporting_cliques(self): self.assertCountEqual( cliques, list(itertools.combinations(domain.attributes, 3)) ) + # Default workload with two columns. + two_column_domain = mbi.Domain(["a", "b"], [3, 3]) + cliques = common.supporting_cliques(two_column_domain, workload=None) + self.assertCountEqual(cliques, [("a", "b")]) # List workload. cliques = common.supporting_cliques(domain, [("a", "b"), ("c", "d")]) self.assertCountEqual(cliques, [("a", "b"), ("c", "d")]) @@ -105,6 +109,14 @@ def test_compiled_workload_with_lists(self): for cl in workload.keys(): self.assertIsInstance(cl, tuple) + def test_compiled_workload_default_with_two_columns(self): + domain = mbi.Domain(["a", "b"], [3, 3]) + workload = common.compiled_workload(domain, None) + self.assertCountEqual( + workload.keys(), + [("a",), ("b",), ("a", "b")], + ) + def test_precompute_marginals_standard_and_jax(self): domain = mbi.Domain(["a", "b"], [3, 4]) data = mbi.Dataset.synthetic(domain, N=100)