|
| 1 | +import inspect |
| 2 | +import math |
1 | 3 | import types |
2 | 4 | import unittest |
3 | 5 | from unittest.mock import MagicMock, patch |
4 | 6 |
|
5 | 7 | import torch |
6 | 8 |
|
7 | 9 | from simpletuner.helpers.models.common import ModelFoundation |
| 10 | +from simpletuner.helpers.models.registry import ModelRegistry |
8 | 11 | from simpletuner.helpers.training.validation import Evaluation |
9 | 12 |
|
10 | 13 |
|
@@ -164,6 +167,70 @@ def set_timesteps(self, num_inference_steps=None, mu=None, **kwargs): |
164 | 167 | with self.assertRaises(ValueError): |
165 | 168 | eval_helper.get_timestep_schedule(scheduler, latents=torch.zeros(1, 4, 4, 4)) |
166 | 169 |
|
| 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 | + |
167 | 234 |
|
168 | 235 | if __name__ == "__main__": |
169 | 236 | unittest.main() |
0 commit comments