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
1 change: 1 addition & 0 deletions src/zeroband/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ class DataConfig(BaseConfig):
data_world_size: int | None = None
reverse_data_files: bool = False
split_by_data_rank: bool = True
dataset_read_retries: int = 3


class AdamConfig(BaseConfig):
Expand Down
33 changes: 29 additions & 4 deletions src/zeroband/data.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from dataclasses import dataclass, asdict
import random
import time
from typing import Any, Generator, Optional, List, Dict, TypedDict, Union
import functools
import threading
Expand Down Expand Up @@ -149,12 +150,31 @@ class ParquetDataset(IterableDataset, Stateful):
* [ ] handle mutli proc dataloader pytorch
"""

def __init__(self, files: List[str], tokenizer: PreTrainedTokenizer):
def __init__(self, files: List[str], tokenizer: PreTrainedTokenizer, read_retries: int = 3):
self.arg_files = files
self.tokenizer = tokenizer
self.read_retries = read_retries

self.state = None

def _read_text_column(self, file: str):
last_error: Exception | None = None
for attempt in range(self.read_retries):
try:
parquet_file = pq.ParquetFile(file)
return parquet_file.read()["text"]
except Exception as e:
last_error = e
if attempt + 1 >= self.read_retries:
break
delay = min(2**attempt, 30)
get_logger().warning(
f"Failed to read parquet {file!r} ({e}); retry {attempt + 1}/{self.read_retries} in {delay}s"
)
time.sleep(delay)
assert last_error is not None
raise last_error

def _lazy_init(self):
worker_info = torch.utils.data.get_worker_info()
if worker_info is not None:
Expand Down Expand Up @@ -185,8 +205,7 @@ def __iter__(self):
while True:
file = self.state.files[self.state.file_index]

parquet_file = pq.ParquetFile(file)
table = parquet_file.read()["text"]
table = self._read_text_column(file)

while True:
row = table[self.state.row_index]
Expand Down Expand Up @@ -452,6 +471,7 @@ def _load_datasets(
streaming: bool = True,
probabilities: Optional[List[float]] = None,
reverse_data_files: bool = False,
dataset_read_retries: int = 3,
) -> InterleaveDataset:
get_logger().debug(dataset_names)
ds_args = []
Expand All @@ -475,7 +495,11 @@ def _load_datasets(
datasets = []
for ds_arg in ds_args:
# logger.debug(f"Loading dataset: {ds_arg['data_files']}")
_ds = ParquetDataset(files=ds_arg["data_files"], tokenizer=tokenizer)
_ds = ParquetDataset(
files=ds_arg["data_files"],
tokenizer=tokenizer,
read_retries=dataset_read_retries,
)
datasets.append(_ds)

if len(datasets) > 1:
Expand Down Expand Up @@ -526,6 +550,7 @@ def load_all_datasets(
probabilities=_get_probabilities(data_config),
reverse_data_files=data_config.reverse_data_files,
tokenizer=tokenizer,
dataset_read_retries=data_config.dataset_read_retries,
)

get_logger().info(f"Train dataset: {ds}")
Expand Down
30 changes: 30 additions & 0 deletions tests/test_data.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import copy
import torch
from unittest.mock import patch
from zeroband.data import InterleaveDataset, ParquetDataset, SequencePackingDataSet, collate_fn
from torch.utils.data import DataLoader
from zeroband.data import load_all_datasets, DataConfig
Expand Down Expand Up @@ -270,3 +271,32 @@ def test_dataloader_parquet_dataset(parquet_files, tokenizer, num_workers):
assert (batch1["input_ids"] == batch2["input_ids"]).all()
assert (batch1["labels"] == batch2["labels"]).all()
assert (batch1["seqlens"] == batch2["seqlens"]).all()


def test_parquet_dataset_retries_transient_read_failure(parquet_files, tokenizer):
dataset = ParquetDataset(parquet_files[:1], tokenizer, read_retries=3)
real_parquet_file = pq.ParquetFile
attempts = 0

def flaky_parquet_file(path):
nonlocal attempts
attempts += 1
if attempts == 1:
raise OSError("HTTP 443 connection reset")
return real_parquet_file(path)

with patch("zeroband.data.pq.ParquetFile", side_effect=flaky_parquet_file):
with patch("zeroband.data.time.sleep"):
table = dataset._read_text_column(parquet_files[0])

assert len(table) == 100
assert attempts == 2


def test_parquet_dataset_read_failure_after_retries(parquet_files, tokenizer):
dataset = ParquetDataset(parquet_files[:1], tokenizer, read_retries=2)

with patch("zeroband.data.pq.ParquetFile", side_effect=OSError("HTTP 443")):
with patch("zeroband.data.time.sleep"):
with pytest.raises(OSError, match="HTTP 443"):
dataset._read_text_column(parquet_files[0])