diff --git a/gliclass/data_processing.py b/gliclass/data_processing.py index ada835f..12569b4 100644 --- a/gliclass/data_processing.py +++ b/gliclass/data_processing.py @@ -256,7 +256,7 @@ def __init__( print("Total labels: ", len(self.dataset_labels)) def get_diversity(self): - return [item.get("_diversity", {}).get("overall_diversity", 0.5) for item in self.data] + return [item.get("_diversity", {}).get("overall_diversity", 0.5) for item in self._data] def collect_dataset_labels(self): dataset_labels = set() diff --git a/tests/test_data_processing.py b/tests/test_data_processing.py index 4e2640d..b526c44 100644 --- a/tests/test_data_processing.py +++ b/tests/test_data_processing.py @@ -3,7 +3,18 @@ import pytest import torch -from gliclass.data_processing import DataCollatorWithPadding, pad_2d_tensor +from gliclass.data_processing import DataCollatorWithPadding, GLiClassDataset, pad_2d_tensor + + +def test_dataset_get_diversity_reads_examples(): + dataset = object.__new__(GLiClassDataset) + dataset._data = [ + {"_diversity": {"overall_diversity": 0.75}}, + {}, + {"_diversity": {"overall_diversity": 0.25}}, + ] + + assert dataset.get_diversity() == [0.75, 0.5, 0.25] def test_collator_stacks_scalar_labels_for_single_label_classification():