資料集文件

與 JAX 搭配使用

Hugging Face's logo
加入 Hugging Face 社群

並獲得增強的文件體驗

開始使用

與 JAX 搭配使用

本文件是關於如何將 datasets 與 JAX 搭配使用的快速入門介紹,特別著重於如何從我們的數據集中取得 jax.Array 物件,以及如何使用它們來訓練 JAX 模型。

要重現上述程式碼,需要安裝 jaxjaxlib,請確保使用 pip install datasets[jax] 進行安裝。

資料集格式

預設情況下,數據集會傳回一般的 Python 物件:整數、浮點數、字串、列表等;字串與二進位物件將保持不變,因為 JAX 僅支援數值。

若要取得 JAX 陣列(類似 NumPy 陣列),您可以將數據集的格式設定為 jax

>>> from datasets import Dataset
>>> data = [[1, 2], [3, 4]]
>>> ds = Dataset.from_dict({"data": data})
>>> ds = ds.with_format("jax")
>>> ds[0]
{'data': DeviceArray([1, 2], dtype=int32)}
>>> ds[:2]
{'data': DeviceArray([
    [1, 2],
    [3, 4]], dtype=int32)}

Dataset 物件是 Arrow 表格的包裝器,它允許從數據集中的陣列快速讀取為 JAX 陣列。

請注意,相同的程序也適用於 DatasetDict 物件,因此當將 DatasetDict 的格式設定為 jax 時,其中的所有 Dataset 也將格式化為 jax

>>> from datasets import DatasetDict
>>> data = {"train": {"data": [[1, 2], [3, 4]]}, "test": {"data": [[5, 6], [7, 8]]}}
>>> dds = DatasetDict.from_dict(data)
>>> dds = dds.with_format("jax")
>>> dds["train"][:2]
{'data': DeviceArray([
    [1, 2],
    [3, 4]], dtype=int32)}

另一件需要考慮的事情是,格式化是在您實際存取資料時才會套用的。因此,如果您想從數據集中取得 JAX 陣列,則需要先存取資料,否則格式將保持不變。

最後,若要將資料載入到您選擇的裝置上,您可以指定 device 引數,但請注意,jaxlib.xla_extension.Device 不受支援,因為它無法使用 pickledill 進行序列化,因此您需要改用其字串識別碼。

>>> import jax
>>> from datasets import Dataset
>>> data = [[1, 2], [3, 4]]
>>> ds = Dataset.from_dict({"data": data})
>>> device = str(jax.devices()[0])  # Not casting to `str` before passing it to `with_format` will raise a `ValueError`
>>> ds = ds.with_format("jax", device=device)
>>> ds[0]
{'data': DeviceArray([1, 2], dtype=int32)}
>>> ds[0]["data"].device()
TFRT_CPU_0
>>> assert ds[0]["data"].device() == jax.devices()[0]
True

請注意,如果 with_format 沒有提供 device 引數,它將使用預設裝置,即 jax.devices()[0]

N 維陣列

如果您的資料集由 N 維陣列組成,您會發現若形狀固定,它們預設會被視為同一個張量

>>> from datasets import Dataset
>>> data = [[[1, 2],[3, 4]], [[5, 6],[7, 8]]]  # fixed shape
>>> ds = Dataset.from_dict({"data": data})
>>> ds = ds.with_format("jax")
>>> ds[0]
{'data': Array([[1, 2],
        [3, 4]], dtype=int32)}
>>> from datasets import Dataset
>>> data = [[[1, 2],[3]], [[4, 5, 6],[7, 8]]]  # varying shape
>>> ds = Dataset.from_dict({"data": data})
>>> ds = ds.with_format("jax")
>>> ds[0]
{'data': [Array([1, 2], dtype=int32), Array([3], dtype=int32)]}

然而,此邏輯通常需要較慢的形狀比較和資料複製。為了避免這種情況,您必須明確使用 Array 特徵類型並指定張量的形狀

>>> from datasets import Dataset, Features, Array2D
>>> data = [[[1, 2],[3, 4]],[[5, 6],[7, 8]]]
>>> features = Features({"data": Array2D(shape=(2, 2), dtype='int32')})
>>> ds = Dataset.from_dict({"data": data}, features=features)
>>> ds = ds.with_format("jax")
>>> ds[0]
{'data': Array([[1, 2],
        [3, 4]], dtype=int32)}
