Skip to content
Open
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
5 changes: 4 additions & 1 deletion src/datasets/formatting/np_formatter.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,10 @@ def __init__(self, features=None, token_per_repo_id=None, **np_array_kwargs):
def _consolidate(self, column):
if isinstance(column, list):
if column and all(
isinstance(x, np.ndarray) and x.shape == column[0].shape and x.dtype == column[0].dtype for x in column
isinstance(x, (np.ndarray, np.number, np.bool_))
and x.shape == column[0].shape
and x.dtype == column[0].dtype
for x in column
):
return np.stack(column)
else:
Expand Down
5 changes: 4 additions & 1 deletion src/datasets/iterable_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -4583,7 +4583,10 @@ def batch(
.with_format(self._formatting.format_type if self._formatting else None)
)
return ds
if self._formatting and self._formatting.is_table:
if self._formatting and (self._ex_iterable.iter_arrow is not None or self._formatting.is_table):
# Batch in arrow when the underlying iterable is arrow-backed (or the output format is a table).
# This avoids converting arrow tables to Python examples only to transpose them back, and matches
# both `IterableDataset.iter(batch_size=...)` and map-style `Dataset.batch()`.
return (
self.with_format("arrow")
.map(
Expand Down
130 changes: 130 additions & 0 deletions tests/test_iterable_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -3412,6 +3412,136 @@ def test_iterable_dataset_batch_with_polars_format():
]


def _indexed_iterable(n=12, num_shards=3):
return Dataset.from_dict(
{
"i": list(range(n)),
"f": [j / 2 for j in range(n)],
"s": [f"r{j}" for j in range(n)],
}
).to_iterable_dataset(num_shards=num_shards)


@pytest.mark.parametrize("batch_size", [2, 5])
def test_iterable_dataset_batch_numpy_format_keeps_dtypes(batch_size):
"""`.batch(n)` must agree with `.iter(batch_size=n)`, which groups the same rows the same way.

Batching through a Python transpose hands `NumpyFormatter._consolidate` a list of numpy
scalars. Those are `np.number` rather than `np.ndarray`, so the `np.stack` branch never
fires and every column came back as `dtype=object`.
"""
ds = _indexed_iterable().with_format("numpy")

batched = list(ds.batch(batch_size=batch_size))
oracle = list(ds.iter(batch_size=batch_size))

assert len(batched) == len(oracle)
for got, want in zip(batched, oracle):
assert sorted(got) == sorted(want)
for column in ("i", "f"):
assert got[column].dtype == want[column].dtype, column
assert got[column].dtype != object, column
assert got[column].shape == want[column].shape, column
assert np.array_equal(got[column], want[column]), column


def test_iterable_dataset_batch_numpy_format_ragged_last_batch():
ds = _indexed_iterable(n=10).with_format("numpy")

batched = list(ds.batch(batch_size=4))

assert [len(batch["i"]) for batch in batched] == [4, 4, 2]
assert all(batch["i"].dtype != object for batch in batched)
assert np.array_equal(np.concatenate([batch["i"] for batch in batched]), np.arange(10))


def test_iterable_dataset_batch_numpy_format_drop_last_batch():
ds = _indexed_iterable(n=10).with_format("numpy")

batched = list(ds.batch(batch_size=4, drop_last_batch=True))

assert [len(batch["i"]) for batch in batched] == [4, 4]
assert all(batch["i"].dtype != object for batch in batched)


@require_torch
@pytest.mark.parametrize("batch_size", [2, 5])
def test_iterable_dataset_batch_torch_format_matches_iter(batch_size):
import torch

ds = _indexed_iterable().with_format("torch")

batched = list(ds.batch(batch_size=batch_size))
oracle = list(ds.iter(batch_size=batch_size))

assert len(batched) == len(oracle)
for got, want in zip(batched, oracle):
for column in ("i", "f"):
assert got[column].dtype == want[column].dtype, column
assert torch.equal(got[column], want[column]), column


def test_iterable_dataset_batch_without_format_is_unchanged():
"""An unformatted dataset has no formatter to mistype, so it keeps the Python transpose."""
ds = _indexed_iterable(n=6, num_shards=2)

batched = list(ds.batch(batch_size=2))

assert len(batched) == 3
assert all(isinstance(batch["i"], list) for batch in batched)
assert [batch["i"] for batch in batched] == [[0, 1], [2, 3], [4, 5]]


def test_iterable_dataset_batch_numpy_format_survives_state_dict_roundtrip():
ds = _indexed_iterable(n=12).with_format("numpy").batch(batch_size=4)

it = iter(ds)
next(it)
state = ds.state_dict()

resumed = _indexed_iterable(n=12).with_format("numpy").batch(batch_size=4)
resumed.load_state_dict(state)
rest = list(resumed)

assert [len(batch["i"]) for batch in rest] == [4, 4]
assert all(batch["i"].dtype != object for batch in rest)
assert np.array_equal(np.concatenate([batch["i"] for batch in rest]), np.arange(4, 12))


def test_iterable_dataset_batch_numpy_format_after_python_map_keeps_dtypes():
ds = _indexed_iterable(n=6).with_format("numpy").map(lambda x: {"i": x["i"] + 1, "b": x["i"] > 2})

batched = list(ds.batch(batch_size=3))
oracle = list(ds.iter(batch_size=3))

assert len(batched) == len(oracle) == 2
for got, want in zip(batched, oracle):
for column in ("i", "b"):
assert got[column].dtype == want[column].dtype, column
assert got[column].dtype != object, column
assert np.array_equal(got[column], want[column]), column


@pytest.mark.parametrize("fmt", ["numpy", "torch"])
def test_iterable_dataset_batch_formatted_python_iterable_nullable_and_empty_lists(fmt):
if fmt == "torch" and not config.TORCH_AVAILABLE:
pytest.skip("requires torch")
data = [
{"x": None, "lst": [], "y": 1},
{"x": 10, "lst": [1, 2], "y": 2},
]
ds = IterableDataset.from_generator(lambda: (r for r in data)).with_format(fmt)

batched = list(ds.batch(batch_size=2))
oracle = list(ds.iter(batch_size=2))

assert len(batched) == len(oracle) == 1
assert batched[0]["x"][0] is None
assert int(batched[0]["x"][1]) == 10
assert list(batched[0]["lst"][0]) == []
assert [int(v) for v in batched[0]["lst"][1]] == [1, 2]


@dataclass
class DecodableFeature:
decode_example_num_calls = 0
Expand Down
Loading