Add easy fluxLoader
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
#credit to comfyanonymous for this module
|
||||
#from https://github.com/comfyanonymous/ComfyUI_bitsandbytes_NF4
|
||||
import comfy.ops
|
||||
import torch
|
||||
import folder_paths
|
||||
from ..libs.utils import install_package
|
||||
|
||||
try:
|
||||
from bitsandbytes.nn.modules import Params4bit, QuantState
|
||||
except ImportError:
|
||||
Params4bit = torch.nn.Parameter
|
||||
raise ImportError("Please install bitsandbytes>=0.43.3")
|
||||
|
||||
def functional_linear_4bits(x, weight, bias):
|
||||
try:
|
||||
install_package("bitsandbytes", "0.43.3", True, "0.43.3")
|
||||
import bitsandbytes as bnb
|
||||
except ImportError:
|
||||
raise ImportError("Please install bitsandbytes>=0.43.3")
|
||||
|
||||
out = bnb.matmul_4bit(x, weight.t(), bias=bias, quant_state=weight.quant_state)
|
||||
out = out.to(x)
|
||||
return out
|
||||
|
||||
|
||||
def copy_quant_state(state, device: torch.device = None):
|
||||
if state is None:
|
||||
return None
|
||||
|
||||
device = device or state.absmax.device
|
||||
|
||||
state2 = (
|
||||
QuantState(
|
||||
absmax=state.state2.absmax.to(device),
|
||||
shape=state.state2.shape,
|
||||
code=state.state2.code.to(device),
|
||||
blocksize=state.state2.blocksize,
|
||||
quant_type=state.state2.quant_type,
|
||||
dtype=state.state2.dtype,
|
||||
)
|
||||
if state.nested
|
||||
else None
|
||||
)
|
||||
|
||||
return QuantState(
|
||||
absmax=state.absmax.to(device),
|
||||
shape=state.shape,
|
||||
code=state.code.to(device),
|
||||
blocksize=state.blocksize,
|
||||
quant_type=state.quant_type,
|
||||
dtype=state.dtype,
|
||||
offset=state.offset.to(device) if state.nested else None,
|
||||
state2=state2,
|
||||
)
|
||||
|
||||
|
||||
class ForgeParams4bit(Params4bit):
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
device, dtype, non_blocking, convert_to_format = torch._C._nn._parse_to(*args, **kwargs)
|
||||
if device is not None and device.type == "cuda" and not self.bnb_quantized:
|
||||
return self._quantize(device)
|
||||
else:
|
||||
n = ForgeParams4bit(
|
||||
torch.nn.Parameter.to(self, device=device, dtype=dtype, non_blocking=non_blocking),
|
||||
requires_grad=self.requires_grad,
|
||||
quant_state=copy_quant_state(self.quant_state, device),
|
||||
blocksize=self.blocksize,
|
||||
compress_statistics=self.compress_statistics,
|
||||
quant_type=self.quant_type,
|
||||
quant_storage=self.quant_storage,
|
||||
bnb_quantized=self.bnb_quantized,
|
||||
module=self.module
|
||||
)
|
||||
self.module.quant_state = n.quant_state
|
||||
self.data = n.data
|
||||
self.quant_state = n.quant_state
|
||||
return n
|
||||
|
||||
class ForgeLoader4Bit(torch.nn.Module):
|
||||
def __init__(self, *, device, dtype, quant_type, **kwargs):
|
||||
super().__init__()
|
||||
self.dummy = torch.nn.Parameter(torch.empty(1, device=device, dtype=dtype))
|
||||
self.weight = None
|
||||
self.quant_state = None
|
||||
self.bias = None
|
||||
self.quant_type = quant_type
|
||||
|
||||
def _save_to_state_dict(self, destination, prefix, keep_vars):
|
||||
super()._save_to_state_dict(destination, prefix, keep_vars)
|
||||
quant_state = getattr(self.weight, "quant_state", None)
|
||||
if quant_state is not None:
|
||||
for k, v in quant_state.as_dict(packed=True).items():
|
||||
destination[prefix + "weight." + k] = v if keep_vars else v.detach()
|
||||
return
|
||||
|
||||
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
|
||||
quant_state_keys = {k[len(prefix + "weight."):] for k in state_dict.keys() if k.startswith(prefix + "weight.")}
|
||||
|
||||
if any('bitsandbytes' in k for k in quant_state_keys):
|
||||
quant_state_dict = {k: state_dict[prefix + "weight." + k] for k in quant_state_keys}
|
||||
|
||||
self.weight = ForgeParams4bit().from_prequantized(
|
||||
data=state_dict[prefix + 'weight'],
|
||||
quantized_stats=quant_state_dict,
|
||||
requires_grad=False,
|
||||
device=self.dummy.device,
|
||||
module=self
|
||||
)
|
||||
self.quant_state = self.weight.quant_state
|
||||
|
||||
if prefix + 'bias' in state_dict:
|
||||
self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
|
||||
|
||||
del self.dummy
|
||||
elif hasattr(self, 'dummy'):
|
||||
if prefix + 'weight' in state_dict:
|
||||
self.weight = ForgeParams4bit(
|
||||
state_dict[prefix + 'weight'].to(self.dummy),
|
||||
requires_grad=False,
|
||||
compress_statistics=True,
|
||||
quant_type=self.quant_type,
|
||||
quant_storage=torch.uint8,
|
||||
module=self,
|
||||
)
|
||||
self.quant_state = self.weight.quant_state
|
||||
|
||||
if prefix + 'bias' in state_dict:
|
||||
self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
|
||||
|
||||
del self.dummy
|
||||
else:
|
||||
super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs)
|
||||
|
||||
current_device = None
|
||||
current_dtype = None
|
||||
current_manual_cast_enabled = False
|
||||
current_bnb_dtype = None
|
||||
|
||||
class OPS(comfy.ops.manual_cast):
|
||||
class Linear(ForgeLoader4Bit):
|
||||
def __init__(self, *args, device=None, dtype=None, **kwargs):
|
||||
super().__init__(device=device, dtype=dtype, quant_type=current_bnb_dtype)
|
||||
self.parameters_manual_cast = current_manual_cast_enabled
|
||||
|
||||
def forward(self, x):
|
||||
self.weight.quant_state = self.quant_state
|
||||
|
||||
if self.bias is not None and self.bias.dtype != x.dtype:
|
||||
# Maybe this can also be set to all non-bnb ops since the cost is very low.
|
||||
# And it only invokes one time, and most linear does not have bias
|
||||
self.bias.data = self.bias.data.to(x.dtype)
|
||||
|
||||
if not self.parameters_manual_cast:
|
||||
return functional_linear_4bits(x, self.weight, self.bias)
|
||||
elif not self.weight.bnb_quantized:
|
||||
assert x.device.type == 'cuda', 'BNB Must Use CUDA as Computation Device!'
|
||||
layer_original_device = self.weight.device
|
||||
self.weight = self.weight._quantize(x.device)
|
||||
bias = self.bias.to(x.device) if self.bias is not None else None
|
||||
out = functional_linear_4bits(x, self.weight, bias)
|
||||
self.weight = self.weight.to(layer_original_device)
|
||||
return out
|
||||
else:
|
||||
weight, bias, signal = weights_manual_cast(self, x, skip_weight_dtype=True, skip_bias_dtype=True)
|
||||
with main_stream_worker(weight, bias, signal):
|
||||
return functional_linear_4bits(x, weight, bias)
|
||||
+73
-2
@@ -918,7 +918,7 @@ class fullLoader:
|
||||
positive, positive_token_normalization, positive_weight_interpretation,
|
||||
negative, negative_token_normalization, negative_weight_interpretation,
|
||||
batch_size, model_override=None, clip_override=None, vae_override=None, optional_lora_stack=None, optional_controlnet_stack=None, a1111_prompt_style=False, prompt=None,
|
||||
my_unique_id=None
|
||||
my_unique_id=None, nf4=False
|
||||
):
|
||||
|
||||
# Clean models from loaded_objects
|
||||
@@ -926,7 +926,7 @@ class fullLoader:
|
||||
|
||||
# Load models
|
||||
log_node_warn("正在加载模型...")
|
||||
model, clip, vae, clip_vision, lora_stack = easyCache.load_main(ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt)
|
||||
model, clip, vae, clip_vision, lora_stack = easyCache.load_main(ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt, nf4=nf4)
|
||||
|
||||
# Create Empty Latent
|
||||
model_type = get_sd_version(model)
|
||||
@@ -1930,6 +1930,60 @@ class kolorsLoader:
|
||||
|
||||
return (chatglm3_model, None, None)
|
||||
|
||||
# Flux Loader
|
||||
class fluxLoader(fullLoader):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
checkpoints = folder_paths.get_filename_list("checkpoints")
|
||||
loras = ["None"] + folder_paths.get_filename_list("loras")
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (checkpoints,),
|
||||
"vae_name": (["Baked VAE"] + folder_paths.get_filename_list("vae"),),
|
||||
"lora_name": (loras,),
|
||||
"lora_model_strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
|
||||
"lora_clip_strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
|
||||
"resolution": (resolution_strings, {"default": "1024 x 1024"}),
|
||||
"empty_latent_width": ("INT", {"default": 1024, "min": 64, "max": MAX_RESOLUTION, "step": 8}),
|
||||
"empty_latent_height": ("INT", {"default": 1024, "min": 64, "max": MAX_RESOLUTION, "step": 8}),
|
||||
|
||||
"positive": ("STRING", {"default": "", "placeholder": "Positive", "multiline": True}),
|
||||
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 64}),
|
||||
},
|
||||
"optional": {
|
||||
"model_override": ("MODEL",),
|
||||
"clip_override": ("CLIP",),
|
||||
"vae_override": ("VAE",),
|
||||
"optional_lora_stack": ("LORA_STACK",),
|
||||
"optional_controlnet_stack": ("LORA_STACK",),
|
||||
},
|
||||
"hidden": {"prompt": "PROMPT", "my_unique_id": "UNIQUE_ID"}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PIPE_LINE", "MODEL", "VAE")
|
||||
RETURN_NAMES = ("pipe", "model", "vae")
|
||||
|
||||
FUNCTION = "fluxloader"
|
||||
CATEGORY = "EasyUse/Loaders"
|
||||
|
||||
def fluxloader(self, ckpt_name, vae_name,
|
||||
lora_name, lora_model_strength, lora_clip_strength,
|
||||
resolution, empty_latent_width, empty_latent_height,
|
||||
positive, batch_size, model_override=None, clip_override=None, vae_override=None, optional_lora_stack=None, optional_controlnet_stack=None,
|
||||
a1111_prompt_style=False, prompt=None,
|
||||
my_unique_id=None):
|
||||
|
||||
return super().adv_pipeloader(ckpt_name, 'Default', vae_name, 0,
|
||||
lora_name, lora_model_strength, lora_clip_strength,
|
||||
resolution, empty_latent_width, empty_latent_height,
|
||||
positive, 'none', 'comfy',
|
||||
'', 'none', 'comfy',
|
||||
batch_size, model_override, clip_override, vae_override, optional_lora_stack=optional_lora_stack,
|
||||
optional_controlnet_stack=optional_controlnet_stack,
|
||||
a1111_prompt_style=a1111_prompt_style, prompt=prompt,
|
||||
my_unique_id=my_unique_id, nf4=True)
|
||||
|
||||
|
||||
# Dit Loader
|
||||
from .dit.utils import string_to_dtype
|
||||
@@ -7510,6 +7564,21 @@ class stableDiffusion3API:
|
||||
|
||||
#---------------------------------------------------------------API 结束----------------------------------------------------------------------
|
||||
|
||||
class CheckpointLoaderNF4:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
|
||||
FUNCTION = "load_checkpoint"
|
||||
|
||||
CATEGORY = "loaders"
|
||||
|
||||
def load_checkpoint(self, ckpt_name):
|
||||
from .bitsandbytes_NF4 import OPS
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"), model_options={"custom_operations": OPS})
|
||||
return out[:3]
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
# seed 随机种
|
||||
@@ -7537,6 +7606,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"easy dynamiCrafterLoader": dynamiCrafterLoader,
|
||||
"easy cascadeLoader": cascadeLoader,
|
||||
"easy kolorsLoader": kolorsLoader,
|
||||
"easy fluxLoader": fluxLoader,
|
||||
"easy pixArtLoader": pixArtLoader,
|
||||
"easy loraStack": loraStack,
|
||||
"easy controlnetStack": controlnetStack,
|
||||
@@ -7657,6 +7727,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"easy dynamiCrafterLoader": "EasyLoader (DynamiCrafter)",
|
||||
"easy cascadeLoader": "EasyCascadeLoader",
|
||||
"easy kolorsLoader": "EasyLoader (Kolors)",
|
||||
"easy fluxLoader": "EasyLoader (Flux)",
|
||||
"easy hunyuanDiTLoader": "EasyLoader (HunyuanDiT)",
|
||||
"easy pixArtLoader": "EasyLoader (PixArt)",
|
||||
"easy loraStack": "EasyLoraStack",
|
||||
|
||||
@@ -12,6 +12,8 @@ def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_
|
||||
log_node_warn("正在进行" + title + "...")
|
||||
|
||||
if model_type in ['hydit', 'flux']:
|
||||
if model_type == 'flux':
|
||||
text = zh_to_en([text])[0] if has_chinese(text) else text
|
||||
embeddings_final, = CLIPTextEncode().encode(clip, text)
|
||||
return (embeddings_final, "", model, clip)
|
||||
|
||||
|
||||
+9
-5
@@ -221,7 +221,7 @@ class easyLoader:
|
||||
del self.loaded_objects[obj_type][item[0]]
|
||||
current_memory = self.get_memory_usage()
|
||||
|
||||
def load_checkpoint(self, ckpt_name, config_name=None, load_vision=False):
|
||||
def load_checkpoint(self, ckpt_name, config_name=None, load_vision=False, nf4=False):
|
||||
cache_name = ckpt_name
|
||||
if config_name not in [None, "Default"]:
|
||||
cache_name = ckpt_name + "_" + config_name
|
||||
@@ -238,7 +238,11 @@ class easyLoader:
|
||||
config_path = folder_paths.get_full_path("configs", config_name)
|
||||
loaded_ckpt = comfy.sd.load_checkpoint(config_path, ckpt_path, output_vae=True, output_clip=output_clip, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
||||
else:
|
||||
loaded_ckpt = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=output_clip, output_clipvision=output_clipvision, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
||||
model_options = {}
|
||||
if nf4:
|
||||
from ..bitsandbytes_NF4 import OPS
|
||||
model_options = {"custom_operations": OPS}
|
||||
loaded_ckpt = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=output_clip, output_clipvision=output_clipvision, embedding_directory=folder_paths.get_folder_paths("embeddings"), model_options=model_options)
|
||||
|
||||
self.add_to_cache("ckpt", cache_name, loaded_ckpt[0])
|
||||
self.add_to_cache("bvae", cache_name, loaded_ckpt[2])
|
||||
@@ -420,7 +424,7 @@ class easyLoader:
|
||||
|
||||
return None
|
||||
|
||||
def load_main(self, ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt):
|
||||
def load_main(self, ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt, nf4=False):
|
||||
model: ModelPatcher | None = None
|
||||
clip: comfy.sd.CLIP | None = None
|
||||
vae: comfy.sd.VAE | None = None
|
||||
@@ -438,7 +442,7 @@ class easyLoader:
|
||||
node = prompt[xy_model_id]
|
||||
if "ckpt_name_1" in node["inputs"]:
|
||||
ckpt_name_1 = node["inputs"]["ckpt_name_1"]
|
||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name_1)
|
||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name_1, nf4=nf4)
|
||||
can_load_lora = False
|
||||
# Load models
|
||||
elif model_override is not None and clip_override is not None and vae_override is not None:
|
||||
@@ -452,7 +456,7 @@ class easyLoader:
|
||||
elif clip_override is not None:
|
||||
raise Exception(f"[ERROR] model or vae is missing")
|
||||
else:
|
||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name, config_name)
|
||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name, config_name, nf4=nf4)
|
||||
|
||||
if optional_lora_stack is not None and can_load_lora:
|
||||
for lora in optional_lora_stack:
|
||||
|
||||
@@ -7,7 +7,7 @@ import { $t } from '../common/i18n.js';
|
||||
import { findWidgetByName, toggleWidget, updateNodeHeight} from "../common/utils.js";
|
||||
|
||||
const seedNodes = ["easy seed", "easy latentNoisy", "easy wildcards", "easy preSampling", "easy preSamplingAdvanced", "easy preSamplingNoiseIn", "easy preSamplingSdTurbo", "easy preSamplingCascade", "easy preSamplingDynamicCFG", "easy preSamplingLayerDiffusion", "easy fullkSampler", "easy fullCascadeKSampler"]
|
||||
const loaderNodes = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader", "easy hunyuanDiTLoader", "easy pixArtLoader"]
|
||||
const loaderNodes = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader", "easy fluxLoader", "easy hunyuanDiTLoader", "easy pixArtLoader"]
|
||||
|
||||
function widgetLogic(node, widget) {
|
||||
if (widget.name === 'lora_name') {
|
||||
@@ -697,6 +697,7 @@ app.registerExtension({
|
||||
switch (node.comfyClass){
|
||||
case "easy fullLoader":
|
||||
case "easy a1111Loader":
|
||||
case "easy fluxLoader":
|
||||
case "easy comfyLoader":
|
||||
case "easy cascadeLoader":
|
||||
case "easy svdLoader":
|
||||
|
||||
Reference in New Issue
Block a user