diff --git a/envs.py b/envs.py new file mode 100644 index 0000000..79c42af --- /dev/null +++ b/envs.py @@ -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] diff --git a/lora_adapter.py b/lora_adapter.py new file mode 100644 index 0000000..0292fd2 --- /dev/null +++ b/lora_adapter.py @@ -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..lora_A.weight' and 'diffusion_model..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())