add lora
This commit is contained in:
@@ -0,0 +1,54 @@
|
||||
import os
|
||||
from functools import lru_cache
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
### https://github.com/ModelTC/LightX2V ####
|
||||
|
||||
|
||||
|
||||
DTYPE_MAP = {
|
||||
"BF16": torch.bfloat16,
|
||||
"FP16": torch.float16,
|
||||
"FP32": torch.float32,
|
||||
"bf16": torch.bfloat16,
|
||||
"fp16": torch.float16,
|
||||
"fp32": torch.float32,
|
||||
"torch.bfloat16": torch.bfloat16,
|
||||
"torch.float16": torch.float16,
|
||||
"torch.float32": torch.float32,
|
||||
}
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def CHECK_ENABLE_PROFILING_DEBUG():
|
||||
ENABLE_PROFILING_DEBUG = os.getenv("ENABLE_PROFILING_DEBUG", "false").lower() == "true"
|
||||
return ENABLE_PROFILING_DEBUG
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def CHECK_ENABLE_GRAPH_MODE():
|
||||
ENABLE_GRAPH_MODE = os.getenv("ENABLE_GRAPH_MODE", "false").lower() == "true"
|
||||
return ENABLE_GRAPH_MODE
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def GET_RUNNING_FLAG():
|
||||
RUNNING_FLAG = os.getenv("RUNNING_FLAG", "infer")
|
||||
return RUNNING_FLAG
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def GET_DTYPE():
|
||||
RUNNING_FLAG = os.getenv("DTYPE", "BF16")
|
||||
assert RUNNING_FLAG in ["BF16", "FP16"]
|
||||
return DTYPE_MAP[RUNNING_FLAG]
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def GET_SENSITIVE_DTYPE():
|
||||
RUNNING_FLAG = os.getenv("SENSITIVE_LAYER_DTYPE", "None")
|
||||
if RUNNING_FLAG == "None":
|
||||
return GET_DTYPE()
|
||||
return DTYPE_MAP[RUNNING_FLAG]
|
||||
+133
@@ -0,0 +1,133 @@
|
||||
import gc
|
||||
import os
|
||||
|
||||
import torch
|
||||
from loguru import logger
|
||||
from safetensors import safe_open
|
||||
|
||||
from .envs import *
|
||||
|
||||
|
||||
### https://github.com/ModelTC/LightX2V ####
|
||||
|
||||
class WanLoraWrapper:
|
||||
def __init__(self, wan_model):
|
||||
self.model = wan_model
|
||||
self.lora_metadata = {}
|
||||
self.override_dict = {} # On CPU
|
||||
|
||||
def load_lora(self, lora_path, lora_name=None):
|
||||
if lora_name is None:
|
||||
lora_name = os.path.basename(lora_path).split(".")[0]
|
||||
|
||||
if lora_name in self.lora_metadata:
|
||||
logger.info(f"LoRA {lora_name} already loaded, skipping...")
|
||||
return lora_name
|
||||
|
||||
self.lora_metadata[lora_name] = {"path": lora_path}
|
||||
logger.info(f"Registered LoRA metadata for: {lora_name} from {lora_path}")
|
||||
|
||||
return lora_name
|
||||
|
||||
def _load_lora_file(self, file_path):
|
||||
with safe_open(file_path, framework="pt") as f:
|
||||
tensor_dict = {key: f.get_tensor(key).to(GET_DTYPE()) for key in f.keys()}
|
||||
return tensor_dict
|
||||
|
||||
def apply_lora(self, lora_name, alpha=1.0):
|
||||
if lora_name not in self.lora_metadata:
|
||||
logger.info(f"LoRA {lora_name} not found. Please load it first.")
|
||||
|
||||
# if not hasattr(self.model, "original_weight_dict"):
|
||||
# logger.error("Model does not have 'original_weight_dict'. Cannot apply LoRA.")
|
||||
# return False
|
||||
|
||||
lora_weights = self._load_lora_file(self.lora_metadata[lora_name]["path"])
|
||||
weight_dict_=self._apply_lora_weights( self.model.state_dict(), lora_weights, alpha)
|
||||
m, u = self.model.load_state_dict(weight_dict_, strict=False)
|
||||
|
||||
logger.info(f"Applied LoRA: {lora_name} with alpha={alpha}")
|
||||
del lora_weights
|
||||
return True
|
||||
|
||||
@torch.no_grad()
|
||||
def _apply_lora_weights(self, weight_dict, lora_weights, alpha):
|
||||
lora_pairs = {}
|
||||
lora_diffs = {}
|
||||
|
||||
def try_lora_pair(key, prefix, suffix_a, suffix_b, target_suffix):
|
||||
if key.endswith(suffix_a):
|
||||
base_name = key[len(prefix) :].replace(suffix_a, target_suffix)
|
||||
pair_key = key.replace(suffix_a, suffix_b)
|
||||
if pair_key in lora_weights:
|
||||
lora_pairs[base_name] = (key, pair_key)
|
||||
|
||||
def try_lora_diff(key, prefix, suffix, target_suffix):
|
||||
if key.endswith(suffix):
|
||||
base_name = key[len(prefix) :].replace(suffix, target_suffix)
|
||||
lora_diffs[base_name] = key
|
||||
|
||||
prefixs = [
|
||||
"", # empty prefix
|
||||
"diffusion_model.",
|
||||
]
|
||||
for prefix in prefixs:
|
||||
for key in lora_weights.keys():
|
||||
if not key.startswith(prefix):
|
||||
continue
|
||||
|
||||
try_lora_pair(key, prefix, "lora_A.weight", "lora_B.weight", "weight")
|
||||
try_lora_pair(key, prefix, "lora_down.weight", "lora_up.weight", "weight")
|
||||
try_lora_diff(key, prefix, "diff", "weight")
|
||||
try_lora_diff(key, prefix, "diff_b", "bias")
|
||||
try_lora_diff(key, prefix, "diff_m", "modulation")
|
||||
|
||||
applied_count = 0
|
||||
for name, param in weight_dict.items():
|
||||
if name in lora_pairs:
|
||||
if name not in self.override_dict:
|
||||
self.override_dict[name] = param.clone().cpu()
|
||||
name_lora_A, name_lora_B = lora_pairs[name]
|
||||
lora_A = lora_weights[name_lora_A].to(param.device, param.dtype)
|
||||
lora_B = lora_weights[name_lora_B].to(param.device, param.dtype)
|
||||
if param.shape == (lora_B.shape[0], lora_A.shape[1]):
|
||||
param += torch.matmul(lora_B, lora_A) * alpha
|
||||
applied_count += 1
|
||||
elif name in lora_diffs:
|
||||
if name not in self.override_dict:
|
||||
self.override_dict[name] = param.clone().cpu()
|
||||
|
||||
name_diff = lora_diffs[name]
|
||||
lora_diff = lora_weights[name_diff].to(param.device, param.dtype)
|
||||
if param.shape == lora_diff.shape:
|
||||
param += lora_diff * alpha
|
||||
applied_count += 1
|
||||
|
||||
logger.info(f"Applied {applied_count} LoRA weight adjustments")
|
||||
if applied_count == 0:
|
||||
logger.info(
|
||||
"Warning: No LoRA weights were applied. Expected naming conventions: 'diffusion_model.<layer_name>.lora_A.weight' and 'diffusion_model.<layer_name>.lora_B.weight'. Please verify the LoRA weight file."
|
||||
)
|
||||
return weight_dict
|
||||
|
||||
@torch.no_grad()
|
||||
def remove_lora(self):
|
||||
logger.info(f"Removing LoRA ...")
|
||||
|
||||
restored_count = 0
|
||||
for k, v in self.override_dict.items():
|
||||
self.model.original_weight_dict[k] = v.to(self.model.device)
|
||||
restored_count += 1
|
||||
|
||||
logger.info(f"LoRA removed, restored {restored_count} weights")
|
||||
|
||||
self.model._init_weights(self.model.original_weight_dict)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
self.lora_metadata = {}
|
||||
self.override_dict = {}
|
||||
|
||||
def list_loaded_loras(self):
|
||||
return list(self.lora_metadata.keys())
|
||||
Reference in New Issue
Block a user