diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 425f2f29235e..a02c13b70d49 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -1288,19 +1288,44 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike | None ) if resolved_model_file is None and not is_sharded: - resolved_model_file = _get_model_file( - pretrained_model_name_or_path, - weights_name=_add_variant(WEIGHTS_NAME, variant), - cache_dir=cache_dir, - force_download=force_download, - proxies=proxies, - local_files_only=local_files_only, - token=token, - revision=revision, - subfolder=subfolder, - user_agent=user_agent, - commit_hash=commit_hash, - ) + pickle_index_file = None + try: + resolved_model_file = _get_model_file( + pretrained_model_name_or_path, + weights_name=_add_variant(WEIGHTS_NAME, variant), + cache_dir=cache_dir, + force_download=force_download, + proxies=proxies, + local_files_only=local_files_only, + token=token, + revision=revision, + subfolder=subfolder, + user_agent=user_agent, + commit_hash=commit_hash, + ) + except EnvironmentError: + if not allow_pickle or use_flashpack: + raise + pickle_index_file_kwargs = {**index_file_kwargs, "use_safetensors": False} + pickle_index_file = _fetch_index_file(**pickle_index_file_kwargs) + if variant is not None and (pickle_index_file is None or not os.path.exists(pickle_index_file)): + pickle_index_file = _fetch_index_file_legacy(**pickle_index_file_kwargs) + if pickle_index_file is None or not pickle_index_file.is_file(): + raise + + if pickle_index_file is not None: + is_sharded = True + resolved_model_file, sharded_metadata = _get_checkpoint_shard_files( + pretrained_model_name_or_path, + pickle_index_file, + cache_dir=cache_dir, + proxies=proxies, + local_files_only=local_files_only, + token=token, + user_agent=user_agent, + revision=revision, + subfolder=subfolder or "", + ) if not isinstance(resolved_model_file, list): resolved_model_file = [resolved_model_file] diff --git a/tests/models/test_modeling_common.py b/tests/models/test_modeling_common.py index b1bdaabbad7a..66954e663433 100644 --- a/tests/models/test_modeling_common.py +++ b/tests/models/test_modeling_common.py @@ -26,7 +26,7 @@ from huggingface_hub import ModelCard, delete_repo, snapshot_download, try_to_load_from_cache from huggingface_hub.utils import HfHubHTTPError, is_jinja_available -from diffusers.models import FluxTransformer2DModel, SD3Transformer2DModel, UNet2DConditionModel +from diffusers.models import FluxTransformer2DModel, SD3Transformer2DModel, UNet2DConditionModel, UNet2DModel from ..others.test_utils import TOKEN, USER, is_staging_test from ..testing_utils import ( @@ -167,6 +167,31 @@ def test_local_files_only_with_sharded_checkpoint(self): f"Expected error about missing shard, got: {error_msg}" ) + @pytest.mark.parametrize( + "variant, index_name", + [ + (None, "diffusion_pytorch_model.bin.index.json"), + ("ema", "diffusion_pytorch_model.bin.index.ema.json"), + ("ema", "diffusion_pytorch_model.bin.ema.index.json"), + ], + ) + def test_sharded_bin_checkpoint_loads_with_default_use_safetensors(self, tmp_path, variant, index_name): + model = UNet2DModel( + sample_size=32, + in_channels=3, + out_channels=3, + block_out_channels=(4, 8), + norm_num_groups=2, + down_block_types=("DownBlock2D", "AttnDownBlock2D"), + up_block_types=("AttnUpBlock2D", "UpBlock2D"), + ) + model.save_pretrained(tmp_path, variant=variant, safe_serialization=False, max_shard_size="50KB") + saved_index = next(tmp_path.glob("*.bin.index*.json")) + saved_index.rename(tmp_path / index_name) + + loaded = UNet2DModel.from_pretrained(tmp_path, variant=variant) + assert all(torch.equal(p1, p2) for p1, p2 in zip(model.parameters(), loaded.parameters())) + @pytest.mark.skip(reason="Flaky behaviour on CI. Re-enable after migrating to new runners") @pytest.mark.skipif(torch_device == "mps", reason="Test not supported for MPS.") def test_one_request_upon_cached(self):