optimization

This commit is contained in:
Tung Nguyen
2023-09-25 22:17:22 +07:00
parent 5ee2c3d48a
commit 103ff66a95
2 changed files with 26 additions and 6 deletions
+2 -2
View File
@@ -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
+24 -4
View File
@@ -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,)