From f63936b0ffbba1be22378d31ceaa3e88cc7ee329 Mon Sep 17 00:00:00 2001 From: dajiaohuang Date: Thu, 24 Sep 2026 03:47:05 +0800 Subject: [PATCH] Fix dataset diversity access --- gliclass/data_processing.py | 2 +- tests/test_data_processing.py | 13 ++++++++++++- 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/gliclass/data_processing.py b/gliclass/data_processing.py index f9bfc12..212472d 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 d71d7b0..78683d3 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 pad_2d_tensor +from gliclass.data_processing import 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] class TestPad2DTensor: