diff --git a/src/diffusers/loaders/single_file.py b/src/diffusers/loaders/single_file.py index 881ff9b96a4c..960b5b24ed7e 100644 --- a/src/diffusers/loaders/single_file.py +++ b/src/diffusers/loaders/single_file.py @@ -270,17 +270,18 @@ class FromSingleFileMixin: @classmethod @validate_hf_hub_args - def from_single_file(cls, pretrained_model_link_or_path, **kwargs) -> Self: + def from_single_file(cls, pretrained_model_link_or_path_or_dict, **kwargs) -> Self: r""" Instantiate a [`DiffusionPipeline`] from pretrained pipeline weights saved in the `.ckpt` or `.safetensors` format. The pipeline is set in evaluation mode (`model.eval()`) by default. Parameters: - pretrained_model_link_or_path (`str` or `os.PathLike`, *optional*): + pretrained_model_link_or_path_or_dict (`str`, *optional*): Can be either: - - A link to the `.ckpt` file (for example - `"https://huggingface.co//blob/main/.ckpt"`) on the Hub. - - A path to a *file* containing all pipeline weights. + - A link to the `.safetensors` or `.ckpt` file (for example + `"https://huggingface.co//blob/main/.safetensors"`) on the Hub. + - A path to a local *file* containing the weights of the component model. + - A state dict containing the component model weights. dtype (`str` or `torch.dtype`, *optional*): Override the default `torch.dtype` and load the model with another dtype. force_download (`bool`, *optional*, defaults to `False`): @@ -389,17 +390,19 @@ def from_single_file(cls, pretrained_model_link_or_path, **kwargs) -> Self: from ..pipelines.pipeline_utils import _get_pipeline_class pipeline_class = _get_pipeline_class(cls, config=None) - - checkpoint = load_single_file_checkpoint( - pretrained_model_link_or_path, - force_download=force_download, - proxies=proxies, - token=token, - cache_dir=cache_dir, - local_files_only=local_files_only, - revision=revision, - disable_mmap=disable_mmap, - ) + if isinstance(pretrained_model_link_or_path_or_dict, dict): + checkpoint = pretrained_model_link_or_path_or_dict + else: + checkpoint = load_single_file_checkpoint( + pretrained_model_link_or_path_or_dict, + force_download=force_download, + proxies=proxies, + token=token, + cache_dir=cache_dir, + local_files_only=local_files_only, + revision=revision, + disable_mmap=disable_mmap, + ) if config is None: config = fetch_diffusers_config(checkpoint)