Update bobs_lora_loader.py
This commit is contained in:
+98
-209
@@ -1,28 +1,18 @@
|
||||
# bobs_lora_loader.py
|
||||
# Bobs LoRA Loader (SDXL + FLUX) — ComfyUI 0.3.x compatible
|
||||
# - SDXL node supports block-weighted LoRA apply (input/middle/output) + Text Encoder strength
|
||||
# - FLUX node applies LoRA with fine-grained block controls mapped to Flux module name patterns
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Dict, Any, Tuple
|
||||
|
||||
from typing import Dict, Any, List
|
||||
import os
|
||||
import torch
|
||||
|
||||
import comfy.utils
|
||||
import comfy.lora
|
||||
import folder_paths
|
||||
import torch
|
||||
from safetensors.torch import load_file as safe_load_file
|
||||
|
||||
LOGGER = logging.getLogger("BobsLoRALoader")
|
||||
LOGGER.setLevel(logging.INFO)
|
||||
# -----------------------------------------------------------------------------#
|
||||
# BLOCK CONSTANTS #
|
||||
# -----------------------------------------------------------------------------#
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Flux block naming maps (for fine-grained control)
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
FLUX_BLOCK_NAME_MAPPING: Dict[str, list] = {
|
||||
FLUX_BLOCK_NAME_MAPPING = {
|
||||
"Text Conditioning": ["txt_in."],
|
||||
"Timestep Embedding": ["time_in."],
|
||||
"Image Hint": ["img_in."],
|
||||
@@ -40,7 +30,7 @@ FLUX_BLOCK_NAME_MAPPING: Dict[str, list] = {
|
||||
"Other Tensors": [],
|
||||
}
|
||||
|
||||
RAW_TO_CONCEPT_MAPPING: Dict[str, str] = {
|
||||
RAW_TO_CONCEPT_MAPPING = {
|
||||
raw: concept
|
||||
for concept, raws in FLUX_BLOCK_NAME_MAPPING.items()
|
||||
for raw in raws
|
||||
@@ -48,10 +38,6 @@ RAW_TO_CONCEPT_MAPPING: Dict[str, str] = {
|
||||
|
||||
ALL_FLUX_BLOCKS = list(FLUX_BLOCK_NAME_MAPPING.keys())
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# SDXL constants and presets
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
SDXL_TEXT_ENCODER = "Text Encoder"
|
||||
SDXL_INPUT_BLOCKS = "Input Blocks"
|
||||
SDXL_MIDDLE_BLOCK = "Middle Block"
|
||||
@@ -63,6 +49,10 @@ ALL_SDXL_BLOCKS = [
|
||||
SDXL_OUTPUT_BLOCKS,
|
||||
]
|
||||
|
||||
# -----------------------------------------------------------------------------#
|
||||
# PRESET DEFINITIONS #
|
||||
# -----------------------------------------------------------------------------#
|
||||
|
||||
LORA_BLOCK_PRESETS = {
|
||||
"FLUX": {
|
||||
"Custom": {},
|
||||
@@ -127,7 +117,7 @@ LORA_BLOCK_PRESETS = {
|
||||
},
|
||||
},
|
||||
},
|
||||
|
||||
|
||||
"SDXL": {
|
||||
"Custom": {},
|
||||
"Full (Normal LoRA)": {
|
||||
@@ -140,7 +130,7 @@ LORA_BLOCK_PRESETS = {
|
||||
SDXL_TEXT_ENCODER: 1.0,
|
||||
SDXL_INPUT_BLOCKS: 1.0,
|
||||
SDXL_MIDDLE_BLOCK: 1.0,
|
||||
SDXL_OUTPUT_BLOCKS: 0.2,
|
||||
SDXL_OUTPUT_BLOCKS:0.2,
|
||||
},
|
||||
},
|
||||
"Style": {
|
||||
@@ -149,7 +139,7 @@ LORA_BLOCK_PRESETS = {
|
||||
SDXL_TEXT_ENCODER: 0.0,
|
||||
SDXL_INPUT_BLOCKS: 0.2,
|
||||
SDXL_MIDDLE_BLOCK: 0.5,
|
||||
SDXL_OUTPUT_BLOCKS: 1.0,
|
||||
SDXL_OUTPUT_BLOCKS:1.0,
|
||||
},
|
||||
},
|
||||
"Concept": {
|
||||
@@ -158,7 +148,7 @@ LORA_BLOCK_PRESETS = {
|
||||
SDXL_TEXT_ENCODER: 1.0,
|
||||
SDXL_INPUT_BLOCKS: 0.8,
|
||||
SDXL_MIDDLE_BLOCK: 0.7,
|
||||
SDXL_OUTPUT_BLOCKS: 0.5,
|
||||
SDXL_OUTPUT_BLOCKS:0.5,
|
||||
},
|
||||
},
|
||||
"Fix Hands/Anatomy": {
|
||||
@@ -167,94 +157,16 @@ LORA_BLOCK_PRESETS = {
|
||||
SDXL_TEXT_ENCODER: 0.2,
|
||||
SDXL_INPUT_BLOCKS: 1.0,
|
||||
SDXL_MIDDLE_BLOCK: 0.4,
|
||||
SDXL_OUTPUT_BLOCKS: 0.0,
|
||||
SDXL_OUTPUT_BLOCKS:0.0,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# -----------------------------------------------------------------------------
|
||||
# -----------------------------------------------------------------------------#
|
||||
# FLUX LOADER (FIXED) #
|
||||
# -----------------------------------------------------------------------------#
|
||||
|
||||
def _build_key_map(model, clip) -> Dict[str, Any]:
|
||||
"""Build key_map compatibly across ComfyUI versions."""
|
||||
# Try unified helper
|
||||
if hasattr(comfy.lora, "model_lora_keys"):
|
||||
try:
|
||||
key_map, _ = comfy.lora.model_lora_keys(model, clip)
|
||||
return key_map
|
||||
except Exception as e:
|
||||
LOGGER.debug(f"Unified model_lora_keys failed, falling back: {e}")
|
||||
|
||||
# Fallback: per-module helpers
|
||||
key_map = {}
|
||||
try:
|
||||
if model is not None and hasattr(comfy.lora, "model_lora_keys_unet"):
|
||||
unet = getattr(model, "model", model)
|
||||
comfy.lora.model_lora_keys_unet(unet, key_map)
|
||||
if clip is not None and hasattr(comfy.lora, "model_lora_keys_clip"):
|
||||
clipm = getattr(clip, "cond_stage_model", clip)
|
||||
comfy.lora.model_lora_keys_clip(clipm, key_map)
|
||||
except Exception as e:
|
||||
LOGGER.warning(f"Failed to build LoRA key_map: {e}")
|
||||
return key_map
|
||||
|
||||
|
||||
def _invert_key_map(key_map: Dict[str, Any]) -> Dict[Any, str]:
|
||||
"""Invert key_map from raw_name -> key_tuple to key_tuple -> raw_name."""
|
||||
inv = {}
|
||||
for raw, key in key_map.items():
|
||||
inv[key] = raw
|
||||
return inv
|
||||
|
||||
|
||||
def _split_sdxl_unet_by_block(loaded_patches: Dict[Any, Any], inv_key_map: Dict[Any, str]):
|
||||
"""Split patches into SDXL input/middle/output/other groups based on raw weight names."""
|
||||
p_in, p_mid, p_out, p_other, p_clip = {}, {}, {}, {}, {}
|
||||
for key_tuple, patch in loaded_patches.items():
|
||||
raw = inv_key_map.get(key_tuple, "")
|
||||
if isinstance(raw, str) and raw.startswith("diffusion_model."):
|
||||
tail = raw[len("diffusion_model."):]
|
||||
if tail.startswith("input_blocks."):
|
||||
p_in[key_tuple] = patch
|
||||
elif tail.startswith("middle_block."):
|
||||
p_mid[key_tuple] = patch
|
||||
elif tail.startswith("output_blocks."):
|
||||
p_out[key_tuple] = patch
|
||||
else:
|
||||
p_other[key_tuple] = patch
|
||||
else:
|
||||
# Not diffusion_model.* -> likely CLIP/text enc
|
||||
p_clip[key_tuple] = patch
|
||||
return p_in, p_mid, p_out, p_other, p_clip
|
||||
|
||||
|
||||
def _group_flux_patches(loaded_patches: Dict[Any, Any], inv_key_map: Dict[Any, str]):
|
||||
"""Group Flux patches by conceptual block using name prefixes."""
|
||||
groups = {name: {} for name in ALL_FLUX_BLOCKS}
|
||||
for key_tuple, patch in loaded_patches.items():
|
||||
raw = inv_key_map.get(key_tuple, "")
|
||||
if not isinstance(raw, str):
|
||||
continue
|
||||
|
||||
if not raw.startswith("diffusion_model."):
|
||||
concept = "Text Conditioning"
|
||||
else:
|
||||
tail = raw[len("diffusion_model."):]
|
||||
concept = "Other Tensors"
|
||||
for prefix, cname in RAW_TO_CONCEPT_MAPPING.items():
|
||||
if tail.startswith(prefix):
|
||||
concept = cname
|
||||
break
|
||||
if concept in groups:
|
||||
groups[concept][key_tuple] = patch
|
||||
return groups
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# FLUX Loader
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
class BobsLoraLoaderFlux:
|
||||
def __init__(self):
|
||||
@@ -275,9 +187,10 @@ class BobsLoraLoaderFlux:
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP")
|
||||
FUNCTION = "apply_lora"
|
||||
CATEGORY = "Bobs/Loaders"
|
||||
CATEGORY = "Bobs_Nodes"
|
||||
|
||||
def apply_lora(self, model, clip, lora_name, strength, preset, **kwargs):
|
||||
|
||||
if lora_name == "None" or strength == 0.0:
|
||||
return model, clip
|
||||
|
||||
@@ -285,62 +198,81 @@ class BobsLoraLoaderFlux:
|
||||
if not lora_path:
|
||||
self.logger.error(f"[FLUX] LoRA file not found: {lora_name}")
|
||||
return model, clip
|
||||
|
||||
|
||||
|
||||
self.logger.info(f"[FLUX] Loading LoRA: {lora_name}")
|
||||
if os.path.splitext(lora_path)[1] == ".safetensors":
|
||||
lora_sd = safe_load_file(lora_path, device="cpu")
|
||||
else:
|
||||
lora_sd = torch.load(lora_path, map_location="cpu")
|
||||
|
||||
|
||||
# preset/custom weights
|
||||
# -------- build per-block final strengths ----------
|
||||
block_strength: Dict[str, float] = {}
|
||||
if preset == "Custom":
|
||||
for blk in ALL_FLUX_BLOCKS:
|
||||
block_strength[blk] = float(kwargs.get(blk, 1.0)) * float(strength)
|
||||
block_strength[blk] = kwargs.get(blk, 1.0) * strength
|
||||
else:
|
||||
cfg = LORA_BLOCK_PRESETS["FLUX"][preset]
|
||||
base = float(strength) * float(cfg.get("strength", 1.0))
|
||||
base = strength * cfg.get("strength", 1.0)
|
||||
weights = cfg.get("block_weights", {})
|
||||
for blk in ALL_FLUX_BLOCKS:
|
||||
block_strength[blk] = float(weights.get(blk, 1.0)) * base
|
||||
block_strength[blk] = weights.get(blk, 1.0) * base
|
||||
|
||||
# map + load patches
|
||||
key_map = _build_key_map(model, clip)
|
||||
if not key_map:
|
||||
self.logger.warning("[FLUX] Empty key_map; LoRA will not be applied.")
|
||||
return model, clip
|
||||
# -------- build key map ---------------------------
|
||||
key_map: Dict[str, Any] = {}
|
||||
key_map = comfy.lora.model_lora_keys_unet(model.model, key_map)
|
||||
key_map.update(comfy.lora.model_lora_keys_clip(clip.cond_stage_model, {}))
|
||||
mod_to_raw = {info[0]: raw for raw, info in key_map.items()}
|
||||
|
||||
loaded_patches = comfy.lora.load_lora(lora_sd, key_map)
|
||||
if not loaded_patches:
|
||||
# -------- load patches -----------------------------
|
||||
all_patches = comfy.lora.load_lora(lora_sd, key_map)
|
||||
if not all_patches:
|
||||
self.logger.warning("[FLUX] No matching keys in LoRA checkpoint.")
|
||||
return model, clip
|
||||
|
||||
inv_key_map = _invert_key_map(key_map)
|
||||
grouped = _group_flux_patches(loaded_patches, inv_key_map)
|
||||
# -------- Group patches by concept -----------------
|
||||
grouped_patches: Dict[str, Dict] = {name: {} for name in ALL_FLUX_BLOCKS}
|
||||
|
||||
for key_tuple, patch in all_patches.items():
|
||||
module = key_tuple[0]
|
||||
raw_key = mod_to_raw.get(module, "")
|
||||
|
||||
|
||||
if not raw_key.startswith("diffusion_model."):
|
||||
concept = "Text Conditioning"
|
||||
else:
|
||||
tail = raw_key[len("diffusion_model.") :]
|
||||
found_concept = "Other Tensors"
|
||||
for prefix, cname in RAW_TO_CONCEPT_MAPPING.items():
|
||||
if tail.startswith(prefix):
|
||||
found_concept = cname
|
||||
break
|
||||
concept = found_concept
|
||||
|
||||
if concept in grouped_patches:
|
||||
grouped_patches[concept][key_tuple] = patch
|
||||
|
||||
# -------- clone & attach patches group by group ------
|
||||
out_model = model.clone()
|
||||
out_clip = clip.clone()
|
||||
|
||||
# Apply per group with its strength
|
||||
for concept, patches_in_group in grouped.items():
|
||||
s = float(block_strength.get(concept, 0.0))
|
||||
if s == 0.0 or not patches_in_group:
|
||||
for concept, patches_in_group in grouped_patches.items():
|
||||
strength_for_group = block_strength.get(concept, 0.0)
|
||||
|
||||
if strength_for_group == 0.0 or not patches_in_group:
|
||||
continue
|
||||
try:
|
||||
out_model.add_patches(patches_in_group, strength_patch=s, strength_model=1.0)
|
||||
except Exception as e:
|
||||
self.logger.debug(f"[FLUX] model add_patches failed for '{concept}': {e}")
|
||||
try:
|
||||
out_clip.add_patches(patches_in_group, strength_patch=s, strength_model=1.0)
|
||||
except Exception as e:
|
||||
self.logger.debug(f"[FLUX] clip add_patches failed for '{concept}': {e}")
|
||||
|
||||
out_model.add_patches(patches_in_group, strength_for_group)
|
||||
out_clip.add_patches(patches_in_group, strength_for_group)
|
||||
|
||||
return out_model, out_clip
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# SDXL Loader
|
||||
# -----------------------------------------------------------------------------
|
||||
# -----------------------------------------------------------------------------#
|
||||
# SDXL LOADER (FIXED) #
|
||||
# -----------------------------------------------------------------------------#
|
||||
|
||||
|
||||
class BobsLoraLoaderSdxl:
|
||||
def __init__(self):
|
||||
@@ -350,7 +282,7 @@ class BobsLoraLoaderSdxl:
|
||||
def INPUT_TYPES(cls):
|
||||
req = {
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP",),
|
||||
"clip": ("CLIP",), # This is the CLIP (Text Encoder) input
|
||||
"lora_name": (["None"] + folder_paths.get_filename_list("loras"),),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -5.0, "max": 5.0, "step": 0.01}),
|
||||
"preset": (list(LORA_BLOCK_PRESETS["SDXL"].keys()),),
|
||||
@@ -359,11 +291,12 @@ class BobsLoraLoaderSdxl:
|
||||
req[blk] = ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.05})
|
||||
return {"required": req}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP")
|
||||
RETURN_TYPES = ("MODEL", "CLIP") # This is the CLIP (Text Encoder) output
|
||||
FUNCTION = "apply_lora"
|
||||
CATEGORY = "Bobs/Loaders"
|
||||
CATEGORY = "Bobs_Nodes"
|
||||
|
||||
def apply_lora(self, model, clip, lora_name, strength, preset, **kwargs):
|
||||
|
||||
if lora_name == "None" or strength == 0.0:
|
||||
return model, clip
|
||||
|
||||
@@ -377,71 +310,35 @@ class BobsLoraLoaderSdxl:
|
||||
lora_sd = safe_load_file(lora_path, device="cpu")
|
||||
else:
|
||||
lora_sd = torch.load(lora_path, map_location="cpu")
|
||||
|
||||
# --- PATCHED CODE STARTS HERE ---
|
||||
# Old, problematic line was: key_map, _ = comfy.lora.model_lora_keys(model, clip)
|
||||
# The new method below is compatible with modern ComfyUI versions.
|
||||
key_map = comfy.lora.model_lora_keys_unet(model.model, {})
|
||||
key_map.update(comfy.lora.model_lora_keys_clip(clip.cond_stage_model, {}))
|
||||
# --- PATCHED CODE ENDS HERE ---
|
||||
|
||||
# Strengths
|
||||
te_slider = float(kwargs.get(SDXL_TEXT_ENCODER, 1.0))
|
||||
in_slider = float(kwargs.get(SDXL_INPUT_BLOCKS, 1.0))
|
||||
mid_slider = float(kwargs.get(SDXL_MIDDLE_BLOCK, 1.0))
|
||||
out_slider = float(kwargs.get(SDXL_OUTPUT_BLOCKS, 1.0))
|
||||
|
||||
if preset != "Custom":
|
||||
cfg = LORA_BLOCK_PRESETS["SDXL"][preset]
|
||||
base = float(strength) * float(cfg.get("strength", 1.0))
|
||||
weights = cfg.get("block_weights", {})
|
||||
te_strength = float(weights.get(SDXL_TEXT_ENCODER, 1.0)) * base
|
||||
in_strength = float(weights.get(SDXL_INPUT_BLOCKS, 1.0)) * base
|
||||
mid_strength = float(weights.get(SDXL_MIDDLE_BLOCK, 1.0)) * base
|
||||
out_strength = float(weights.get(SDXL_OUTPUT_BLOCKS, 1.0)) * base
|
||||
block_strength: Dict[str, float] = {}
|
||||
if preset == "Custom":
|
||||
for blk in ALL_SDXL_BLOCKS:
|
||||
block_strength[blk] = kwargs.get(blk, 1.0) * strength
|
||||
else:
|
||||
base = float(strength)
|
||||
te_strength = te_slider * base
|
||||
in_strength = in_slider * base
|
||||
mid_strength = mid_slider * base
|
||||
out_strength = out_slider * base
|
||||
cfg = LORA_BLOCK_PRESETS["SDXL"][preset]
|
||||
base = strength * cfg.get("strength", 1.0)
|
||||
weights = cfg.get("block_weights", {})
|
||||
for blk in ALL_SDXL_BLOCKS:
|
||||
block_strength[blk] = weights.get(blk, 1.0) * base
|
||||
|
||||
# Map + load patches
|
||||
key_map = _build_key_map(model, clip)
|
||||
if not key_map:
|
||||
self.logger.warning("[SDXL] Empty key_map; LoRA will not be applied.")
|
||||
return model, clip
|
||||
|
||||
loaded_patches = comfy.lora.load_lora(lora_sd, key_map)
|
||||
if not loaded_patches:
|
||||
self.logger.warning("[SDXL] No matching keys in LoRA checkpoint.")
|
||||
return model, clip
|
||||
|
||||
inv_key_map = _invert_key_map(key_map)
|
||||
|
||||
p_in, p_mid, p_out, p_other, p_clip = _split_sdxl_unet_by_block(loaded_patches, inv_key_map)
|
||||
try:
|
||||
self.logger.info(f"[SDXL] LoRA keys: input={len(p_in)} middle={len(p_mid)} output={len(p_out)} other={len(p_other)} clip={len(p_clip)}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
out_model = model.clone()
|
||||
out_clip = clip.clone()
|
||||
|
||||
# Apply UNet patches
|
||||
if p_in and in_strength != 0.0:
|
||||
out_model.add_patches(p_in, strength_patch=in_strength, strength_model=1.0)
|
||||
if p_mid and mid_strength != 0.0:
|
||||
out_model.add_patches(p_mid, strength_patch=mid_strength, strength_model=1.0)
|
||||
if p_out and out_strength != 0.0:
|
||||
out_model.add_patches(p_out, strength_patch=out_strength, strength_model=1.0)
|
||||
if p_other:
|
||||
# time_embed etc. get base strength
|
||||
out_model.add_patches(p_other, strength_patch=float(strength), strength_model=1.0)
|
||||
|
||||
# Apply CLIP patches (Text Encoder)
|
||||
if p_clip and te_strength != 0.0:
|
||||
out_clip.add_patches(p_clip, strength_patch=te_strength, strength_model=1.0)
|
||||
|
||||
return out_model, out_clip
|
||||
|
||||
model_new, clip_new = comfy.lora.load_lora_for_models_with_block_weights(
|
||||
model, clip, comfy.lora.load_lora(lora_sd, key_map), 1.0, 1.0, block_strength
|
||||
)
|
||||
return (model_new, clip_new)
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Registration
|
||||
# -----------------------------------------------------------------------------
|
||||
# -----------------------------------------------------------------------------#
|
||||
# COMFYUI REGISTRATION #
|
||||
# -----------------------------------------------------------------------------#
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"BobsLoraLoaderFlux": BobsLoraLoaderFlux,
|
||||
@@ -452,11 +349,3 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"BobsLoraLoaderFlux": "Bobs LoRA Loader (FLUX)",
|
||||
"BobsLoraLoaderSdxl": "Bobs LoRA Loader (SDXL)",
|
||||
}
|
||||
|
||||
def _announce():
|
||||
try:
|
||||
print("✨ Bobs LoRA Loader nodes loaded! ✨")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
_announce()
|
||||
|
||||
Reference in New Issue
Block a user