1111import threading
1212import time
1313import unittest
14+ from types import SimpleNamespace
1415from typing import Any , Dict
1516from unittest .mock import MagicMock , Mock , patch
1617
1718import torch
1819
1920from simpletuner .helpers .data_backend .runtime import BatchFetcher
21+ from simpletuner .helpers .data_backend .runtime import batch_fetcher as batch_fetcher_module
2022from simpletuner .helpers .data_backend .runtime import dataloader_iterator as dataloader_module
2123from 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