diff --git a/app.py b/app.py index 26e72416202..8acf3388440 100644 --- a/app.py +++ b/app.py @@ -198,6 +198,19 @@ def get_min_max_step(data): return min_value, max_value, step +def combine_conds(conds, num_samples): + """Combine the per-OP boolean conditions into a single mask over samples. + + ``conds`` is empty whenever none of the OPs in the config is covered by + ``op_stats_dict``. ``np.all`` over an empty list collapses to a scalar, + which ``DataFrame.loc`` then reads as a label instead of a boolean mask and + raises ``KeyError: True``, so keep all the samples explicitly in that case. + """ + if not conds: + return np.ones(num_samples, dtype=bool) + return np.all([list(cond.values())[0] for cond in conds], axis=0) + + op_stats_dict = { "alphanumeric_filter": [StatsKeys.alpha_token_ratio, StatsKeys.alnum_ratio], "average_line_length_filter": [StatsKeys.avg_line_length], @@ -336,7 +349,12 @@ def set_sliders(total_stats, ordered): if ordered: all_conds = [True if i in filtered_stats.index else False for i in range(len(stats))] else: - all_conds = np.all([list(cond.values())[0] for cond in conds], axis=0) + if not conds: + st.warning( + "None of the OPs in this config provides stats that this demo can tune, " + "so no sample is filtered out below." + ) + all_conds = combine_conds(conds, len(stats)) ds = pd.DataFrame(dataset) Visualize.display_dataset(ds, all_conds, show_num, "Retained samples", "docs") st.download_button( diff --git a/tests/tools/test_app.py b/tests/tools/test_app.py new file mode 100644 index 00000000000..e42a0a0d8b3 --- /dev/null +++ b/tests/tools/test_app.py @@ -0,0 +1,44 @@ +import importlib.util +import unittest +from pathlib import Path + +import numpy as np +import pandas as pd + +from data_juicer.utils.unittest_utils import DataJuicerTestCaseBase + +ROOT = Path(__file__).resolve().parents[2] + + +class AppTest(DataJuicerTestCaseBase): + def setUp(self): + super().setUp() + spec = importlib.util.spec_from_file_location("dj_app", ROOT / "app.py") + self.app = importlib.util.module_from_spec(spec) + spec.loader.exec_module(self.app) + + def test_combine_conds_without_cond(self): + # None of the OPs in the config is covered by op_stats_dict, so no + # condition is collected. The combined mask must still cover every + # sample: a scalar would be read as a label by DataFrame.loc and raise + # KeyError instead of selecting rows. + dataframe = pd.DataFrame({"text": ["a", "b", "c"]}) + all_conds = self.app.combine_conds([], len(dataframe)) + self.assertEqual(len(all_conds), len(dataframe)) + self.assertTrue(np.all(all_conds)) + self.assertEqual(len(dataframe.loc[all_conds]), 3) + self.assertEqual(len(dataframe.loc[np.invert(all_conds)]), 0) + + def test_combine_conds_with_conds(self): + dataframe = pd.DataFrame({"text": ["a", "b", "c"]}) + conds = [ + {("1 text_length_filter", "text_len"): [True, True, False]}, + {("2 words_num_filter", "num_words"): [True, False, True]}, + ] + all_conds = self.app.combine_conds(conds, len(dataframe)) + self.assertEqual(list(all_conds), [True, False, False]) + self.assertEqual(len(dataframe.loc[all_conds]), 1) + + +if __name__ == "__main__": + unittest.main()