From ea5ff7d089da3a1b5fa260e5a90254f00fc18975 Mon Sep 17 00:00:00 2001 From: Fabien Casenave Date: Sat, 1 Aug 2026 13:01:55 +0200 Subject: [PATCH 1/3] :bug: (storage/hf_datasets/writer) fix progress bar --- src/plaid/storage/hf_datasets/writer.py | 41 +++++++++++++++------- tests/storage/test_hf_datasets_writer.py | 44 ++++++++++++++++++++++++ 2 files changed, 72 insertions(+), 13 deletions(-) create mode 100644 tests/storage/test_hf_datasets_writer.py diff --git a/src/plaid/storage/hf_datasets/writer.py b/src/plaid/storage/hf_datasets/writer.py index 77276c2b..314af91b 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,18 +95,14 @@ 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 - min_num_shards = min(num_shards.values()) - if min_num_shards < num_proc: - logger.warning( - f"num_proc changed from {num_proc} to 1 to safely adapt for num_shards={num_shards}" - ) - num_proc = 1 - del kwargs["num_proc"] + # Do not pass ``num_proc=1`` here. Hugging Face treats every non-None value + # as a request to spawn a process pool, and progress updates from that pool + # can arrive only when a complete shard has been written. The in-process + # path reports each Arrow write batch and therefore advances continuously. + kwargs.pop("num_proc", None) hf_datasetdict.save_to_disk( - str(Path(path) / "data"), num_shards=num_shards, num_proc=num_proc, **kwargs + str(Path(path) / "data"), num_shards=num_shards, num_proc=None, **kwargs ) @@ -97,7 +112,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 +126,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, @@ -119,7 +134,7 @@ def generate_datasetdict_to_disk( gen_kwargs=gen_kwargs, processes_number=num_proc, ) - save_datasetdict_to_disk(output_folder, hf_datasetdict, num_proc=num_proc) + save_datasetdict_to_disk(output_folder, hf_datasetdict) del hf_datasetdict gc.collect() diff --git a/tests/storage/test_hf_datasets_writer.py b/tests/storage/test_hf_datasets_writer.py new file mode 100644 index 00000000..155e132b --- /dev/null +++ b/tests/storage/test_hf_datasets_writer.py @@ -0,0 +1,44 @@ +from unittest.mock import MagicMock + +import pytest +from datasets.utils import logging as datasets_logging + +from plaid.storage.hf_datasets import writer + + +def test_save_datasetdict_uses_in_process_progress(tmp_path): + datasetdict = MagicMock() + dataset = MagicMock() + dataset.__len__.return_value = 10 + dataset.data.nbytes = 1 + datasetdict.items.return_value = [("train", dataset)] + + writer.save_datasetdict_to_disk(tmp_path, datasetdict, num_proc=4) + + datasetdict.save_to_disk.assert_called_once_with( + str(tmp_path / "data"), num_shards={"train": 1}, num_proc=None + ) + + +@pytest.mark.parametrize("verbose", [False, True]) +def test_generate_controls_and_restores_hf_progress(monkeypatch, tmp_path, 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()) + + initially_enabled = datasets_logging.is_progress_bar_enabled() + 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 From dd82d747f68b3df969135bc160f14a8a87b495e1 Mon Sep 17 00:00:00 2001 From: Fabien Casenave Date: Sat, 1 Aug 2026 13:10:55 +0200 Subject: [PATCH 2/3] wip --- tests/storage/test_hf_datasets_writer.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/tests/storage/test_hf_datasets_writer.py b/tests/storage/test_hf_datasets_writer.py index 155e132b..c7864446 100644 --- a/tests/storage/test_hf_datasets_writer.py +++ b/tests/storage/test_hf_datasets_writer.py @@ -20,8 +20,11 @@ def test_save_datasetdict_uses_in_process_progress(tmp_path): ) +@pytest.mark.parametrize("initially_enabled", [False, True]) @pytest.mark.parametrize("verbose", [False, True]) -def test_generate_controls_and_restores_hf_progress(monkeypatch, tmp_path, verbose): +def test_generate_controls_and_restores_hf_progress( + monkeypatch, tmp_path, initially_enabled, verbose +): datasetdict = MagicMock() observed = [] @@ -32,7 +35,11 @@ def fake_generate(*_args, **_kwargs): monkeypatch.setattr(writer, "generator_to_datasetdict", fake_generate) monkeypatch.setattr(writer, "save_datasetdict_to_disk", MagicMock()) - initially_enabled = datasets_logging.is_progress_bar_enabled() + if initially_enabled: + datasets_logging.enable_progress_bar() + else: + datasets_logging.disable_progress_bar() + writer.generate_datasetdict_to_disk( tmp_path, generators={"train": MagicMock()}, From 52d1d93364792d6ed7d0aab6bf1f22691341f105 Mon Sep 17 00:00:00 2001 From: Fabien Casenave Date: Sat, 1 Aug 2026 13:27:09 +0200 Subject: [PATCH 3/3] fix parallel writes --- src/plaid/storage/hf_datasets/writer.py | 21 ++++++++++++++------- tests/storage/test_hf_datasets_writer.py | 22 ++++++++++++++++++---- 2 files changed, 32 insertions(+), 11 deletions(-) diff --git a/src/plaid/storage/hf_datasets/writer.py b/src/plaid/storage/hf_datasets/writer.py index 314af91b..df92818b 100644 --- a/src/plaid/storage/hf_datasets/writer.py +++ b/src/plaid/storage/hf_datasets/writer.py @@ -95,14 +95,21 @@ def save_datasetdict_to_disk( None """ num_shards = _compute_num_shards(hf_datasetdict) - # Do not pass ``num_proc=1`` here. Hugging Face treats every non-None value - # as a request to spawn a process pool, and progress updates from that pool - # can arrive only when a complete shard has been written. The in-process - # path reports each Arrow write batch and therefore advances continuously. - kwargs.pop("num_proc", None) + 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()) + 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 {requested_num_proc} to " + f"{effective_num_proc} to safely adapt for num_shards={num_shards}" + ) hf_datasetdict.save_to_disk( - str(Path(path) / "data"), num_shards=num_shards, num_proc=None, **kwargs + str(Path(path) / "data"), num_shards=num_shards, num_proc=num_proc, **kwargs ) @@ -134,7 +141,7 @@ def generate_datasetdict_to_disk( gen_kwargs=gen_kwargs, processes_number=num_proc, ) - save_datasetdict_to_disk(output_folder, hf_datasetdict) + save_datasetdict_to_disk(output_folder, hf_datasetdict, num_proc=num_proc) del hf_datasetdict gc.collect() diff --git a/tests/storage/test_hf_datasets_writer.py b/tests/storage/test_hf_datasets_writer.py index c7864446..1cdcd99d 100644 --- a/tests/storage/test_hf_datasets_writer.py +++ b/tests/storage/test_hf_datasets_writer.py @@ -6,17 +6,31 @@ from plaid.storage.hf_datasets import writer -def test_save_datasetdict_uses_in_process_progress(tmp_path): +@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 = 1 + dataset.data.nbytes = dataset_size datasetdict.items.return_value = [("train", dataset)] - writer.save_datasetdict_to_disk(tmp_path, datasetdict, num_proc=4) + 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": 1}, num_proc=None + str(tmp_path / "data"), + num_shards={"train": min(10, max(1, dataset_size // (500 * 1024 * 1024)))}, + num_proc=expected_num_proc, )