Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 30 additions & 8 deletions src/plaid/storage/hf_datasets/writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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.

Expand Down Expand Up @@ -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
Expand All @@ -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.

Expand All @@ -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,
Expand Down
65 changes: 65 additions & 0 deletions tests/storage/test_hf_datasets_writer.py
Original file line number Diff line number Diff line change
@@ -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
Loading