From ed85d6ffcb3f930a0c4543ba5c8e033798f44c7b Mon Sep 17 00:00:00 2001 From: Sh0rck_Wang Date: Fri, 4 Sep 2026 14:04:26 +0800 Subject: [PATCH 1/2] fix: prevent silent double-application of LoRA weights Wire up the previously dead `num_fused_loras` counter in LoRABaseMixin to track merged LoRA identities. Before each merge, the loader checks whether the same (path, weight_name) was already applied and skips with a warning if so. This prevents accidental weight corruption when users list duplicate LoRA paths in config or re-run merge calls in notebooks. --- src/maxdiffusion/loaders/lora_base.py | 11 +++++++++++ src/maxdiffusion/loaders/ltx2_lora_nnx_loader.py | 5 +++++ src/maxdiffusion/loaders/wan_lora_nnx_loader.py | 15 +++++++++++++-- 3 files changed, 29 insertions(+), 2 deletions(-) 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) From 4962a2195dc20c28de4b4e5f69df08d4fc292102 Mon Sep 17 00:00:00 2001 From: Sh0rck_Wang Date: Fri, 4 Sep 2026 14:20:45 +0800 Subject: [PATCH 2/2] chore: retrigger CLA check