Skip to content
Open
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
51 changes: 38 additions & 13 deletions src/diffusers/models/modeling_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
27 changes: 26 additions & 1 deletion tests/models/test_modeling_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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):
Expand Down
Loading