diff --git a/simpletuner/helpers/data_backend/factory.py b/simpletuner/helpers/data_backend/factory.py index 7e9d6bbad..05cffac64 100644 --- a/simpletuner/helpers/data_backend/factory.py +++ b/simpletuner/helpers/data_backend/factory.py @@ -64,6 +64,16 @@ def _as_bool(value: Any) -> bool: return bool(value) +def _normalise_vae_cache_config(config: Dict[str, Any]) -> Tuple[bool, bool]: + """Normalize optional VAE cache booleans and preserve disable-implies-ondemand.""" + + vae_cache_disable = _as_bool(config.get("vae_cache_disable", False)) + vae_cache_ondemand = _as_bool(config.get("vae_cache_ondemand", False)) or vae_cache_disable + config["vae_cache_disable"] = vae_cache_disable + config["vae_cache_ondemand"] = vae_cache_ondemand + return vae_cache_disable, vae_cache_ondemand + + def _coerce_bucket_keys(indices: Dict[Any, Iterable]) -> Dict[Any, list]: """Return a copy of aspect ratio bucket indices with numeric keys coerced to float.""" coerced: Dict[Any, list] = {} @@ -3997,7 +4007,8 @@ def _count_entries(collection) -> int: f"{len(init_backend['conditioning_image_embed_cache'].image_path_to_embed_path)} entries." ) - if not init_backend["config"]["vae_cache_ondemand"]: + _, vae_cache_ondemand = _normalise_vae_cache_config(init_backend["config"]) + if not vae_cache_ondemand: pending_files = init_backend["conditioning_image_embed_cache"].discover_unprocessed_files() logger.info(f"Conditioning image embed cache has {len(pending_files)} unprocessed files.") if is_i2v_video: @@ -4135,8 +4146,7 @@ def _configure_vae_cache( if vae_batch_size is None: vae_batch_size = self.args.vae_batch_size - dataset_vae_cache_disable = init_backend["config"]["vae_cache_disable"] - dataset_vae_cache_ondemand = init_backend["config"]["vae_cache_ondemand"] + dataset_vae_cache_disable, dataset_vae_cache_ondemand = _normalise_vae_cache_config(init_backend["config"]) init_backend["vaecache"] = VAECache( id=init_backend["id"], diff --git a/simpletuner/helpers/training/trainer.py b/simpletuner/helpers/training/trainer.py index 42d01b29c..37e74a088 100644 --- a/simpletuner/helpers/training/trainer.py +++ b/simpletuner/helpers/training/trainer.py @@ -5537,6 +5537,20 @@ def _epoch_rollover(self, epoch): gradient_accumulation_steps=self.config.gradient_accumulation_steps, apply_padding=(self.config.overrode_max_train_steps or self.config.allow_dataset_oversubscription), ) + local_sample_count = sum( + len(bucket) for bucket in backend["metadata_backend"].aspect_ratio_bucket_indices.values() + ) + if local_sample_count == 0: + raise ValueError( + f"(id={backend_id}) Dataset produced no usable samples. The epoch rollover" + f" re-split left this rank with zero samples, so epoch {epoch} would train" + f" against an empty schedule here.\n" + f"This usually means per-epoch re-bucketing (e.g., crop_aspect=random) produced" + f" buckets too small to divide across the data-parallel ranks, or that samples" + f" were filtered out of the cache during the previous epoch.\n" + f"Enable --allow_dataset_oversubscription so short buckets are padded across" + f" ranks, use fewer GPUs, or add more samples to the dataset." + ) # we have to rebuild the VAE cache if it exists. if "vaecache" in backend: logger.info("Rebuilding VAE cache..") diff --git a/tests/test_factory_edge_cases.py b/tests/test_factory_edge_cases.py index 0aacb9628..591c3c37c 100644 --- a/tests/test_factory_edge_cases.py +++ b/tests/test_factory_edge_cases.py @@ -249,12 +249,13 @@ def test_dataset_vae_cache_modes_use_dataset_ondemand_mode(self): model=self.model, ) modes = [ - ({"vae_cache_ondemand": True}, False), - ({"vae_cache_disable": True}, True), + ({"vae_cache_ondemand": True}, False, ()), + ({"vae_cache_disable": True}, True, ()), + ({"vae_cache_ondemand": True}, False, ("vae_cache_disable",)), ] - for backend_mode, expected_disable in modes: - with self.subTest(backend_mode=backend_mode): + for backend_mode, expected_disable, missing_config_keys in modes: + with self.subTest(backend_mode=backend_mode, missing_config_keys=missing_config_keys): backend = { "id": "dataset-ondemand-cache", "type": "local", @@ -264,6 +265,8 @@ def test_dataset_vae_cache_modes_use_dataset_ondemand_mode(self): **backend_mode, } init_backend = init_backend_config(backend, self.args, self.accelerator) + for key in missing_config_keys: + init_backend["config"].pop(key) init_backend.update( { "data_backend": MagicMock(id=backend["id"], type="local"), diff --git a/tests/test_trainer.py b/tests/test_trainer.py index b9aea9736..221e87d7f 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -10,6 +10,7 @@ import time import types import unittest +from contextlib import contextmanager from pathlib import Path from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch @@ -1862,6 +1863,7 @@ def test_epoch_rollover_reuses_initial_bucket_padding_policy(self): allow_dataset_oversubscription=allow_oversubscription, ) metadata_backend = MagicMock(read_only=True) + metadata_backend.aspect_ratio_bucket_indices = {"1.0": ["image-0.jpg"]} backends = {"train": {"metadata_backend": metadata_backend}} backend_config = { "crop": True, @@ -1886,6 +1888,72 @@ def test_epoch_rollover_reuses_initial_bucket_padding_policy(self): apply_padding=expected, ) + def _rollover_fixture(self, bucket_contents, usable_batches): + trainer = object.__new__(Trainer) + trainer.state = {"first_epoch": 1, "current_epoch": 1} + trainer.accelerator = MagicMock(is_main_process=True) + trainer.extra_lr_scheduler_kwargs = {} + trainer.get_steps_per_epoch_for_epoch = MagicMock(return_value=100) + trainer.config = SimpleNamespace( + num_train_epochs=5, + aspect_bucket_disable_rebuild=False, + lr_scheduler="constant", + num_update_steps_per_epoch=100, + gradient_accumulation_steps=1, + overrode_max_train_steps=False, + allow_dataset_oversubscription=False, + ) + metadata_backend = MagicMock(read_only=True, batch_size=4) + metadata_backend.aspect_ratio_bucket_indices = bucket_contents + metadata_backend.__len__.return_value = usable_batches + vaecache = MagicMock() + backends = {"train": {"metadata_backend": metadata_backend, "vaecache": vaecache}} + return trainer, metadata_backend, vaecache, backends + + @contextmanager + def _rollover_patches(self, backends): + with ( + patch("simpletuner.helpers.training.trainer.StateTracker.set_epoch"), + patch( + "simpletuner.helpers.training.trainer.StateTracker.get_data_backends", + return_value=backends, + ), + patch( + "simpletuner.helpers.training.trainer.StateTracker.get_data_backend_config", + return_value={"crop": True, "crop_aspect": "random"}, + ), + ): + yield + + def test_epoch_rollover_rejects_a_rank_the_resplit_left_empty(self): + trainer, metadata_backend, vaecache, backends = self._rollover_fixture( + bucket_contents={"1.0": [], "1.5": []}, usable_batches=0 + ) + + with self._rollover_patches(backends): + with self.assertRaises(ValueError) as context: + trainer._epoch_rollover(2) + + self.assertIn("Dataset produced no usable samples", str(context.exception)) + # The rejection must land after the re-split and before anything downstream + # consumes the new schedule. + metadata_backend.split_buckets_between_processes.assert_called_once() + vaecache.rebuild_cache.assert_not_called() + + def test_epoch_rollover_keeps_a_rank_that_cannot_fill_a_batch(self): + # One sample against batch_size=4, so the startup guard's len() would read 0 here. + # This guard asks only whether the shard is empty, and this rank keeps training on + # recycled samples exactly as it does today. + trainer, metadata_backend, vaecache, backends = self._rollover_fixture( + bucket_contents={"1.0": ["image-0.jpg"]}, usable_batches=0 + ) + + with self._rollover_patches(backends): + trainer._epoch_rollover(2) + + metadata_backend.split_buckets_between_processes.assert_called_once() + vaecache.rebuild_cache.assert_called_once() + @patch( "simpletuner.helpers.training.trainer.Trainer.parse_arguments", return_value=Mock(),