diff --git a/src/diffusers/loaders/single_file_utils.py b/src/diffusers/loaders/single_file_utils.py index 296f32f891f0..27d65728b5ff 100644 --- a/src/diffusers/loaders/single_file_utils.py +++ b/src/diffusers/loaders/single_file_utils.py @@ -149,6 +149,10 @@ "net.pos_embedder.dim_spatial_range", ], "flux2": ["model.diffusion_model.single_stream_modulation.lin.weight", "single_stream_modulation.lin.weight"], + "chroma": [ + "model.diffusion_model.distilled_guidance_layer.in_proj.bias", + "distilled_guidance_layer.in_proj.bias", + ], "ltx2": [ "model.diffusion_model.av_ca_a2v_gate_adaln_single.emb.timestep_embedder.linear_1.weight", "vae.per_channel_statistics.mean-of-means", @@ -204,6 +208,7 @@ "flux-depth": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.1-Depth-dev"}, "flux-schnell": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.1-schnell"}, "flux-2-dev": {"pretrained_model_name_or_path": "black-forest-labs/FLUX.2-dev"}, + "chroma": {"pretrained_model_name_or_path": "lodestones/Chroma1-HD"}, "ltx-video": {"pretrained_model_name_or_path": "diffusers/LTX-Video-0.9.0"}, "ltx-video-0.9.1": {"pretrained_model_name_or_path": "diffusers/LTX-Video-0.9.1"}, "ltx-video-0.9.5": {"pretrained_model_name_or_path": "Lightricks/LTX-Video-0.9.5"}, @@ -683,6 +688,9 @@ def infer_diffusers_model_type(checkpoint): elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["flux2"]): model_type = "flux-2-dev" + elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["chroma"]): + model_type = "chroma" + elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["flux"]): if any( g in checkpoint for g in ["guidance_in.in_layer.bias", "model.diffusion_model.guidance_in.in_layer.bias"] @@ -754,17 +762,19 @@ def infer_diffusers_model_type(checkpoint): elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["wan"]): if "model.diffusion_model.patch_embedding.weight" in checkpoint: - target_key = "model.diffusion_model.patch_embedding.weight" + prefix = "model.diffusion_model." else: - target_key = "patch_embedding.weight" + prefix = "" + + target_key = f"{prefix}patch_embedding.weight" - if CHECKPOINT_KEY_NAMES["wan_vace"] in checkpoint: + if f"{prefix}{CHECKPOINT_KEY_NAMES['wan_vace']}" in checkpoint: if checkpoint[target_key].shape[0] == 1536: model_type = "wan-vace-1.3B" elif checkpoint[target_key].shape[0] == 5120: model_type = "wan-vace-14B" - if CHECKPOINT_KEY_NAMES["wan_animate"] in checkpoint: + elif f"{prefix}{CHECKPOINT_KEY_NAMES['wan_animate']}" in checkpoint: model_type = "wan-animate-14B" elif checkpoint[target_key].shape[0] == 1536: diff --git a/tests/models/autoencoders/test_models_autoencoder_dc.py b/tests/models/autoencoders/test_models_autoencoder_dc.py index a5ce0a975a11..c1743623021d 100644 --- a/tests/models/autoencoders/test_models_autoencoder_dc.py +++ b/tests/models/autoencoders/test_models_autoencoder_dc.py @@ -20,7 +20,13 @@ from diffusers.utils.torch_utils import randn_tensor from ...testing_utils import IS_GITHUB_ACTIONS, enable_full_determinism, torch_device -from ..testing_utils import BaseModelTesterConfig, MemoryTesterMixin, ModelTesterMixin, TrainingTesterMixin +from ..testing_utils import ( + BaseModelTesterConfig, + MemoryTesterMixin, + ModelTesterMixin, + SingleFileTesterMixin, + TrainingTesterMixin, +) from .testing_utils import NewAutoencoderTesterMixin @@ -102,3 +108,13 @@ def test_layerwise_casting_memory(self): class TestAutoencoderDCSlicingTiling(AutoencoderDCTesterConfig, NewAutoencoderTesterMixin): """Slicing and tiling tests for AutoencoderDC.""" + + +class TestAutoencoderDCSingleFile(AutoencoderDCTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/mit-han-lab/dc-ae-f32c32-sana-1.0/blob/main/model.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "mit-han-lab/dc-ae-f32c32-sana-1.0-diffusers" diff --git a/tests/models/autoencoders/test_models_autoencoder_kl.py b/tests/models/autoencoders/test_models_autoencoder_kl.py index 1872820269a1..acd15485f4b7 100644 --- a/tests/models/autoencoders/test_models_autoencoder_kl.py +++ b/tests/models/autoencoders/test_models_autoencoder_kl.py @@ -35,7 +35,13 @@ torch_all_close, torch_device, ) -from ..testing_utils import BaseModelTesterConfig, MemoryTesterMixin, ModelTesterMixin, TrainingTesterMixin +from ..testing_utils import ( + BaseModelTesterConfig, + MemoryTesterMixin, + ModelTesterMixin, + SingleFileTesterMixin, + TrainingTesterMixin, +) from .testing_utils import NewAutoencoderTesterMixin @@ -419,3 +425,19 @@ def test_stable_diffusion_encode_sample(self, seed, expected_slice): tolerance = 3e-3 if torch_device != "mps" else 1e-2 assert torch_all_close(output_slice, expected_output_slice, atol=tolerance) + + +class TestAutoencoderKLSingleFile(AutoencoderKLTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/stabilityai/sd-vae-ft-mse-original/blob/main/vae-ft-mse-840000-ema-pruned.safetensors" + + @property + def pretrained_model_name_or_path(self): + # `from_single_file` resolves a bare VAE checkpoint against the SD 1.5 repo, so compare against that + # rather than against `stabilityai/sd-vae-ft-mse`, whose own config differs (sample_size 256 vs 512). + return "stable-diffusion-v1-5/stable-diffusion-v1-5" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "vae"} diff --git a/tests/models/autoencoders/test_models_autoencoder_ltx_video.py b/tests/models/autoencoders/test_models_autoencoder_ltx_video.py index d93c1e566712..b8d992c63d33 100644 --- a/tests/models/autoencoders/test_models_autoencoder_ltx_video.py +++ b/tests/models/autoencoders/test_models_autoencoder_ltx_video.py @@ -20,7 +20,13 @@ from diffusers.utils.torch_utils import randn_tensor from ...testing_utils import enable_full_determinism, torch_device -from ..testing_utils import BaseModelTesterConfig, MemoryTesterMixin, ModelTesterMixin, TrainingTesterMixin +from ..testing_utils import ( + BaseModelTesterConfig, + MemoryTesterMixin, + ModelTesterMixin, + SingleFileTesterMixin, + TrainingTesterMixin, +) from .testing_utils import NewAutoencoderTesterMixin @@ -179,3 +185,17 @@ def test_gradient_checkpointing_is_applied(self): class TestAutoencoderKLLTXVideo091Memory(AutoencoderKLLTXVideo091TesterConfig, MemoryTesterMixin): """Memory optimization tests for AutoencoderKLLTXVideo (0.9.1 config).""" + + +class TestAutoencoderKLLTXVideoSingleFile(AutoencoderKLLTXVideo090TesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Lightricks/LTX-Video/blob/main/ltx-video-2b-v0.9.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "diffusers/LTX-Video-0.9.0" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "vae"} diff --git a/tests/models/autoencoders/test_models_autoencoder_wan.py b/tests/models/autoencoders/test_models_autoencoder_wan.py index 89e5c58e4a27..ee67cc87a9c3 100644 --- a/tests/models/autoencoders/test_models_autoencoder_wan.py +++ b/tests/models/autoencoders/test_models_autoencoder_wan.py @@ -20,7 +20,13 @@ from diffusers.utils.torch_utils import randn_tensor from ...testing_utils import enable_full_determinism, torch_device -from ..testing_utils import BaseModelTesterConfig, MemoryTesterMixin, ModelTesterMixin, TrainingTesterMixin +from ..testing_utils import ( + BaseModelTesterConfig, + MemoryTesterMixin, + ModelTesterMixin, + SingleFileTesterMixin, + TrainingTesterMixin, +) from .testing_utils import NewAutoencoderTesterMixin @@ -86,3 +92,17 @@ def test_layerwise_casting_training(self): class TestAutoencoderKLWanSlicingTiling(AutoencoderKLWanTesterConfig, NewAutoencoderTesterMixin): """Slicing and tiling tests for AutoencoderKLWan.""" + + +class TestAutoencoderKLWanSingleFile(AutoencoderKLWanTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/blob/main/split_files/vae/wan_2.1_vae.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "vae"} diff --git a/tests/models/testing_utils/single_file.py b/tests/models/testing_utils/single_file.py index 6e9443aa088c..adbc17c7355a 100644 --- a/tests/models/testing_utils/single_file.py +++ b/tests/models/testing_utils/single_file.py @@ -13,10 +13,15 @@ # See the License for the specific language governing permissions and # limitations under the License. +import functools import gc +import huggingface_hub +import pytest import torch -from huggingface_hub import hf_hub_download, snapshot_download +from accelerate import init_empty_weights +from huggingface_hub import HfApi +from safetensors.torch import _getdtype from diffusers.loaders.single_file_utils import _extract_repo_id_and_weights_name @@ -24,40 +29,30 @@ backend_empty_cache, is_single_file, nightly, - require_torch_accelerator, torch_device, ) from .common import check_device_map_is_respected -def download_single_file_checkpoint(pretrained_model_name_or_path, filename, tmpdir): - """Download a single file checkpoint from the Hub to a temporary directory.""" - path = hf_hub_download(pretrained_model_name_or_path, filename=filename, local_dir=tmpdir) - return path - - -def download_diffusers_config(pretrained_model_name_or_path, tmpdir): - """Download diffusers config files (excluding weights) from a repository.""" - path = snapshot_download( - pretrained_model_name_or_path, - ignore_patterns=[ - "**/*.ckpt", - "*.ckpt", - "**/*.bin", - "*.bin", - "**/*.pt", - "*.pt", - "**/*.safetensors", - "*.safetensors", - ], - allow_patterns=["**/*.json", "*.json", "*.txt", "**/*.txt"], - local_dir=tmpdir, - ) - return path - - -@nightly -@require_torch_accelerator +# Config entries that legitimately differ between a pretrained load and a single file load. +PARAMS_TO_IGNORE = ["torch_dtype", "_name_or_path", "_use_default_values", "_diffusers_version"] + + +@functools.lru_cache(maxsize=None) +def fetch_checkpoint_metadata(ckpt_path): + """ + Fetch a single file checkpoint's keys, shapes and dtypes without downloading its weights. + + Only the safetensors header crosses the network, and the result is cached for the session because every test + in a class asks for the same one. Returns `None` for checkpoints that are not safetensors. + """ + pretrained_model_name_or_path, weight_name = _extract_repo_id_and_weights_name(ckpt_path) + if not weight_name.endswith(".safetensors"): + return None + + return HfApi().parse_safetensors_file_metadata(pretrained_model_name_or_path, weight_name) + + @is_single_file class SingleFileTesterMixin: """ @@ -68,7 +63,6 @@ class SingleFileTesterMixin: Optional properties: - torch_dtype: torch dtype to use for testing (default: None) - - alternate_ckpt_paths: List of alternate checkpoint paths for variant testing (default: None) Expected from config mixin: - model_class: The model class to test @@ -93,11 +87,6 @@ def torch_dtype(self) -> torch.dtype | None: """torch dtype to use for single file testing.""" return None - @property - def alternate_ckpt_paths(self) -> list[str] | None: - """List of alternate checkpoint paths for variant testing.""" - return None - def setup_method(self): gc.collect() backend_empty_cache(torch_device) @@ -106,163 +95,135 @@ def teardown_method(self): gc.collect() backend_empty_cache(torch_device) - def test_single_file_model_config(self): - pretrained_kwargs = {"device_map": "auto", **self.pretrained_model_kwargs} - single_file_kwargs = {"device_map": "auto"} + def get_dummy_single_file_state_dict(self): + """ + A stand-in for the single file checkpoint's state dict. - if self.torch_dtype: - pretrained_kwargs["torch_dtype"] = self.torch_dtype - single_file_kwargs["torch_dtype"] = self.torch_dtype + Keys, shapes and dtypes match the real checkpoint exactly, while the tensors themselves are on meta and + hold no data. Pass it to `from_single_file` with `device="meta"` to exercise the loading path without + materializing a weight. The conversion functions consume the dict they are handed, so build a fresh one + for every load. + """ + metadata = fetch_checkpoint_metadata(self.ckpt_path) + if metadata is None: + pytest.skip(f"{self.ckpt_path} is not a safetensors checkpoint") - model = self.model_class.from_pretrained(self.pretrained_model_name_or_path, **pretrained_kwargs) - model_single_file = self.model_class.from_single_file(self.ckpt_path, **single_file_kwargs) + return { + name: torch.empty(tuple(tensor.shape), dtype=_getdtype(tensor.dtype), device="meta") + for name, tensor in metadata.tensors.items() + } + + def test_single_file_model_config(self): + # An empty model supplies the registered defaults (`out_channels` and friends) that a bare config dict + # would be missing. + with init_empty_weights(): + model = self.model_class.from_config(self.pretrained_model_name_or_path, **self.pretrained_model_kwargs) + + config = model.config + model_single_file = self.model_class.from_single_file(self.get_dummy_single_file_state_dict(), device="meta") - PARAMS_TO_IGNORE = ["torch_dtype", "_name_or_path", "_use_default_values", "_diffusers_version"] for param_name, param_value in model_single_file.config.items(): if param_name in PARAMS_TO_IGNORE: continue - assert model.config[param_name] == param_value, ( + assert config[param_name] == param_value, ( f"{param_name} differs between pretrained loading and single file loading: " - f"pretrained={model.config[param_name]}, single_file={param_value}" + f"pretrained={config[param_name]}, single_file={param_value}" ) - def test_single_file_model_parameters(self): - pretrained_kwargs = {"device_map": "auto", **self.pretrained_model_kwargs} - single_file_kwargs = {"device_map": "auto"} - - if self.torch_dtype: - pretrained_kwargs["torch_dtype"] = self.torch_dtype - single_file_kwargs["torch_dtype"] = self.torch_dtype - - # Load pretrained model, get state dict on CPU, then free GPU memory - model = self.model_class.from_pretrained(self.pretrained_model_name_or_path, **pretrained_kwargs) - state_dict = {k: v.cpu() for k, v in model.state_dict().items()} - del model - gc.collect() - backend_empty_cache(torch_device) + def test_single_file_loading_local_files_only(self, monkeypatch): + state_dict = self.get_dummy_single_file_state_dict() - # Load single file model, get state dict on CPU - model_single_file = self.model_class.from_single_file(self.ckpt_path, **single_file_kwargs) - state_dict_single_file = {k: v.cpu() for k, v in model_single_file.state_dict().items()} - del model_single_file - gc.collect() - backend_empty_cache(torch_device) + # The checkpoint is already in memory, so the config is the only thing left to resolve. Fetch it into the + # cache the way a previous run would have, and keep it as the reference to compare against. + config = self.model_class.load_config(self.pretrained_model_name_or_path, **self.pretrained_model_kwargs) - assert set(state_dict.keys()) == set(state_dict_single_file.keys()), ( - "Model parameters keys differ between pretrained and single file loading. " - f"Missing in single file: {set(state_dict.keys()) - set(state_dict_single_file.keys())}. " - f"Extra in single file: {set(state_dict_single_file.keys()) - set(state_dict.keys())}" - ) + # Cut the Hub off, so anything the load does not find locally raises instead of quietly downloading. + monkeypatch.setattr(huggingface_hub.constants, "HF_HUB_OFFLINE", True) - for key in state_dict.keys(): - param = state_dict[key] - param_single_file = state_dict_single_file[key] + model_single_file = self.model_class.from_single_file(state_dict, local_files_only=True, device="meta") - assert param.shape == param_single_file.shape, ( - f"Parameter shape mismatch for {key}: " - f"pretrained {param.shape} vs single file {param_single_file.shape}" + # Resolving from the cache has to land on the same config as resolving from the Hub. + for param_name, param_value in config.items(): + if param_name in PARAMS_TO_IGNORE: + continue + assert model_single_file.config[param_name] == param_value, ( + f"{param_name} differs when loading with local_files_only=True: " + f"pretrained={param_value}, single_file={model_single_file.config[param_name]}" ) - assert torch.equal(param, param_single_file), f"Parameter values differ for {key}" - - def test_single_file_loading_local_files_only(self, tmp_path): - single_file_kwargs = {} - - if self.torch_dtype: - single_file_kwargs["torch_dtype"] = self.torch_dtype - - pretrained_model_name_or_path, weight_name = _extract_repo_id_and_weights_name(self.ckpt_path) - local_ckpt_path = download_single_file_checkpoint(pretrained_model_name_or_path, weight_name, str(tmp_path)) - - model_single_file = self.model_class.from_single_file( - local_ckpt_path, local_files_only=True, **single_file_kwargs - ) - - assert model_single_file is not None, "Failed to load model with local_files_only=True" - def test_single_file_loading_with_diffusers_config(self): - single_file_kwargs = {} - - if self.torch_dtype: - single_file_kwargs["torch_dtype"] = self.torch_dtype - single_file_kwargs.update(self.pretrained_model_kwargs) - - # Load with config parameter model_single_file = self.model_class.from_single_file( - self.ckpt_path, config=self.pretrained_model_name_or_path, **single_file_kwargs + self.get_dummy_single_file_state_dict(), + config=self.pretrained_model_name_or_path, + device="meta", + **self.pretrained_model_kwargs, ) - # Load pretrained for comparison - pretrained_kwargs = {**self.pretrained_model_kwargs} - if self.torch_dtype: - pretrained_kwargs["torch_dtype"] = self.torch_dtype + # An empty model supplies the registered defaults (`out_channels` and friends) that a bare config dict + # would be missing. + with init_empty_weights(): + model = self.model_class.from_config(self.pretrained_model_name_or_path, **self.pretrained_model_kwargs) - model = self.model_class.from_pretrained(self.pretrained_model_name_or_path, **pretrained_kwargs) + config = model.config # Compare configs - PARAMS_TO_IGNORE = ["torch_dtype", "_name_or_path", "_use_default_values", "_diffusers_version"] for param_name, param_value in model_single_file.config.items(): if param_name in PARAMS_TO_IGNORE: continue - assert model.config[param_name] == param_value, ( - f"{param_name} differs: pretrained={model.config[param_name]}, single_file={param_value}" + assert config[param_name] == param_value, ( + f"{param_name} differs: pretrained={config[param_name]}, single_file={param_value}" ) - def test_single_file_loading_with_diffusers_config_local_files_only(self, tmp_path): - single_file_kwargs = {} - - if self.torch_dtype: - single_file_kwargs["torch_dtype"] = self.torch_dtype - single_file_kwargs.update(self.pretrained_model_kwargs) - - pretrained_model_name_or_path, weight_name = _extract_repo_id_and_weights_name(self.ckpt_path) - local_ckpt_path = download_single_file_checkpoint(pretrained_model_name_or_path, weight_name, str(tmp_path)) - local_diffusers_config = download_diffusers_config(self.pretrained_model_name_or_path, str(tmp_path)) - + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16], ids=["bf16", "fp16"]) + def test_single_file_loading_dtype(self, dtype): model_single_file = self.model_class.from_single_file( - local_ckpt_path, config=local_diffusers_config, local_files_only=True, **single_file_kwargs + self.get_dummy_single_file_state_dict(), device="meta", dtype=dtype ) + assert model_single_file.dtype == dtype, f"Expected dtype {dtype}, got {model_single_file.dtype}" - assert model_single_file is not None, "Failed to load model with config and local_files_only=True" - - def test_single_file_loading_dtype(self): - for dtype in [torch.float32, torch.float16]: - if torch_device == "mps" and dtype == torch.bfloat16: - continue - - model_single_file = self.model_class.from_single_file(self.ckpt_path, torch_dtype=dtype) - - assert model_single_file.dtype == dtype, f"Expected dtype {dtype}, got {model_single_file.dtype}" + @nightly + def test_single_file_model_parameters(self): + pretrained_kwargs = {**self.pretrained_model_kwargs} + single_file_kwargs = {} - # Cleanup - del model_single_file - gc.collect() - backend_empty_cache(torch_device) + if self.torch_dtype: + pretrained_kwargs["dtype"] = self.torch_dtype + single_file_kwargs["dtype"] = self.torch_dtype - def test_checkpoint_variant_loading(self): - if not self.alternate_ckpt_paths: - return + # Both models load onto the CPU, so keep only the state dicts and drop the models themselves. + model = self.model_class.from_pretrained(self.pretrained_model_name_or_path, **pretrained_kwargs) + state_dict = model.state_dict() + del model + gc.collect() - for ckpt_path in self.alternate_ckpt_paths: - backend_empty_cache(torch_device) + model_single_file = self.model_class.from_single_file(self.ckpt_path, **single_file_kwargs) + state_dict_single_file = model_single_file.state_dict() + del model_single_file + gc.collect() - single_file_kwargs = {} - if self.torch_dtype: - single_file_kwargs["torch_dtype"] = self.torch_dtype + assert set(state_dict.keys()) == set(state_dict_single_file.keys()), ( + "Model parameters keys differ between pretrained and single file loading. " + f"Missing in single file: {set(state_dict.keys()) - set(state_dict_single_file.keys())}. " + f"Extra in single file: {set(state_dict_single_file.keys()) - set(state_dict.keys())}" + ) - model = self.model_class.from_single_file(ckpt_path, **single_file_kwargs) + for key in state_dict.keys(): + param = state_dict[key] + param_single_file = state_dict_single_file[key] - assert model is not None, f"Failed to load checkpoint from {ckpt_path}" + assert param.shape == param_single_file.shape, ( + f"Parameter shape mismatch for {key}: " + f"pretrained {param.shape} vs single file {param_single_file.shape}" + ) - del model - gc.collect() - backend_empty_cache(torch_device) + assert torch.equal(param, param_single_file), f"Parameter values differ for {key}" + @nightly def test_single_file_loading_with_device_map(self): single_file_kwargs = {"device_map": "auto"} if self.torch_dtype: - single_file_kwargs["torch_dtype"] = self.torch_dtype + single_file_kwargs["dtype"] = self.torch_dtype model = self.model_class.from_single_file(self.ckpt_path, **single_file_kwargs) diff --git a/tests/models/transformers/test_models_transformer_chroma.py b/tests/models/transformers/test_models_transformer_chroma.py index d63aee32cbf6..dc300fbbe716 100644 --- a/tests/models/transformers/test_models_transformer_chroma.py +++ b/tests/models/transformers/test_models_transformer_chroma.py @@ -24,6 +24,7 @@ LoraHotSwappingForModelTesterMixin, LoraTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TrainingTesterMixin, ) @@ -155,3 +156,17 @@ def get_dummy_inputs(self, height: int = 4, width: int = 4) -> dict[str, torch.T "txt_ids": randn_tensor((sequence_length, num_image_channels), device=torch_device), "timestep": torch.tensor([1.0]).to(torch_device).expand(batch_size), } + + +class TestChromaTransformer2DSingleFile(ChromaTransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/lodestones/Chroma1-HD/blob/main/Chroma1-HD.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "lodestones/Chroma1-HD" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_flux.py b/tests/models/transformers/test_models_transformer_flux.py index 0214c1f65cab..2d65b60b448c 100644 --- a/tests/models/transformers/test_models_transformer_flux.py +++ b/tests/models/transformers/test_models_transformer_flux.py @@ -338,10 +338,6 @@ class TestFluxSingleFile(FluxTransformerTesterConfig, SingleFileTesterMixin): def ckpt_path(self): return "https://huggingface.co/black-forest-labs/FLUX.1-dev/blob/main/flux1-dev.safetensors" - @property - def alternate_ckpt_paths(self): - return ["https://huggingface.co/Comfy-Org/flux1-dev/blob/main/flux1-dev-fp8.safetensors"] - @property def pretrained_model_name_or_path(self): return "black-forest-labs/FLUX.1-dev" diff --git a/tests/models/transformers/test_models_transformer_flux2.py b/tests/models/transformers/test_models_transformer_flux2.py index 9546fdb5d969..6308d27085dc 100644 --- a/tests/models/transformers/test_models_transformer_flux2.py +++ b/tests/models/transformers/test_models_transformer_flux2.py @@ -36,6 +36,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchAoCompileTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, @@ -622,3 +623,17 @@ def test_no_kv_cache_mode_returns_no_cache(self): output = model(**base_config.get_dummy_inputs()) assert output.kv_cache is None + + +class TestFlux2Transformer2DSingleFile(Flux2TransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/black-forest-labs/FLUX.2-dev/blob/main/flux2-dev.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "black-forest-labs/FLUX.2-dev" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_hidream.py b/tests/models/transformers/test_models_transformer_hidream.py index c789850c0174..6489d82f50b1 100644 --- a/tests/models/transformers/test_models_transformer_hidream.py +++ b/tests/models/transformers/test_models_transformer_hidream.py @@ -22,6 +22,7 @@ from ..testing_utils import ( BaseModelTesterConfig, ModelTesterMixin, + SingleFileTesterMixin, TrainingTesterMixin, ) @@ -106,3 +107,17 @@ class TestHiDreamTransformerTraining(HiDreamTransformerTesterConfig, TrainingTes def test_gradient_checkpointing_is_applied(self): expected_set = {"HiDreamImageTransformer2DModel"} super().test_gradient_checkpointing_is_applied(expected_set=expected_set) + + +class TestHiDreamImageTransformer2DSingleFile(HiDreamTransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/HiDream-I1_ComfyUI/blob/main/split_files/diffusion_models/hidream_i1_full_fp16.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "HiDream-ai/HiDream-I1-Dev" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_ltx.py b/tests/models/transformers/test_models_transformer_ltx.py index 48397030633e..685734475b11 100644 --- a/tests/models/transformers/test_models_transformer_ltx.py +++ b/tests/models/transformers/test_models_transformer_ltx.py @@ -22,6 +22,7 @@ BaseModelTesterConfig, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -116,3 +117,17 @@ class TestLTXTransformerCompile(LTXTransformerTesterConfig, TorchCompileTesterMi # TODO: Add pretrained_model_name_or_path once a tiny LTX model is available on the Hub # class TestLTXTransformerTorchAo(LTXTransformerTesterConfig, TorchAoTesterMixin): # """TorchAo quantization tests for LTX Video Transformer.""" + + +class TestLTXVideoTransformer3DSingleFile(LTXTransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Lightricks/LTX-Video/blob/main/ltx-video-2b-v0.9.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "diffusers/LTX-Video-0.9.0" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_lumina2.py b/tests/models/transformers/test_models_transformer_lumina2.py index 7e50902392f4..1f54fb591e81 100644 --- a/tests/models/transformers/test_models_transformer_lumina2.py +++ b/tests/models/transformers/test_models_transformer_lumina2.py @@ -22,6 +22,7 @@ from ..testing_utils import ( BaseModelTesterConfig, ModelTesterMixin, + SingleFileTesterMixin, TrainingTesterMixin, ) @@ -95,3 +96,17 @@ class TestLumina2TransformerTraining(Lumina2TransformerTesterConfig, TrainingTes def test_gradient_checkpointing_is_applied(self): expected_set = {"Lumina2Transformer2DModel"} super().test_gradient_checkpointing_is_applied(expected_set=expected_set) + + +class TestLumina2Transformer2DSingleFile(Lumina2TransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/Lumina_Image_2.0_Repackaged/blob/main/split_files/diffusion_models/lumina_2_model_bf16.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "Alpha-VLLM/Lumina-Image-2.0" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_mochi.py b/tests/models/transformers/test_models_transformer_mochi.py index 8530dd98cf90..2c282a4138e4 100644 --- a/tests/models/transformers/test_models_transformer_mochi.py +++ b/tests/models/transformers/test_models_transformer_mochi.py @@ -22,6 +22,7 @@ from ..testing_utils import ( BaseModelTesterConfig, ModelTesterMixin, + SingleFileTesterMixin, TrainingTesterMixin, ) @@ -98,3 +99,17 @@ class TestMochiTransformerTraining(MochiTransformerTesterConfig, TrainingTesterM def test_gradient_checkpointing_is_applied(self): expected_set = {"MochiTransformer3DModel"} super().test_gradient_checkpointing_is_applied(expected_set=expected_set) + + +class TestMochiTransformer3DSingleFile(MochiTransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/mochi_preview_repackaged/blob/main/split_files/diffusion_models/mochi_preview_bf16.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "genmo/mochi-1-preview" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_sana.py b/tests/models/transformers/test_models_transformer_sana.py index 41d581c29090..c162dabd69cd 100644 --- a/tests/models/transformers/test_models_transformer_sana.py +++ b/tests/models/transformers/test_models_transformer_sana.py @@ -23,6 +23,7 @@ BaseModelTesterConfig, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TrainingTesterMixin, ) @@ -113,3 +114,17 @@ def test_gradient_checkpointing_is_applied(self): class TestSanaTransformerAttention(SanaTransformerTesterConfig, AttentionTesterMixin): pass + + +class TestSanaTransformer2DSingleFile(SanaTransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Efficient-Large-Model/Sana_1600M_1024px/blob/main/checkpoints/Sana_1600M_1024px.pth" + + @property + def pretrained_model_name_or_path(self): + return "Efficient-Large-Model/Sana_1600M_1024px_diffusers" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_sd3.py b/tests/models/transformers/test_models_transformer_sd3.py index e38c7853a613..0129960d985c 100644 --- a/tests/models/transformers/test_models_transformer_sd3.py +++ b/tests/models/transformers/test_models_transformer_sd3.py @@ -23,6 +23,7 @@ BaseModelTesterConfig, BitsAndBytesTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, @@ -224,3 +225,17 @@ class TestSD35TransformerBitsAndBytes(SD35TransformerTesterConfig, BitsAndBytesT class TestSD35TransformerTorchAo(SD35TransformerTesterConfig, TorchAoTesterMixin): """TorchAO quantization tests for SD3.5 Transformer.""" + + +class TestSD3Transformer2DSingleFile(SD3TransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/stabilityai/stable-diffusion-3-medium/blob/main/sd3_medium.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "stabilityai/stable-diffusion-3-medium-diffusers" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_wan.py b/tests/models/transformers/test_models_transformer_wan.py index c4c64eabcca7..26210eb3d547 100644 --- a/tests/models/transformers/test_models_transformer_wan.py +++ b/tests/models/transformers/test_models_transformer_wan.py @@ -26,6 +26,7 @@ GGUFTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, @@ -228,3 +229,35 @@ def get_dummy_inputs(self): ), "timestep": torch.tensor([1.0]).to(torch_device, self.torch_dtype), } + + +class TestWanTransformer3DText2VideoSingleFile(WanTransformer3DTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/blob/main/split_files/diffusion_models/wan2.1_t2v_1.3B_bf16.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} + + +class TestWanTransformer3DImage2VideoSingleFile(WanTransformer3DTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/blob/main/split_files/diffusion_models/wan2.1_i2v_480p_14B_fp8_e4m3fn.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} + + @property + def torch_dtype(self): + return torch.float8_e4m3fn diff --git a/tests/models/transformers/test_models_transformer_wan_vace.py b/tests/models/transformers/test_models_transformer_wan_vace.py index ac980b1ea7c0..cc0a3e34e12c 100644 --- a/tests/models/transformers/test_models_transformer_wan_vace.py +++ b/tests/models/transformers/test_models_transformer_wan_vace.py @@ -27,6 +27,7 @@ GGUFTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchAoTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, @@ -251,3 +252,17 @@ def get_dummy_inputs(self): ), "timestep": torch.tensor([1.0]).to(torch_device, self.torch_dtype), } + + +class TestWanVACETransformer3DSingleFile(WanVACETransformer3DTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/blob/main/split_files/diffusion_models/wan2.1_vace_1.3B_fp16.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "Wan-AI/Wan2.1-VACE-1.3B-diffusers" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/transformers/test_models_transformer_z_image.py b/tests/models/transformers/test_models_transformer_z_image.py index 85aeb34c25c4..06c4ee507fb0 100644 --- a/tests/models/transformers/test_models_transformer_z_image.py +++ b/tests/models/transformers/test_models_transformer_z_image.py @@ -28,6 +28,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -344,3 +345,17 @@ def _test_torch_compile_with_group_offload(self, config_kwargs, use_stream=False output = output[0] if isinstance(output, (list, tuple)) else output assert output is not None, "Model output is None" assert not torch.isnan(output).any(), "Model output contains NaN" + + +class TestZImageTransformer2DSingleFile(ZImageTransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/Comfy-Org/z_image_turbo/blob/main/split_files/diffusion_models/z_image_turbo_bf16.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "Tongyi-MAI/Z-Image-Turbo" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} diff --git a/tests/models/unets/test_models_unet_2d_condition.py b/tests/models/unets/test_models_unet_2d_condition.py index fc50a0a46125..f31d148caf2a 100644 --- a/tests/models/unets/test_models_unet_2d_condition.py +++ b/tests/models/unets/test_models_unet_2d_condition.py @@ -60,6 +60,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -1443,3 +1444,17 @@ def test_stabilityai_sd_v2_fp16(self, seed, timestep, expected_slice): expected_output_slice = torch.tensor(expected_slice) assert torch_all_close(output_slice, expected_output_slice, atol=5e-3) + + +class TestUNet2DConditionSingleFile(UNet2DConditionTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5/blob/main/v1-5-pruned-emaonly.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "stable-diffusion-v1-5/stable-diffusion-v1-5" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "unet"}