huggingface / huggingface/datasets

Memory explosion when trying to access 4d tensors in datasets cast to torch or np

Open
#5,165 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
22k
Forks
3.4k
Avg merge
5d 7h
Merged PRs (30d)
17

Description

### Describe the bug

When trying to access an item by index, in a datasets.Dataset cast to torch/np using `set_format` or `with_format`, we get a memory explosion if the item contains 4d (or above) tensors.

### Steps to reproduce the bug

MWE:

```python
from datasets import load_dataset
import numpy as np

def create_4d_tensor(item):
i = item["num_nodes"]
item["x_big"] = np.random.rand(i, 2*i, int(i/2), 1) + 1 # we create a big 4d tensor
return item

if __name__ == "__main__":
dataset = load_dataset(path=f"graphs-datasets/PROTEINS")

# This works
print(dataset["train"].format)
print(dataset["train"][0].keys())

dataset = dataset.map(
create_4d_tensor,
batched=False,
writer_batch_size=100,
)

# This works
print(dataset["train"].format)
print(dataset["train"][0].keys())

dataset.set_format("torch")

print(dataset["train"].format)
# This gets killed :(
print(dataset["train"][0].keys())
```

The problem likely comes from `format_table` [here](https://cs.github.com/huggingface/datasets/blob/f09f781be3278156ce3aa6ec90c1926b1846a78f/src/datasets/arrow_dataset.py#L2328)

### Expected behavior

No memory explosion when trying to access dataset items after cast.

### Environment info

- `datasets` version: 2.3.2
- Platform: Linux-5.14.0-1054-oem-x86_64-with-glibc2.29
- Python version: 3.8.10
- PyArrow version: 8.0.0
- Pandas version: 1.4.3

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.