Skip to content

Commit f1ce2d5

Browse files
authored
Merge pull request #3121 from bghira/issue/3099
Make DDP dataset selection deterministic
2 parents d3049f4 + 80dc99c commit f1ce2d5

3 files changed

Lines changed: 103 additions & 17 deletions

File tree

simpletuner/helpers/data_backend/runtime/batch_fetcher.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99

1010
from simpletuner.helpers.training.multi_process import should_log
1111

12+
from .dataloader_iterator import _set_epoch_step_for_training_step
1213
from .dataloader_iterator import random_dataloader_iterator as _runtime_random_iterator
1314

1415
DEFAULT_KEEP_RUNNING_PROPERTY = None
@@ -49,6 +50,7 @@ def _resolve_random_iterator():
4950

5051

5152
class BatchFetcher:
53+
"""Prefetch batches using the last consumed training step as the first schedule hint."""
5254

5355
def __init__(self, step: int, max_size: int = 10, datasets: Optional[Dict[str, Any]] = None) -> None:
5456
if DEFAULT_KEEP_RUNNING_PROPERTY is not None:
@@ -57,6 +59,7 @@ def __init__(self, step: int, max_size: int = 10, datasets: Optional[Dict[str, A
5759
self.datasets = datasets or {}
5860
self._keep_running = True
5961
self.step = step
62+
self._next_fetch_step = step
6063

6164
def start_fetching(self) -> threading.Thread:
6265
thread = threading.Thread(target=self.fetch_responses)
@@ -72,7 +75,14 @@ def fetch_responses(self) -> None:
7275
if self.queue.qsize() < self.queue.maxsize:
7376
prefetch_log_debug(f"Queue size: {self.queue.qsize()}. Fetching more data.")
7477
try:
75-
item = iterator_fn(self.step, self.datasets)
78+
if iterator_fn is DEFAULT_RANDOM_ITERATOR:
79+
item = iterator_fn(
80+
self._next_fetch_step,
81+
self.datasets,
82+
update_epoch_step=False,
83+
)
84+
else:
85+
item = iterator_fn(self._next_fetch_step, self.datasets)
7686
except ValueError:
7787
prefetch_log_debug("No datasets available during prefetch; stopping fetch thread.")
7888
self._keep_running = False
@@ -84,6 +94,7 @@ def fetch_responses(self) -> None:
8494
logger.debug(f"BatchFetcher encountered exception: {exc}")
8595
break
8696
self.queue.put(item)
97+
self._next_fetch_step += 1
8798
if self.queue.qsize() >= self.queue.maxsize:
8899
prefetch_log_debug("Completed fetching data. Queue is full.")
89100
if threading.current_thread() is threading.main_thread():
@@ -107,7 +118,9 @@ def next_response(self, step: int) -> Any:
107118
while self.queue.empty():
108119
continue
109120
prefetch_log_debug("Queue has data. Yielding next item.")
110-
return self.queue.get()
121+
item = self.queue.get()
122+
_set_epoch_step_for_training_step(step)
123+
return item
111124

112125
def stop_fetching(self) -> None:
113126
self._keep_running = False

simpletuner/helpers/data_backend/runtime/dataloader_iterator.py

Lines changed: 33 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
prefetch_log.setLevel(logging.ERROR)
2525

2626
_SCALED_SAMPLERS: dict[int, "ScaledDatasetSampler"] = {}
27+
_MAX_TORCH_SEED = 2**63 - 1
2728

2829

2930
def prefetch_log_debug(message: str) -> None:
@@ -32,14 +33,33 @@ def prefetch_log_debug(message: str) -> None:
3233
prefetch_log.debug(f"{rank_info()} {message}")
3334

3435

36+
def _set_epoch_step_for_training_step(step: int) -> None:
37+
args = StateTracker.get_args()
38+
grad_steps = getattr(args, "gradient_accumulation_steps", 1) if args is not None else 1
39+
if isinstance(grad_steps, (int, float)):
40+
gradient_accumulation_steps = max(1, int(grad_steps))
41+
else:
42+
gradient_accumulation_steps = 1
43+
StateTracker.set_epoch_step(int(step / gradient_accumulation_steps))
44+
45+
46+
def _selection_generator(step: int) -> torch.Generator:
47+
args = StateTracker.get_args()
48+
base_seed = getattr(args, "seed", 0) if args is not None else 0
49+
generator = torch.Generator(device="cpu")
50+
generator.manual_seed((int(base_seed or 0) + int(step)) % _MAX_TORCH_SEED)
51+
return generator
52+
53+
3554
def select_dataloader_index(step: int, backends: Dict[str, Any]) -> Optional[str]:
3655
if not backends:
3756
raise ValueError("No data backends available for selection.")
3857

3958
weights = []
4059
backend_ids = []
4160
inactive_ids = []
42-
for backend_id, backend in backends.items():
61+
for backend_id in sorted(backends):
62+
backend = backends[backend_id]
4363
weight = get_backend_weight(backend_id, backend, step)
4464
if weight <= 0:
4565
inactive_ids.append(backend_id)
@@ -69,7 +89,7 @@ def select_dataloader_index(step: int, backends: Dict[str, Any]) -> Optional[str
6989
raise ValueError("All backend sampling weights are zero.")
7090
weights_tensor /= total
7191

72-
chosen_index = torch.multinomial(weights_tensor, 1).item()
92+
chosen_index = torch.multinomial(weights_tensor, 1, generator=_selection_generator(step)).item()
7393
chosen_backend_id = backend_ids[chosen_index]
7494

7595
return chosen_backend_id
@@ -171,7 +191,8 @@ def __init__(self, slider_key: str = "slider_strength") -> None:
171191

172192
def rebuild(self, backends: Dict[str, Any]) -> None:
173193
self._groups = {"positive": [], "negative": [], "neutral": []}
174-
for backend_id, backend in backends.items():
194+
for backend_id in sorted(backends):
195+
backend = backends[backend_id]
175196
strength = _read_slider_strength(backend_id, backend, self.slider_key)
176197
if strength is None:
177198
self._groups["neutral"].append(backend_id)
@@ -226,7 +247,7 @@ def _choose_from_group(self, group: str, step: int, backends: Dict[str, Any]) ->
226247
if total <= 0:
227248
return None
228249
weights_tensor /= total
229-
chosen_index = torch.multinomial(weights_tensor, 1).item()
250+
chosen_index = torch.multinomial(weights_tensor, 1, generator=_selection_generator(step)).item()
230251
return filtered_ids[chosen_index]
231252

232253
def next_backend_id(self, step: int, backends: Dict[str, Any]) -> Optional[str]:
@@ -264,21 +285,20 @@ def _get_scaled_sampler(backends: Dict[str, Any]) -> Optional[ScaledDatasetSampl
264285
return sampler
265286

266287

267-
def random_dataloader_iterator(step: int, backends: Dict[str, Any]) -> Union[Any, bool]:
288+
def random_dataloader_iterator(
289+
step: int,
290+
backends: Dict[str, Any],
291+
*,
292+
update_epoch_step: bool = True,
293+
) -> Union[Any, bool]:
268294
if not backends:
269295
raise ValueError("No data backends provided to iterator.")
270296

271297
prefetch_log_debug("Random dataloader iterator launched.")
272-
args = StateTracker.get_args()
273-
grad_steps = getattr(args, "gradient_accumulation_steps", 1) if args is not None else 1
274-
if isinstance(grad_steps, (int, float)):
275-
gradient_accumulation_steps = max(1, int(grad_steps))
276-
else:
277-
gradient_accumulation_steps = 1
278298
logger.debug(f"Backends to select from {backends}")
279299
while backends:
280-
epoch_step = int(step / gradient_accumulation_steps)
281-
StateTracker.set_epoch_step(epoch_step)
300+
if update_epoch_step:
301+
_set_epoch_step_for_training_step(step)
282302

283303
sampler = _get_scaled_sampler(backends)
284304
chosen_backend_id = None

tests/test_runtime_components.py

Lines changed: 55 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,12 +11,14 @@
1111
import threading
1212
import time
1313
import unittest
14+
from types import SimpleNamespace
1415
from typing import Any, Dict
1516
from unittest.mock import MagicMock, Mock, patch
1617

1718
import torch
1819

1920
from simpletuner.helpers.data_backend.runtime import BatchFetcher
21+
from simpletuner.helpers.data_backend.runtime import batch_fetcher as batch_fetcher_module
2022
from simpletuner.helpers.data_backend.runtime import dataloader_iterator as dataloader_module
2123
from simpletuner.helpers.data_backend.runtime import get_backend_weight, random_dataloader_iterator, select_dataloader_index
2224

@@ -110,6 +112,29 @@ def check_iteration():
110112
# Verify iterator was called multiple times
111113
self.assertGreater(mock_iterator.call_count, 1)
112114

115+
@patch("simpletuner.helpers.data_backend.runtime.batch_fetcher.random_dataloader_iterator")
116+
def test_prefetch_starts_from_last_consumed_schedule_hint(self, mock_iterator):
117+
mock_iterator.side_effect = lambda step, datasets: {"step": step}
118+
fetcher = BatchFetcher(step=100, max_size=3, datasets=self.datasets)
119+
120+
fetcher.fetch_responses()
121+
122+
self.assertEqual(
123+
[call.args[0] for call in mock_iterator.call_args_list],
124+
[100, 101, 102],
125+
)
126+
127+
def test_runtime_iterator_does_not_advance_epoch_state_during_prefetch(self):
128+
mock_iterator = Mock(return_value={"batch": "data"})
129+
with (
130+
patch.object(batch_fetcher_module, "DEFAULT_RANDOM_ITERATOR", mock_iterator),
131+
patch.object(batch_fetcher_module, "random_dataloader_iterator", mock_iterator),
132+
):
133+
fetcher = BatchFetcher(step=100, max_size=1, datasets=self.datasets)
134+
fetcher.fetch_responses()
135+
136+
mock_iterator.assert_called_once_with(100, self.datasets, update_epoch_step=False)
137+
113138
@patch("simpletuner.helpers.data_backend.runtime.batch_fetcher.random_dataloader_iterator")
114139
def test_fetch_responses_queue_full_behavior(self, mock_iterator):
115140
"""Test fetch_responses behavior when queue is full"""
@@ -134,7 +159,8 @@ def limited_keep_running():
134159
# Queue should be at maximum capacity
135160
self.assertEqual(fetcher.queue.qsize(), 1)
136161

137-
def test_next_response_with_data_available(self):
162+
@patch("simpletuner.helpers.data_backend.runtime.batch_fetcher._set_epoch_step_for_training_step")
163+
def test_next_response_with_data_available(self, mock_set_epoch_step):
138164
"""Test next_response method when data is available"""
139165
fetcher = BatchFetcher(step=100, max_size=5, datasets=self.datasets)
140166

@@ -150,8 +176,10 @@ def test_next_response_with_data_available(self):
150176

151177
# Verify correct data was returned
152178
self.assertEqual(result, test_data)
179+
mock_set_epoch_step.assert_called_once_with(101)
153180

154-
def test_next_response_with_empty_queue(self):
181+
@patch("simpletuner.helpers.data_backend.runtime.batch_fetcher._set_epoch_step_for_training_step")
182+
def test_next_response_with_empty_queue(self, mock_set_epoch_step):
155183
"""Test next_response method blocks when queue is empty"""
156184
fetcher = BatchFetcher(step=100, max_size=5, datasets=self.datasets)
157185

@@ -179,6 +207,7 @@ def add_data_later():
179207

180208
# Verify step was updated
181209
self.assertEqual(fetcher.step, 101)
210+
mock_set_epoch_step.assert_called_once_with(101)
182211

183212
def test_stop_fetching(self):
184213
"""Test stop_fetching method"""
@@ -326,6 +355,19 @@ def test_select_dataloader_index_multiple_backends(self):
326355
self.assertTrue(len(results) >= 1) # At least one backend selected
327356
self.assertTrue(all(r in ["backend1", "backend2"] for r in results))
328357

358+
@patch("simpletuner.helpers.data_backend.runtime.dataloader_iterator.StateTracker.get_args")
359+
def test_select_dataloader_index_is_independent_of_global_rng(self, mock_get_args):
360+
mock_get_args.return_value = SimpleNamespace(seed=1234, data_backend_sampling="uniform")
361+
backends = {"backend1": self.mock_backend1, "backend2": self.mock_backend2}
362+
363+
torch.manual_seed(1)
364+
first_sequence = [select_dataloader_index(step=step, backends=backends) for step in range(32)]
365+
torch.manual_seed(987654)
366+
torch.rand(127)
367+
second_sequence = [select_dataloader_index(step=step, backends=backends) for step in range(32)]
368+
369+
self.assertEqual(first_sequence, second_sequence)
370+
329371
def test_select_dataloader_index_empty_backends(self):
330372
"""Test select_dataloader_index with empty backends"""
331373
backends = {}
@@ -394,6 +436,17 @@ def test_random_dataloader_iterator_single_backend(self, mock_select):
394436
# Verify correct dataloader was returned
395437
self.assertEqual(result, "dataloader1")
396438

439+
@patch("simpletuner.helpers.data_backend.runtime.dataloader_iterator.StateTracker")
440+
@patch("simpletuner.helpers.data_backend.runtime.dataloader_iterator.select_dataloader_index")
441+
def test_random_dataloader_iterator_can_defer_epoch_state_update(self, mock_select, mock_state_tracker):
442+
mock_select.return_value = "backend1"
443+
backends = {"backend1": self.mock_backend1}
444+
445+
result = random_dataloader_iterator(step=100, backends=backends, update_epoch_step=False)
446+
447+
self.assertEqual(result, "dataloader1")
448+
mock_state_tracker.set_epoch_step.assert_not_called()
449+
397450
@patch("simpletuner.helpers.data_backend.runtime.dataloader_iterator.select_dataloader_index")
398451
def test_random_dataloader_iterator_multiple_backends(self, mock_select):
399452
"""Test random_dataloader_iterator with multiple backends"""

0 commit comments

Comments
 (0)