diff --git a/videox_fun/models/cogvideox_transformer3d.py b/videox_fun/models/cogvideox_transformer3d.py index e12d479..dd40721 100755 --- a/videox_fun/models/cogvideox_transformer3d.py +++ b/videox_fun/models/cogvideox_transformer3d.py @@ -708,6 +708,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): try: import re + from diffusers import __version__ as diffusers_version from diffusers.models.modeling_utils import \ load_model_dict_into_meta from diffusers.utils import is_accelerate_available @@ -733,32 +734,44 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): for key in _state_dict: state_dict[key] = _state_dict[key] model._convert_deprecated_attention_blocks(state_dict) - # move the params from meta device to cpu - missing_keys = set(model.state_dict().keys()) - set(state_dict.keys()) - if len(missing_keys) > 0: - raise ValueError( - f"Cannot load {cls} from {pretrained_model_path} because the following keys are" - f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass" - " `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize" - " those weights or else make sure your checkpoint file is correct." + + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + model, + state_dict, + dtype=torch_dtype, + model_name_or_path=pretrained_model_path, + ) + else: + # move the params from meta device to cpu + missing_keys = set(model.state_dict().keys()) - set(state_dict.keys()) + if len(missing_keys) > 0: + raise ValueError( + f"Cannot load {cls} from {pretrained_model_path} because the following keys are" + f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass" + " `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize" + " those weights or else make sure your checkpoint file is correct." + ) + + unexpected_keys = load_model_dict_into_meta( + model, + state_dict, + device=param_device, + dtype=torch_dtype, + model_name_or_path=pretrained_model_path, ) - unexpected_keys = load_model_dict_into_meta( - model, - state_dict, - device=param_device, - dtype=torch_dtype, - model_name_or_path=pretrained_model_path, - ) + if cls._keys_to_ignore_on_load_unexpected is not None: + for pat in cls._keys_to_ignore_on_load_unexpected: + unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - if cls._keys_to_ignore_on_load_unexpected is not None: - for pat in cls._keys_to_ignore_on_load_unexpected: - unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - - if len(unexpected_keys) > 0: - print( - f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" - ) + if len(unexpected_keys) > 0: + print( + f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" + ) + return model except Exception as e: print( diff --git a/videox_fun/models/wan_text_encoder.py b/videox_fun/models/wan_text_encoder.py index 49fc936..9bc41fb 100755 --- a/videox_fun/models/wan_text_encoder.py +++ b/videox_fun/models/wan_text_encoder.py @@ -316,6 +316,7 @@ class WanT5EncoderModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): try: import re + from diffusers import __version__ as diffusers_version from diffusers.models.modeling_utils import \ load_model_dict_into_meta from diffusers.utils import is_accelerate_available @@ -332,32 +333,44 @@ class WanT5EncoderModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): state_dict = load_file(pretrained_model_path) else: state_dict = torch.load(pretrained_model_path, map_location="cpu") - # move the params from meta device to cpu - missing_keys = set(model.state_dict().keys()) - set(state_dict.keys()) - if len(missing_keys) > 0: - raise ValueError( - f"Cannot load {cls} from {pretrained_model_path} because the following keys are" - f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass" - " `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize" - " those weights or else make sure your checkpoint file is correct." + + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + model, + state_dict, + dtype=torch_dtype, + model_name_or_path=pretrained_model_path, + ) + else: + # move the params from meta device to cpu + missing_keys = set(model.state_dict().keys()) - set(state_dict.keys()) + if len(missing_keys) > 0: + raise ValueError( + f"Cannot load {cls} from {pretrained_model_path} because the following keys are" + f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass" + " `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize" + " those weights or else make sure your checkpoint file is correct." + ) + + unexpected_keys = load_model_dict_into_meta( + model, + state_dict, + device=param_device, + dtype=torch_dtype, + model_name_or_path=pretrained_model_path, ) - unexpected_keys = load_model_dict_into_meta( - model, - state_dict, - device=param_device, - dtype=torch_dtype, - model_name_or_path=pretrained_model_path, - ) + if cls._keys_to_ignore_on_load_unexpected is not None: + for pat in cls._keys_to_ignore_on_load_unexpected: + unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - if cls._keys_to_ignore_on_load_unexpected is not None: - for pat in cls._keys_to_ignore_on_load_unexpected: - unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - - if len(unexpected_keys) > 0: - print( - f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" - ) + if len(unexpected_keys) > 0: + print( + f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" + ) + return model except Exception as e: print( diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index bc1ad0a..a66a405 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -1086,6 +1086,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): try: import re + from diffusers import __version__ as diffusers_version from diffusers.models.modeling_utils import \ load_model_dict_into_meta from diffusers.utils import is_accelerate_available @@ -1111,33 +1112,45 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): _state_dict = load_file(_model_file_safetensors) for key in _state_dict: state_dict[key] = _state_dict[key] - model._convert_deprecated_attention_blocks(state_dict) - # move the params from meta device to cpu - missing_keys = set(model.state_dict().keys()) - set(state_dict.keys()) - if len(missing_keys) > 0: - raise ValueError( - f"Cannot load {cls} from {pretrained_model_path} because the following keys are" - f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass" - " `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize" - " those weights or else make sure your checkpoint file is correct." + + if diffusers_version >= "0.33.0": + # Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit: + # https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785. + load_model_dict_into_meta( + model, + state_dict, + dtype=torch_dtype, + model_name_or_path=pretrained_model_path, + ) + else: + model._convert_deprecated_attention_blocks(state_dict) + # move the params from meta device to cpu + missing_keys = set(model.state_dict().keys()) - set(state_dict.keys()) + if len(missing_keys) > 0: + raise ValueError( + f"Cannot load {cls} from {pretrained_model_path} because the following keys are" + f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass" + " `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize" + " those weights or else make sure your checkpoint file is correct." + ) + + unexpected_keys = load_model_dict_into_meta( + model, + state_dict, + device=param_device, + dtype=torch_dtype, + model_name_or_path=pretrained_model_path, ) - unexpected_keys = load_model_dict_into_meta( - model, - state_dict, - device=param_device, - dtype=torch_dtype, - model_name_or_path=pretrained_model_path, - ) + if cls._keys_to_ignore_on_load_unexpected is not None: + for pat in cls._keys_to_ignore_on_load_unexpected: + unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - if cls._keys_to_ignore_on_load_unexpected is not None: - for pat in cls._keys_to_ignore_on_load_unexpected: - unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None] - - if len(unexpected_keys) > 0: - print( - f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" - ) + if len(unexpected_keys) > 0: + print( + f"Some weights of the model checkpoint were not used when initializing {cls.__name__}: \n {[', '.join(unexpected_keys)]}" + ) + return model except Exception as e: print(