Skip to content

Commit 809a2c7

Browse files
authored
Merge pull request #3090 from bghira/agent/dataset-batch-cache-versioning
Keep dataset batch size current across cache reloads
2 parents 1a0c94f + 5fcd668 commit 809a2c7

2 files changed

Lines changed: 76 additions & 0 deletions

File tree

simpletuner/helpers/data_backend/factory.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3376,6 +3376,7 @@ def _handle_bucket_operations(
33763376

33773377
def _handle_config_versioning(self, backend: Dict[str, Any], init_backend: Dict[str, Any]) -> None:
33783378
"""Handle configuration versioning and validation."""
3379+
runtime_mutable_keys = ("train_batch_size",)
33793380
excluded_keys = [
33803381
"probability",
33813382
"repeats",
@@ -3393,6 +3394,7 @@ def _handle_config_versioning(self, backend: Dict[str, Any], init_backend: Dict[
33933394
"start_epoch",
33943395
"hash_filenames", # always enabled, not user-configurable
33953396
"_s2v_audio_autoinjected", # runtime flag, not user-configurable
3397+
*runtime_mutable_keys,
33963398
]
33973399
_latest_config_version = latest_config_version()
33983400
current_config_version = _latest_config_version
@@ -3427,6 +3429,14 @@ def _handle_config_versioning(self, backend: Dict[str, Any], init_backend: Dict[
34273429
)
34283430
init_backend["config"][key] = prev_config[key]
34293431

3432+
metadata_backend_config = getattr(init_backend["metadata_backend"], "config", None)
3433+
if isinstance(metadata_backend_config, dict):
3434+
for key in runtime_mutable_keys:
3435+
if key in init_backend["config"]:
3436+
metadata_backend_config[key] = init_backend["config"][key]
3437+
else:
3438+
metadata_backend_config.pop(key, None)
3439+
34303440
runtime_linkage_keys = (
34313441
"conditioning_data",
34323442
"conditioning",

tests/test_factory_edge_cases.py

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)