diff --git a/src/maxdiffusion/loaders/lora_base.py b/src/maxdiffusion/loaders/lora_base.py index f22696d3c..f52ef1531 100644 --- a/src/maxdiffusion/loaders/lora_base.py +++ b/src/maxdiffusion/loaders/lora_base.py @@ -24,6 +24,17 @@ class LoRABaseMixin: _lora_lodable_modules = [] num_fused_loras = 0 + def __init__(self): + self._fused_lora_keys = set() + + def _check_and_record_lora(self, lora_key): + """Return True if this LoRA was already merged (duplicate). Records it otherwise.""" + if lora_key in self._fused_lora_keys: + return True + self._fused_lora_keys.add(lora_key) + self.num_fused_loras += 1 + return False + def load_lora_weights(self, **kwargs): raise NotImplementedError("`load_lora_weights()` is not implemented.") diff --git a/src/maxdiffusion/loaders/ltx2_lora_nnx_loader.py b/src/maxdiffusion/loaders/ltx2_lora_nnx_loader.py index a3c4d0d38..7261fa917 100644 --- a/src/maxdiffusion/loaders/ltx2_lora_nnx_loader.py +++ b/src/maxdiffusion/loaders/ltx2_lora_nnx_loader.py @@ -54,6 +54,11 @@ def translate_fn(nnx_path_str): max_logging.log("No LoRA weight name provided; skipping LoRA load.") return pipeline + lora_key = (lora_model_path, transformer_weight_name) + if self._check_and_record_lora(lora_key): + max_logging.log(f"WARNING: LoRA '{lora_model_path}' already merged — skipping to avoid double-application.") + return pipeline + h_state_dict, _ = lora_loader.lora_state_dict(lora_model_path, weight_name=transformer_weight_name, **kwargs) transformer_state_dict = {} connector_state_dict = {} diff --git a/src/maxdiffusion/loaders/wan_lora_nnx_loader.py b/src/maxdiffusion/loaders/wan_lora_nnx_loader.py index a34c0f1a1..7e6598730 100644 --- a/src/maxdiffusion/loaders/wan_lora_nnx_loader.py +++ b/src/maxdiffusion/loaders/wan_lora_nnx_loader.py @@ -50,6 +50,11 @@ def load_lora_weights( def translate_fn(nnx_path_str): return lora_conversion_utils.translate_wan_nnx_path_to_diffusers_lora(nnx_path_str, scan_layers=scan_layers) + lora_key = (lora_model_path, transformer_weight_name) + if self._check_and_record_lora(lora_key): + max_logging.log(f"WARNING: LoRA '{lora_model_path}' already merged — skipping to avoid double-application.") + return pipeline + if hasattr(pipeline, "transformer") and transformer_weight_name: max_logging.log(f"Merging LoRA into transformer with rank={rank}") h_state_dict, _ = lora_loader.lora_state_dict(lora_model_path, weight_name=transformer_weight_name, **kwargs) @@ -91,7 +96,10 @@ def translate_fn(nnx_path_str: str): return lora_conversion_utils.translate_wan_nnx_path_to_diffusers_lora(nnx_path_str, scan_layers=scan_layers) # Handle high noise model - if hasattr(pipeline, "high_noise_transformer") and high_noise_weight_name: + high_key = (lora_model_path, high_noise_weight_name, "high_noise") + if self._check_and_record_lora(high_key): + max_logging.log(f"WARNING: LoRA '{lora_model_path}' already merged into high_noise_transformer — skipping.") + elif hasattr(pipeline, "high_noise_transformer") and high_noise_weight_name: max_logging.log(f"Merging LoRA into high_noise_transformer with rank={rank}") h_state_dict, _ = lora_loader.lora_state_dict(lora_model_path, weight_name=high_noise_weight_name, **kwargs) h_state_dict = lora_conversion_utils.preprocess_wan_lora_dict(h_state_dict) @@ -100,7 +108,10 @@ def translate_fn(nnx_path_str: str): max_logging.log("high_noise_transformer not found or no weight name provided for LoRA.") # Handle low noise model - if hasattr(pipeline, "low_noise_transformer") and low_noise_weight_name: + low_key = (lora_model_path, low_noise_weight_name, "low_noise") + if self._check_and_record_lora(low_key): + max_logging.log(f"WARNING: LoRA '{lora_model_path}' already merged into low_noise_transformer — skipping.") + elif hasattr(pipeline, "low_noise_transformer") and low_noise_weight_name: max_logging.log(f"Merging LoRA into low_noise_transformer with rank={rank}") l_state_dict, _ = lora_loader.lora_state_dict(lora_model_path, weight_name=low_noise_weight_name, **kwargs) l_state_dict = lora_conversion_utils.preprocess_wan_lora_dict(l_state_dict)