Skip to content

Commit 747f8a9

Browse files
author
bghira
committed
Add dynamic shift model invariants
1 parent 3d9dd95 commit 747f8a9

2 files changed

Lines changed: 68 additions & 0 deletions

File tree

simpletuner/helpers/models/flux2/model.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,6 +87,7 @@ class Flux2(ImageModelFoundation):
8787
ENABLED_IN_WIZARD = True
8888
PREDICTION_TYPE = PredictionTypes.FLOW_MATCHING
8989
MODEL_TYPE = ModelTypes.TRANSFORMER
90+
USES_DYNAMIC_SHIFT = True
9091
AUTO_LORA_FORMAT_DETECTION = True
9192
NATIVE_COMFYUI_LORA_SUPPORT = True # Flux2 has native ComfyUI LoRA support, no conversion needed
9293
COMFYUI_LORA_PRESERVE_COMPONENT_PREFIXES = {"transformer"}

tests/test_validation_dynamic_shift.py

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,13 @@
1+
import inspect
2+
import math
13
import types
24
import unittest
35
from unittest.mock import MagicMock, patch
46

57
import torch
68

79
from simpletuner.helpers.models.common import ModelFoundation
10+
from simpletuner.helpers.models.registry import ModelRegistry
811
from simpletuner.helpers.training.validation import Evaluation
912

1013

@@ -164,6 +167,70 @@ def set_timesteps(self, num_inference_steps=None, mu=None, **kwargs):
164167
with self.assertRaises(ValueError):
165168
eval_helper.get_timestep_schedule(scheduler, latents=torch.zeros(1, 4, 4, 4))
166169

170+
def _dynamic_shift_model_classes(self):
171+
dynamic_shift_model_classes = {}
172+
for family, registry_entry in ModelRegistry.model_families().items():
173+
if hasattr(registry_entry, "get_real_class"):
174+
model_cls = registry_entry.get_real_class()
175+
else:
176+
model_cls = registry_entry
177+
if getattr(model_cls, "USES_DYNAMIC_SHIFT", False):
178+
dynamic_shift_model_classes[family] = model_cls
179+
return dynamic_shift_model_classes
180+
181+
def _model_class_declares_patch_size(self, model_cls):
182+
component_cls = getattr(model_cls, "MODEL_CLASS", None)
183+
self.assertIsNotNone(component_cls, f"{model_cls.__name__} must define MODEL_CLASS")
184+
init_method = getattr(component_cls, "__init__", None)
185+
try:
186+
parameters = inspect.signature(init_method).parameters
187+
except (TypeError, ValueError):
188+
return False
189+
return "patch_size" in parameters
190+
191+
def test_dynamic_shift_models_expose_patch_geometry_or_sequence_length(self):
192+
model_classes = self._dynamic_shift_model_classes()
193+
self.assertGreater(len(model_classes), 0)
194+
195+
for family, model_cls in model_classes.items():
196+
with self.subTest(family=family):
197+
has_model_sequence_length = "_latent_sequence_length" in model_cls.__dict__
198+
has_component_patch_size = self._model_class_declares_patch_size(model_cls)
199+
self.assertTrue(
200+
has_model_sequence_length or has_component_patch_size,
201+
f"{model_cls.__name__} uses dynamic shift but does not expose patch geometry",
202+
)
203+
204+
model = model_cls.__new__(model_cls)
205+
model.config = types.SimpleNamespace(controlnet=False)
206+
model.accelerator = None
207+
model.model = types.SimpleNamespace(config=types.SimpleNamespace())
208+
if not has_model_sequence_length:
209+
model.model.config.patch_size = 2
210+
model.model.config.patch_size_t = 1
211+
212+
latent_channels = getattr(model_cls, "LATENT_CHANNEL_COUNT", 4)
213+
latents = torch.zeros(1, latent_channels, 8, 8)
214+
mu = model.calculate_dynamic_shift_mu(self.scheduler, latents)
215+
216+
self.assertTrue(math.isfinite(mu), f"{model_cls.__name__} produced non-finite dynamic shift mu")
217+
218+
def test_dynamic_scheduler_configs_declare_dynamic_shift_flag(self):
219+
for family, registry_entry in ModelRegistry.model_families().items():
220+
if hasattr(registry_entry, "get_real_class"):
221+
model_cls = registry_entry.get_real_class()
222+
else:
223+
model_cls = registry_entry
224+
scheduler_config = getattr(inspect.getmodule(model_cls), "SCHEDULER_CONFIG", None)
225+
if not isinstance(scheduler_config, dict) or not scheduler_config.get("use_dynamic_shifting"):
226+
continue
227+
228+
with self.subTest(family=family):
229+
self.assertTrue(
230+
getattr(model_cls, "USES_DYNAMIC_SHIFT", False),
231+
f"{model_cls.__name__} scheduler config enables dynamic shift without USES_DYNAMIC_SHIFT",
232+
)
233+
167234

168235
if __name__ == "__main__":
169236
unittest.main()

0 commit comments

Comments
 (0)