資料集文件

Dataset 與 IterableDataset 之間的差異

Hugging Face's logo
加入 Hugging Face 社群

並獲得增強的文件體驗

開始使用

Dataset 與 IterableDataset 之間的差異

資料集物件主要有兩種型別:DatasetIterableDataset。您應選擇使用或建立哪種型別的資料集,取決於資料集的大小。一般而言,由於 IterableDataset 具備「惰性」(lazy) 特性且擁有速度優勢,因此它是大型資料集(例如數百 GB!)的理想選擇;而對於其他所有情況,Dataset 都是很好的選擇。本頁面將比較 DatasetIterableDataset 之間的差異,協助您挑選最適合的資料集物件。

下載與串流

當您擁有一般的 Dataset 時,可以使用 my_dataset[0] 來存取它。這提供了對資料列的隨機存取。這類資料集也稱為「映射型」(map-style) 資料集。例如,您可以這樣下載 ImageNet-1k 並存取其中任何一行資料:

from datasets import load_dataset

imagenet = load_dataset("timm/imagenet-1k-wds", split="train")  # downloads the full dataset
print(imagenet[0])

但一個缺點是,您必須將整個資料集儲存在磁碟或記憶體中,這會限制您存取大於磁碟容量的資料集。由於對於大型資料集來說這很不方便,因此存在另一種資料集型別:IterableDataset。當您擁有 IterableDataset 時,可以使用 for 迴圈進行迭代,並在迭代過程中逐步載入資料。透過這種方式,記憶體中只會載入一小部分的範例,且您不需要將任何內容寫入磁碟。

例如,您可以串流 ImageNet-1k 資料集,而無需將其下載到磁碟中:

from datasets import load_dataset

imagenet = load_dataset("timm/imagenet-1k-wds", split="train", streaming=True)  # will start loading the data when iterated over
for example in imagenet:
    print(example)
    break

串流可以在不將任何檔案寫入磁碟的情況下讀取線上資料。例如,您可以串流由多個分片 (shards) 組成的資料集,每個分片的大小可能高達數百 GB,例如 C4LAION-2B。在 資料集串流指南 中了解更多關於如何串流資料集的資訊。

不過這並非唯一的差異,因為 IterableDataset 的「惰性」特性也體現在資料集的建立與處理過程中。

建立映射型與可迭代型資料集

您可以使用列表或字典建立 Dataset,資料會被完全轉換為 Arrow 格式,以便您可以輕鬆存取任何一行資料:

my_dataset = Dataset.from_dict({"col_1": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]})
print(my_dataset[0])

另一方面,要建立 IterableDataset,您必須提供一種「惰性」方式來載入資料。在 Python 中,我們通常使用生成器 (generator) 函式。這些函式一次 yield(產出)一個範例,這意味著您無法像普通 Dataset 那樣透過切片 (slicing) 來存取資料列。

def my_generator(n):
    for i in range(n):
        yield {"col_1": i}

my_iterable_dataset = IterableDataset.from_generator(my_generator, gen_kwargs={"n": 10})
for example in my_iterable_dataset:
    print(example)
    break

完整與逐步載入本機檔案

可以使用 load_dataset() 將本機或遠端資料檔案轉換為 Arrow Dataset

data_files = {"train": ["path/to/data.csv"]}
my_dataset = load_dataset("csv", data_files=data_files, split="train")
print(my_dataset[0])

然而,這需要一個從 CSV 轉換為 Arrow 格式的步驟;如果您的資料集很大,這會耗費時間與磁碟空間。

為了節省磁碟空間並跳過轉換步驟,您可以透過直接串流本機檔案來定義 IterableDataset。這樣一來,資料會在您迭代資料集時,從本機檔案中逐步讀取。

data_files = {"train": ["path/to/data.csv"]}
my_iterable_dataset = load_dataset("csv", data_files=data_files, split="train", streaming=True)
for example in my_iterable_dataset:  # this reads the CSV file progressively as you iterate over the dataset
    print(example)
    break

