@@ -537,6 +537,72 @@ def test_eval_backend_config_forces_train_batch_size_one(self):
537537
538538 self .assertEqual (result ["config" ]["train_batch_size" ], 1 )
539539
540+ def _apply_cached_train_batch_size (self , backend , cached_train_batch_size ):
541+ from simpletuner .helpers .data_backend .factory import FactoryRegistry , init_backend_config
542+
543+ init_backend = init_backend_config (backend , self .args , self .accelerator )
544+ effective_train_batch_size = init_backend ["config" ]["train_batch_size" ]
545+ metadata_backend = MagicMock ()
546+ metadata_backend .config = {"train_batch_size" : cached_train_batch_size }
547+ metadata_backend .batch_size = effective_train_batch_size
548+ metadata_backend .__len__ .return_value = 1
549+ init_backend ["metadata_backend" ] = metadata_backend
550+ init_backend ["data_backend" ] = MagicMock ()
551+
552+ factory = FactoryRegistry (
553+ args = self .args ,
554+ accelerator = self .accelerator ,
555+ text_encoders = self .text_encoders ,
556+ tokenizers = self .tokenizers ,
557+ model = self .model ,
558+ )
559+ sampler = MagicMock (caption_strategy = "filename" )
560+ with (
561+ patch ("simpletuner.helpers.data_backend.factory.StateTracker.set_data_backend_config" ) as set_config ,
562+ patch ("simpletuner.helpers.data_backend.factory.print_bucket_info" ),
563+ patch ("simpletuner.helpers.data_backend.factory.MultiAspectDataset" ),
564+ patch ("simpletuner.helpers.data_backend.factory.MultiAspectSampler" , return_value = sampler ) as sampler_cls ,
565+ patch ("simpletuner.helpers.data_backend.factory.torch.utils.data.DataLoader" ),
566+ ):
567+ factory ._handle_config_versioning (backend , init_backend )
568+ factory ._create_dataset_and_sampler (backend , init_backend , conditioning_type = None )
569+
570+ set_config .assert_called_once_with (init_backend ["id" ], init_backend ["config" ])
571+ self .assertEqual (init_backend ["config" ]["train_batch_size" ], effective_train_batch_size )
572+ self .assertEqual (metadata_backend .batch_size , effective_train_batch_size )
573+ self .assertEqual (metadata_backend .config ["train_batch_size" ], effective_train_batch_size )
574+ self .assertEqual (sampler_cls .call_args .kwargs ["batch_size" ], effective_train_batch_size )
575+ return init_backend
576+
577+ def test_cached_batch_size_does_not_override_changed_global_default (self ):
578+ self .args .train_batch_size = 4
579+ backend = {
580+ "id" : "image-global-batch" ,
581+ "type" : "local" ,
582+ "dataset_type" : "image" ,
583+ "instance_data_dir" : self .temp_dir ,
584+ }
585+
586+ result = self ._apply_cached_train_batch_size (backend , cached_train_batch_size = 2 )
587+
588+ self .assertNotIn ("train_batch_size" , backend )
589+ self .assertEqual (result ["config" ]["train_batch_size" ], 4 )
590+
591+ def test_cached_batch_size_does_not_override_changed_dataset_value (self ):
592+ self .args .train_batch_size = 8
593+ backend = {
594+ "id" : "image-dataset-batch" ,
595+ "type" : "local" ,
596+ "dataset_type" : "image" ,
597+ "train_batch_size" : 4 ,
598+ "instance_data_dir" : self .temp_dir ,
599+ }
600+
601+ result = self ._apply_cached_train_batch_size (backend , cached_train_batch_size = 2 )
602+
603+ self .assertEqual (backend ["train_batch_size" ], 4 )
604+ self .assertEqual (result ["config" ]["train_batch_size" ], 4 )
605+
540606 def test_inline_conditioning_auto_generation_for_image_dataset (self ):
541607 """Inline conditioning blocks on image datasets should spawn auto-generated conditioning datasets."""
542608 from simpletuner .helpers .data_backend .factory import FactoryRegistry
0 commit comments