@@ -214,6 +214,44 @@ def test_cache_transformers(self, mock_create_model, mock_accelerator, mock_gree
214214 ],
215215 )
216216
217+ @patch ("lighteval.models.transformers.transformers_model.TransformersModel._padded_greedy_until" )
218+ @patch ("lighteval.models.transformers.transformers_model.Accelerator" )
219+ @patch ("lighteval.models.transformers.transformers_model.TransformersModel._create_auto_model" )
220+ def test_cache_only_main_process_writes (self , mock_create_model , mock_accelerator , mock_greedy_until ):
221+ """Regression test for #1102. Under a data-parallel (accelerate) launch every rank holds the same
222+ gathered results and would write the same parquet concurrently, corrupting the cache. Only the main
223+ process must write; other ranks must wait at a barrier (so the file exists before they read it)."""
224+ from lighteval .models .transformers .transformers_model import TransformersModel , TransformersModelConfig
225+
226+ mock_create_model = Mock () # noqa F841
227+ mock_accelerator_instance = Mock ()
228+ mock_accelerator_instance .device = torch .device ("cpu" )
229+ mock_accelerator .return_value = mock_accelerator_instance
230+ mock_greedy_until .return_value = self .model_responses
231+
232+ with tempfile .TemporaryDirectory () as temp_dir :
233+ config = TransformersModelConfig (model_name = "Qwen/Qwen3-0.6B" , cache_dir = temp_dir )
234+ model = TransformersModel (config )
235+ cache : SampleCache = model ._cache
236+ task_id = cache .get_task_id (self .task_name , SamplingMethod .GENERATIVE )
237+ cache_file = cache .get_cache_path (task_id )
238+
239+ # Non-main process: must NOT write the cache file, but must hit the barrier and still return
240+ # results (in a real run the main process has written the file by the time the barrier clears,
241+ # which we emulate by patching the cache read).
242+ mock_accelerator_instance .is_main_process = False
243+ mock_accelerator_instance .wait_for_everyone .reset_mock ()
244+ with patch .object (cache , "get_samples_from_cache" , return_value = self .model_responses ):
245+ results = model .greedy_until (self .docs )
246+ self .assertFalse (cache_file .exists (), "Non-main process must not write the cache file (#1102)" )
247+ mock_accelerator_instance .wait_for_everyone .assert_called ()
248+ self .assertEqual (len (results ), len (self .docs ))
249+
250+ # Main process: must write the cache file.
251+ mock_accelerator_instance .is_main_process = True
252+ model .greedy_until (self .docs )
253+ self .assertTrue (cache_file .exists (), "Main process must write the cache file" )
254+
217255 @patch ("lighteval.models.vllm.vllm_model.VLLMModel._loglikelihood_tokens" )
218256 @patch ("lighteval.models.vllm.vllm_model.VLLMModel._greedy_until" )
219257 @patch ("lighteval.models.vllm.vllm_model.VLLMModel._create_auto_model" )
0 commit comments