Skip to content

Batch in arrow when the source is arrow-backed, and let _consolidate type numpy scalars - #8665

Open
Yonghui-Lee wants to merge 1 commit into
huggingface:mainfrom
Yonghui-Lee:fix/ah-batch-arrow
Open

Yonghui-Lee wants to merge 1 commit into
huggingface:mainfrom
Yonghui-Lee:fix/ah-batch-arrow

Conversation

@Yonghui-Lee

Copy link
Copy Markdown

Fixes #8663

1. IterableDataset.batch() takes arrow-backed data out of arrow for no reason
(iterable_dataset.py). It regroups via a Python transpose, _batch_fn, which converts arrow
to Python examples and transposes them back. #8126 already added an arrow path for this, gated
on self._formatting.is_table; this extends the gate to any arrow-backed source:

-        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):

2. _consolidate cannot type numpy scalars (np_formatter.py):

-                isinstance(x, np.ndarray) and x.shape == column[0].shape and x.dtype == column[0].dtype
+                isinstance(x, (np.ndarray, np.number, np.bool_))
+                and x.shape == column[0].shape
+                and x.dtype == column[0].dtype

Performance

200k rows x 4 columns, in-memory source, single process, pinned with taskset. Absolute wall time for a full pass of .batch(n).

Arrow-backed sources:

format bs before (s) after (s) speedup
numpy 32 6.663 2.775 2.4x
numpy 256 4.430 0.383 11.6x
numpy 1000 4.147 0.105 39.5x
numpy (string) 1000 3.633 1.306 2.8x
torch 32 12.020 3.128 3.8x
torch 256 8.830 0.423 20.9x
torch 1000 8.577 0.127 67.5x
no format 32 / 256 / 1000 0.961 / 0.622 / 0.579 0.991 / 0.611 / 0.574 unchanged

Generator-backed sources

format bs before (s) after (s)
numpy 32 / 256 / 1000 13.253 / 12.272 / 12.516 13.382 / 12.372 / 12.556
torch 1000 18.820 18.309
no format 32 / 1000 0.335 / 0.280 0.330 / 0.273

@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

IterableDataset numpy format returns object instead of the declared dtype

2 participants