>>> ds[:2]
{'data': Array([[[1, 2],
         [3, 4]],
 
        [[5, 6],
         [7, 8]]], dtype=int32)}

其他特徵類型

ClassLabel 資料會被正確地轉換為陣列。

>>> from datasets import Dataset, Features, ClassLabel
>>> labels = [0, 0, 1]
>>> features = Features({"label": ClassLabel(names=["negative", "positive"])})
>>> ds = Dataset.from_dict({"label": labels}, features=features)
>>> ds = ds.with_format("jax")
>>> ds[:3]
{'label': DeviceArray([0, 0, 1], dtype=int32)}

字串與二進位物件將保持不變,因為 JAX 僅支援數值。

ImageAudio 特徵類型也受到支援。

若要使用 Image 特徵類型,您需要安裝 vision 額外套件,指令為 pip install datasets[vision]

>>> from datasets import Dataset, Features, Image
>>> images = ["path/to/image.png"] * 10
>>> features = Features({"image": Image()})
>>> ds = Dataset.from_dict({"image": images}, features=features)
>>> ds = ds.with_format("jax")
>>> ds[0]["image"].shape
(512, 512, 3)
>>> ds[0]
{'image': DeviceArray([[[ 255, 255, 255],
              [ 255, 255, 255],
              ...,
              [ 255, 255, 255],
              [ 255, 255, 255]]], dtype=uint8)}
>>> ds[:2]["image"].shape
(2, 512, 512, 3)
>>> ds[:2]
{'image': DeviceArray([[[[ 255, 255, 255],
              [ 255, 255, 255],
              ...,
              [ 255, 255, 255],
              [ 255, 255, 255]]]], dtype=uint8)}

若要使用 Audio 特徵類型,您需要安裝 audio 額外套件,指令為 pip install datasets[audio]

>>> from datasets import Dataset, Features, Audio
>>> audio = ["path/to/audio.wav"] * 10
>>> features = Features({"audio": Audio()})
>>> ds = Dataset.from_dict({"audio": audio}, features=features)
>>> ds = ds.with_format("jax")
>>> ds[0]["audio"]["array"]
DeviceArray([-0.059021  , -0.03894043, -0.00735474, ...,  0.0133667 ,
              0.01809692,  0.00268555], dtype=float32)
>>> ds[0]["audio"]["sampling_rate"]
DeviceArray(44100, dtype=int32, weak_type=True)

資料載入

JAX 沒有任何內建的資料載入功能,因此您需要使用像 PyTorchDataLoaderTensorFlowtf.data.Dataset 等函式庫來載入資料。引用關於此主題的 JAX 文件:「JAX 專注於程式轉換與加速器支援的 NumPy,因此我們不在 JAX 函式庫中包含資料載入或處理功能。市面上已經有很多優秀的資料載入器,所以讓我們直接使用它們,而不是重複造輪子。我們將使用 PyTorch 的資料載入器,並製作一個微小的墊片(shim)使其能與 NumPy 陣列協作。」。

這就是為什麼 datasets 中的 JAX 格式化功能如此有用的原因,它讓您可以使用 HuggingFace Hub 上的任何模型並搭配 JAX 使用,而無需擔心資料載入的部分。

使用 with_format('jax')

從數據集中取得 JAX 陣列最簡單的方法是使用 with_format('jax') 方法。假設我們想要在 HuggingFace Hub 上提供的 MNIST 數據集(位於 https://huggingface.co/datasets/ylecun/mnist)上訓練神經網路。

>>> from datasets import load_dataset
>>> ds = load_dataset("ylecun/mnist")
>>> ds = ds.with_format("jax")
>>> ds["train"][0]
{'image': DeviceArray([[  0,   0,   0, ...],
                       [  0,   0,   0, ...],
                       ...,
                       [  0,   0,   0, ...],
                       [  0,   0,   0, ...]], dtype=uint8),
 'label': DeviceArray(5, dtype=int32)}

一旦設定好格式,我們就可以使用 Dataset.iter() 方法以批次(batch)方式將數據集餵入 JAX 模型。

>>> for epoch in range(epochs):
...     for batch in ds["train"].iter(batch_size=32):
...         x, y = batch["image"], batch["label"]
...         ...
在 GitHub 上更新

© . This site is unofficial and not affiliated with Hugging Face, Inc.