Skip to content

Commit 405b2bd

Browse files
committed
make log clear when dataset is run out and re-looped
1 parent d1b7507 commit 405b2bd

1 file changed

Lines changed: 15 additions & 4 deletions

File tree

torchtitan/hf_datasets/text_datasets.py

Lines changed: 15 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,7 @@ def __init__(
154154
ds = dataset_loader(path)
155155

156156
self.dataset_name = dataset_name
157+
self.dataset_path = dataset_path
157158
self._data = split_dataset_by_node(ds, dp_rank, dp_world_size)
158159
self._tokenizer = tokenizer
159160
self.infinite = infinite
@@ -201,12 +202,16 @@ def __iter__(self):
201202
yield sample_tokens
202203

203204
if not self.infinite:
204-
logger.warning(f"Dataset {self.dataset_name} has run out of data")
205+
logger.warning(
206+
f"HuggingFaceDataset {self.dataset_name} from {self.dataset_path} has run out of data"
207+
)
205208
break
206209
else:
207210
# Reset offset for the next iteration
208211
self._sample_idx = 0
209-
logger.warning(f"Dataset {self.dataset_name} is being re-looped")
212+
logger.warning(
213+
f"HuggingFaceDataset {self.dataset_name} from {self.dataset_path} is being re-looped"
214+
)
210215
# Ensures re-looping a dataset loaded from a checkpoint works correctly
211216
if not isinstance(self._data, Dataset):
212217
if hasattr(self._data, "set_epoch") and hasattr(
@@ -344,6 +349,10 @@ def __init__(
344349
def dataset_name(self):
345350
return self._data.dataset_name
346351

352+
@property
353+
def dataset_path(self):
354+
return self._data.dataset_path
355+
347356
def _get_data_iter(self):
348357
# We don't use the sample index because we defer skipping to the
349358
# sub-dataset.
@@ -367,13 +376,15 @@ def __iter__(self):
367376

368377
if not self.infinite:
369378
logger.warning(
370-
f"Packed dataset {self.dataset_name} has run out of data"
379+
f"GreedyPackedDataset {self.dataset_name} from {self.dataset_path} has run out of data"
371380
)
372381
break
373382
else:
374383
# Reset offset for the next iteration
375384
self._sample_idx = 0
376-
logger.warning(f"Packed dataset {self.dataset_name} is being re-looped")
385+
logger.warning(
386+
f"GreedyPackedDataset {self.dataset_name} from {self.dataset_path} is being re-looped"
387+
)
377388
# Ensures re-looping a dataset loaded from a checkpoint works correctly
378389
if not isinstance(self._data, Dataset):
379390
if hasattr(self._data, "set_epoch") and hasattr(

0 commit comments

Comments
 (0)