許多檔案格式都受到支援,如 CSV、JSONL 和 Parquet,以及圖像和音訊檔案。您可以在對應的指南中找到更多資訊,分別是載入 表格文字視覺音訊 資料集。

即時資料處理與惰性資料處理

當您使用 Dataset.map() 處理 Dataset 物件時,整個資料集會立即被處理並回傳。這與 pandas 的運作方式類似。

my_dataset = my_dataset.map(process_fn)  # process_fn is applied on all the examples of the dataset
print(my_dataset[0])

另一方面,由於 IterableDataset 的「惰性」本質,呼叫 IterableDataset.map() 並不會將您的 map 函式套用到整個資料集上。相反地,您的 map 函式是「隨取隨用」(on-the-fly) 地被套用。

正因如此,您可以串接多個處理步驟,當您開始迭代資料集時,它們將會一次全部執行。

my_iterable_dataset = my_iterable_dataset.map(process_fn_1)
my_iterable_dataset = my_iterable_dataset.filter(filter_fn)
my_iterable_dataset = my_iterable_dataset.map(process_fn_2)

# process_fn_1, filter_fn and process_fn_2 are applied on-the-fly when iterating over the dataset
for example in my_iterable_dataset:  
    print(example)
    break

精確洗牌與快速近似洗牌

當您使用 Dataset.shuffle()Dataset 進行洗牌時,您執行的是資料集的精確洗牌。它的運作方式是取一個索引列表 [0, 1, 2, ... len(my_dataset) - 1] 並對此列表進行洗牌。接著,存取 my_dataset[0] 會回傳由洗牌後索引對應的第一個元素所定義的資料列。

my_dataset = my_dataset.shuffle(seed=42)
print(my_dataset[0])

由於在 IterableDataset 的情況下,我們無法對資料列進行隨機存取,因此無法使用洗牌後的索引列表來存取任意位置的資料列。這使得無法使用精確洗牌。相反地,IterableDataset.shuffle() 使用的是快速近似洗牌。它利用洗牌緩衝區 (shuffle buffer) 從資料集中迭代地抽樣隨機範例。由於資料集依然是迭代讀取的,因此這能提供極佳的速度效能。

my_iterable_dataset = my_iterable_dataset.shuffle(seed=42, buffer_size=100)
for example in my_iterable_dataset:
    print(example)
    break

但僅使用洗牌緩衝區並不足以提供機器學習模型訓練所需的滿意洗牌效果。因此,如果您的資料集由多個檔案或來源組成,IterableDataset.shuffle() 也會對資料集的分片進行洗牌。

# Stream from the internet
my_iterable_dataset = load_dataset("deepmind/code_contests", split="train", streaming=True)
my_iterable_dataset.num_shards  # 39

# Stream from local files
data_files = {"train": [f"path/to/data_{i}.csv" for i in range(1024)]}
my_iterable_dataset = load_dataset("csv", data_files=data_files, split="train", streaming=True)
my_iterable_dataset.num_shards  # 1024

# From a generator function
def my_generator(n, sources):
    for source in sources:
        for example_id_for_current_source in range(n):
            yield {"example_id": f"{source}_{example_id_for_current_source}"}

gen_kwargs = {"n": 10, "sources": [f"path/to/data_{i}" for i in range(1024)]}
my_iterable_dataset = IterableDataset.from_generator(my_generator, gen_kwargs=gen_kwargs)
my_iterable_dataset.num_shards  # 1024

速度差異

一般的 Dataset 物件基於 Arrow,它提供了對資料列的快速隨機存取。得益於記憶體映射 (memory mapping) 以及 Arrow 是一種記憶體內格式的事實,從磁碟讀取資料不會產生昂貴的系統呼叫與去序列化。當使用 for 迴圈進行迭代時,它透過迭代連續的 Arrow 記錄批次 (record batches),提供了更快的資料載入速度。

