diff --git a/src/plaid/storage/hf_datasets/writer.py b/src/plaid/storage/hf_datasets/writer.py index 77276c2b..df92818b 100644 --- a/src/plaid/storage/hf_datasets/writer.py +++ b/src/plaid/storage/hf_datasets/writer.py @@ -14,10 +14,12 @@ import gc import logging import tempfile +from contextlib import contextmanager from pathlib import Path from typing import Any, Callable, Generator, Optional, Union import yaml +from datasets.utils import logging as datasets_logging from huggingface_hub import DatasetCard, hf_hub_download from plaid.storage.hf_datasets.bridge import generator_to_datasetdict @@ -29,6 +31,23 @@ logger = logging.getLogger(__name__) +@contextmanager +def _hf_progress_bars(enabled: bool): + """Temporarily configure Hugging Face progress bars.""" + was_enabled = datasets_logging.is_progress_bar_enabled() + if enabled: + datasets_logging.enable_progress_bar() + else: + datasets_logging.disable_progress_bar() + try: + yield + finally: + if was_enabled: + datasets_logging.enable_progress_bar() + else: + datasets_logging.disable_progress_bar() + + def _compute_num_shards(hf_dataset_dict: Any) -> dict[str, int]: """Computes the number of shards for each split in a DatasetDict. @@ -76,15 +95,18 @@ def save_datasetdict_to_disk( None """ num_shards = _compute_num_shards(hf_datasetdict) - num_proc = kwargs.get("num_proc", None) - if num_proc is not None: # pragma: no cover + requested_num_proc = kwargs.pop("num_proc", None) + if requested_num_proc is None or requested_num_proc <= 1: + num_proc = None + else: min_num_shards = min(num_shards.values()) - if min_num_shards < num_proc: + effective_num_proc = min(requested_num_proc, min_num_shards) + num_proc = effective_num_proc if effective_num_proc > 1 else None + if effective_num_proc < requested_num_proc: logger.warning( - f"num_proc changed from {num_proc} to 1 to safely adapt for num_shards={num_shards}" + f"num_proc changed from {requested_num_proc} to " + f"{effective_num_proc} to safely adapt for num_shards={num_shards}" ) - num_proc = 1 - del kwargs["num_proc"] hf_datasetdict.save_to_disk( str(Path(path) / "data"), num_shards=num_shards, num_proc=num_proc, **kwargs @@ -97,7 +119,7 @@ def generate_datasetdict_to_disk( variable_schema: dict[str, dict], gen_kwargs: Optional[dict[str, dict[str, Any]]] = None, num_proc: int = 1, - verbose: bool = False, # noqa: ARG001 + verbose: bool = False, ) -> None: """Generates and saves a DatasetDict to disk from sample generators. @@ -111,7 +133,7 @@ def generate_datasetdict_to_disk( num_proc (int): Number of processes for generation. verbose (bool): Whether to enable verbose output. """ - with tempfile.TemporaryDirectory() as tmpdirname: + with _hf_progress_bars(verbose), tempfile.TemporaryDirectory() as tmpdirname: hf_datasetdict = generator_to_datasetdict( generators, variable_schema, diff --git a/tests/storage/test_hf_datasets_writer.py b/tests/storage/test_hf_datasets_writer.py new file mode 100644 index 00000000..1cdcd99d --- /dev/null +++ b/tests/storage/test_hf_datasets_writer.py @@ -0,0 +1,65 @@ +from unittest.mock import MagicMock + +import pytest +from datasets.utils import logging as datasets_logging + +from plaid.storage.hf_datasets import writer + + +@pytest.mark.parametrize( + ("requested_num_proc", "dataset_size", "expected_num_proc"), + [ + (None, 1, None), + (1, 1, None), + (4, 1, None), + (4, 2 * 500 * 1024 * 1024, 2), + (2, 4 * 500 * 1024 * 1024, 2), + ], +) +def test_save_datasetdict_adapts_parallelism( + tmp_path, requested_num_proc, dataset_size, expected_num_proc +): + datasetdict = MagicMock() + dataset = MagicMock() + dataset.__len__.return_value = 10 + dataset.data.nbytes = dataset_size + datasetdict.items.return_value = [("train", dataset)] + + writer.save_datasetdict_to_disk(tmp_path, datasetdict, num_proc=requested_num_proc) + + datasetdict.save_to_disk.assert_called_once_with( + str(tmp_path / "data"), + num_shards={"train": min(10, max(1, dataset_size // (500 * 1024 * 1024)))}, + num_proc=expected_num_proc, + ) + + +@pytest.mark.parametrize("initially_enabled", [False, True]) +@pytest.mark.parametrize("verbose", [False, True]) +def test_generate_controls_and_restores_hf_progress( + monkeypatch, tmp_path, initially_enabled, verbose +): + datasetdict = MagicMock() + observed = [] + + def fake_generate(*_args, **_kwargs): + observed.append(datasets_logging.is_progress_bar_enabled()) + return datasetdict + + monkeypatch.setattr(writer, "generator_to_datasetdict", fake_generate) + monkeypatch.setattr(writer, "save_datasetdict_to_disk", MagicMock()) + + if initially_enabled: + datasets_logging.enable_progress_bar() + else: + datasets_logging.disable_progress_bar() + + writer.generate_datasetdict_to_disk( + tmp_path, + generators={"train": MagicMock()}, + variable_schema={}, + verbose=verbose, + ) + + assert observed == [verbose] + assert datasets_logging.is_progress_bar_enabled() is initially_enabled