From de829684400c68a4a40bc2f91de04add2f6a9f10 Mon Sep 17 00:00:00 2001 From: Abhishek Enaguthi Date: Wed, 8 Jul 2026 18:41:26 -0700 Subject: [PATCH] fix: retry transient parquet read failures in training ParquetDataset read errors (e.g. transient HTTP/IO while reading hub-cached parquet) crash long training runs. Retry with exponential backoff; configurable via data.dataset_read_retries (default 3). Fixes #100 --- src/zeroband/config.py | 1 + src/zeroband/data.py | 33 +++++++++++++++++++++++++++++---- tests/test_data.py | 30 ++++++++++++++++++++++++++++++ 3 files changed, 60 insertions(+), 4 deletions(-) diff --git a/src/zeroband/config.py b/src/zeroband/config.py index 11c27af5..9f7b7a0a 100644 --- a/src/zeroband/config.py +++ b/src/zeroband/config.py @@ -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): diff --git a/src/zeroband/data.py b/src/zeroband/data.py index 50ff1f58..f42751b7 100644 --- a/src/zeroband/data.py +++ b/src/zeroband/data.py @@ -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 @@ -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: @@ -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] @@ -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 = [] @@ -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: @@ -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}") diff --git a/tests/test_data.py b/tests/test_data.py index fff9d537..82bbcd37 100644 --- a/tests/test_data.py +++ b/tests/test_data.py @@ -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 @@ -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])