diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index ec3d77f..321f6c4 100644 --- a/animatediff/model_utils.py +++ b/animatediff/model_utils.py @@ -49,7 +49,7 @@ def get_lora_path(lora_name): def get_model_hash(file_path): with open(file_path, "rb") as f: - bytes = f.read() # read entire file as bytes + bytes = f.read(1024 * 1024) # read entire file as bytes return hashlib.sha256(bytes).hexdigest() @@ -93,7 +93,7 @@ def load_lora(lora_name: str): weight_down = state_dict[key] weight_up = state_dict[up_key] - updated_state_dict[combined_key] = torch.mm(weight_up, weight_down) + updated_state_dict[combined_key] = torch.mm(weight_up, weight_down).to("cpu") motion_loras[lora_hash] = updated_state_dict diff --git a/animatediff/nodes.py b/animatediff/nodes.py index a59584f..ea5645d 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -53,6 +53,20 @@ class AnimateDiffModuleLoader: curr_layer.weight.data += alpha * state_dict[key].to(curr_layer.weight.data.device) + def eject_loras(self, motion_module: MotionWrapper, lora_stack: List[Tuple[float, Dict[str, Tensor]]]): + for lora in lora_stack.reverse(): + (alpha, state_dict) = lora + + for key in state_dict: + layer_infos = key.split(".") + + curr_layer = motion_module + while len(layer_infos) > 0: + temp_name = layer_infos.pop(0) + curr_layer = curr_layer.__getattr__(temp_name) + + curr_layer.weight.data -= alpha * state_dict[key].to(curr_layer.weight.data.device) + def load_motion_module( self, model_name: str, @@ -61,11 +75,17 @@ class AnimateDiffModuleLoader: motion_module = load_motion_module(model_name) # inject loras - if isinstance(lora_stack, list): - if motion_module.is_v2: + if motion_module.is_v2: + if hasattr(motion_module, "lora_stack"): + self.eject_loras(motion_module, motion_module.lora_stack) + delattr(motion_module, "lora_stack") + + if isinstance(lora_stack, list): self.inject_loras(motion_module, lora_stack) - else: - logger.warning("LoRA is provided but only motion module v2 is supported.") + setattr(motion_module, "lora_stack", lora_stack) + + elif isinstance(lora_stack, list): + logger.warning("LoRA is provided but only motion module v2 is supported.") return (motion_module,)