然而,一旦您的 Dataset 具有索引映射(例如透過 Dataset.shuffle()),速度可能會慢上 10 倍。這是因為需要額外的步驟來透過索引映射取得讀取資料列的指標,最重要的是,您不再讀取連續的資料區塊。為了恢復速度,您需要使用 Dataset.flatten_indices() 將整個資料集重新寫入磁碟,這會移除索引映射。根據資料集的大小,這可能會花費相當多的時間。

my_dataset[0]  # fast
my_dataset = my_dataset.shuffle(seed=42)
my_dataset[0]  # up to 10x slower
my_dataset = my_dataset.flatten_indices()  # rewrite the shuffled dataset on disk as contiguous chunks of data
my_dataset[0]  # fast again

在這種情況下,我們建議切換到 IterableDataset,並利用其快速近似洗牌方法 IterableDataset.shuffle()。它僅對分片順序進行洗牌並將洗牌緩衝區添加到您的資料集中,這能讓資料集保持最佳速度。您也可以輕鬆地重新洗牌資料集。

for example in enumerate(my_iterable_dataset):  # fast
    pass

shuffled_iterable_dataset = my_iterable_dataset.shuffle(seed=42, buffer_size=100)

for example in enumerate(shuffled_iterable_dataset):  # as fast as before
    pass

shuffled_iterable_dataset = my_iterable_dataset.shuffle(seed=1337, buffer_size=100)  # reshuffling using another seed is instantaneous

for example in enumerate(shuffled_iterable_dataset):  # still as fast as before
    pass

如果您在多個 Epoch 中使用資料集,用於洗牌緩衝區中分片順序的有效種子 (seed) 為 seed + epoch。這使得在 Epoch 之間重新洗牌資料集變得很容易。

for epoch in range(n_epochs):
    my_iterable_dataset.set_epoch(epoch)
    for example in my_iterable_dataset:  # fast + reshuffled at each epoch using `effective_seed = seed + epoch`
        pass

若要重啟映射型資料集的迭代,您只需跳過前幾個範例即可。

my_dataset = my_dataset.select(range(start_index, len(dataset)))

但如果您使用帶有 SamplerDataLoader,則應該儲存 Sampler 的狀態(您可能需要撰寫一個支援恢復狀態的自訂 Sampler)。

另一方面,可迭代資料集不提供針對特定範例索引的隨機存取來恢復進度。但您可以使用 IterableDataset.state_dict()IterableDataset.load_state_dict() 從檢查點 (checkpoint) 恢復,這與您處理模型和優化器的方式類似。

>>> iterable_dataset = Dataset.from_dict({"a": range(6)}).to_iterable_dataset(num_shards=3)
>>> # save in the middle of training
>>> state_dict = iterable_dataset.state_dict()
>>> # and resume later
>>> iterable_dataset.load_state_dict(state_dict)

在底層,可迭代資料集會追蹤目前正在讀取的分片以及該分片中的範例索引,並將此資訊儲存在 state_dict 中。

為了從檢查點恢復,資料集會跳過所有先前讀取過的分片,以便從目前的分片重新開始。然後它會讀取該分片並跳過範例,直到到達檢查點所在的確切範例。

因此,重啟資料集相當快,因為它不會重新讀取已經迭代過的分片。不過,恢復資料集通常並非即時的,因為它必須從目前分片的開頭重新開始讀取,並跳過範例直到到達檢查點位置。

此功能可以與 torchdata 中的 StatefulDataLoader 搭配使用,請參閱 使用 PyTorch DataLoader 進行串流

從映射型切換為可迭代型

如果您希望利用 IterableDataset 的「惰性」特性或其速度優勢,您可以將映射型的 Dataset 切換為 IterableDataset

my_iterable_dataset = my_dataset.to_iterable_dataset()

如果您想要洗牌您的資料集或 將其與 PyTorch DataLoader 一起使用,我們建議生成一個分片式的 IterableDataset

my_iterable_dataset = my_dataset.to_iterable_dataset(num_shards=1024)
my_iterable_dataset.num_shards  # 1024
在 GitHub 上更新

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