Skip to content
Merged

merge #2998

Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 13 additions & 3 deletions simpletuner/helpers/data_backend/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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] = {}
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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"],
Expand Down
14 changes: 14 additions & 0 deletions simpletuner/helpers/training/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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..")
Expand Down
11 changes: 7 additions & 4 deletions tests/test_factory_edge_cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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"),
Expand Down
68 changes: 68 additions & 0 deletions tests/test_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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(),
Expand Down
Loading