This commit is contained in:
smthemex
2026-01-12 12:39:16 +08:00
parent e1ff13a957
commit 4064e1ad20
26 changed files with 565 additions and 423 deletions
+20 -13
View File
@@ -38,6 +38,7 @@ class LucidFlux_SM_Model(io.ComfyNode):
inputs=[
io.Combo.Input("LucidFlux",options= ["none"] + [i for i in folder_paths.get_filename_list("LucidFlux") if "lucid" in i.lower()]),
io.Combo.Input("diffusion_models",options= ["none"] + folder_paths.get_filename_list("diffusion_models")),
io.Boolean.Input("block_offload", default=True),
io.Boolean.Input("use_accelerate", default=True),
io.Boolean.Input("use_quantize", default=False),
io.Boolean.Input("use_mmgp", default=False),
@@ -51,7 +52,7 @@ class LucidFlux_SM_Model(io.ComfyNode):
],
)
@classmethod
def execute(cls, LucidFlux,diffusion_models,use_accelerate,use_quantize,use_mmgp,mmgp_quantize,profile_number,cf_model=None) -> io.NodeOutput:
def execute(cls, LucidFlux,diffusion_models,block_offload,use_accelerate,use_quantize,use_mmgp,mmgp_quantize,profile_number,cf_model=None) -> io.NodeOutput:
is_dev="flux-dev" if "dev" in diffusion_models.lower() else "flux-schnell"
if cf_model is not None:
if "guidance_in.in_layer.weight" in cf_model.model.diffusion_model.state_dict().keys():
@@ -71,18 +72,24 @@ class LucidFlux_SM_Model(io.ComfyNode):
"checkpoint":LucidFlux_path,
}
args=OmegaConf.create(origin_dict)
model,state=load_lucidflux_model(args,ckpt_path,cf_model,use_accelerate,use_quantize,device,)
if use_quantize:
mmgp_quantize=False
if use_accelerate and use_mmgp:
print("if use accelerate ,can not use_mmgp")
use_mmgp=False
if use_mmgp:
from mmgp import offload
pipeline = { "transformer": model }
offload.profile(pipeline, quantizeTransformer = mmgp_quantize, profile_no = profile_number ) # uncomment this line and comment the previous one if you have 24 GB of VRAM and wants faster generation
state["use_mmgp"]=use_mmgp
model.use_mmgp=use_mmgp
model,state=load_lucidflux_model(args,ckpt_path,cf_model,use_accelerate,use_quantize,block_offload,device,)
if block_offload:
model.block_offload=True
model.use_mmgp=False
else:
model.block_offload=False
if use_quantize:
mmgp_quantize=False
if use_accelerate and use_mmgp:
print("if use accelerate ,can not use_mmgp")
use_mmgp=False
if use_mmgp:
from mmgp import offload
pipeline = { "transformer": model }
offload.profile(pipeline, quantizeTransformer = mmgp_quantize, profile_no = profile_number ) # uncomment this line and comment the previous one if you have 24 GB of VRAM and wants faster generation
state["use_mmgp"]=use_mmgp
model.use_mmgp=use_mmgp
return io.NodeOutput(model,state)
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 704 KiB

After

Width:  |  Height:  |  Size: 510 KiB

@@ -2,153 +2,8 @@
"id": "6e50dd9c-3346-49b3-b8fe-69c77b45da78",
"revision": 0,
"last_node_id": 44,
"last_link_id": 93,
"last_link_id": 95,
"nodes": [
{
"id": 28,
"type": "CheckpointLoaderSimple",
"pos": [
955.4044799804688,
-35.34019088745117
],
"size": [
270,
98
],
"flags": {},
"order": 0,
"mode": 4,
"inputs": [],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": []
},
{
"name": "CLIP",
"type": "CLIP",
"links": []
},
{
"name": "VAE",
"type": "VAE",
"links": []
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"flux1-dev-fp8.safetensors"
]
},
{
"id": 8,
"type": "DualCLIPLoader",
"pos": [
630.3765869140625,
14.808533668518066
],
"size": [
270,
130
],
"flags": {},
"order": 1,
"mode": 4,
"inputs": [],
"outputs": [
{
"name": "CLIP",
"type": "CLIP",
"links": [
59
]
}
],
"properties": {
"Node name for S&R": "DualCLIPLoader"
},
"widgets_values": [
"clip_l.safetensors",
"t5xxl_fp8_e4m3fn.safetensors",
"flux",
"default"
]
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
959.9752807617188,
135.83877563476562
],
"size": [
210,
115
],
"flags": {},
"order": 8,
"mode": 4,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 59
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": []
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"restore this image into high-quality, clean, high-resolution result"
]
},
{
"id": 5,
"type": "LoadImage",
"pos": [
306.2716979980469,
409.0683898925781
],
"size": [
406.7806701660156,
482.0671081542969
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
76
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"pasted/image (1).png",
"image"
]
},
{
"id": 9,
"type": "CLIPVisionLoader",
@@ -161,7 +16,7 @@
58
],
"flags": {},
"order": 3,
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
@@ -192,7 +47,7 @@
118
],
"flags": {},
"order": 13,
"order": 11,
"mode": 0,
"inputs": [
{
@@ -233,86 +88,6 @@
"prompt_embeddings.pt"
]
},
{
"id": 43,
"type": "LoadImage",
"pos": [
1299.713610202764,
454.24399790884786
],
"size": [
406.7806701660156,
482.0671081542969
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
88
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"pasted/image (2).png",
"image"
]
},
{
"id": 37,
"type": "LucidFlux_SM_Diffbir",
"pos": [
1252.0393931817712,
309.8838689453488
],
"size": [
270,
102
],
"flags": {},
"order": 10,
"mode": 4,
"inputs": [
{
"name": "model",
"type": "LucidFlux_SM_diff",
"link": 75
},
{
"name": "image",
"type": "IMAGE",
"link": 76
}
],
"outputs": [
{
"name": "Image",
"type": "IMAGE",
"links": [
78
]
}
],
"properties": {
"Node name for S&R": "LucidFlux_SM_Diffbir"
},
"widgets_values": [
1024,
1024
]
},
{
"id": 24,
"type": "VAELoader",
@@ -325,7 +100,7 @@
58
],
"flags": {},
"order": 5,
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
@@ -362,7 +137,7 @@
{
"name": "image",
"type": "IMAGE",
"link": 88
"link": 95
}
],
"outputs": [
@@ -383,6 +158,85 @@
1
]
},
{
"id": 25,
"type": "SaveImage",
"pos": [
2244.7917362971957,
-71.49938210778471
],
"size": [
628.4735705793846,
497.93119292626915
],
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 58
}
],
"outputs": [],
"properties": {},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 1,
"type": "LucidFlux_SM_Model",
"pos": [
1259.7284617132782,
-73.76634934610702
],
"size": [
270,
246
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "cf_model",
"shape": 7,
"type": "MODEL",
"link": null
}
],
"outputs": [
{
"name": "model",
"type": "LucidFlux_SM",
"links": [
69
]
},
{
"name": "state",
"type": "LucidFlux_SD",
"links": [
90
]
}
],
"properties": {
"Node name for S&R": "LucidFlux_SM_Model"
},
"widgets_values": [
"lucidflux.pth",
"flux1-kj-dev-fp8.safetensors",
true,
false,
true,
false,
2,
2
]
},
{
"id": 34,
"type": "LucidFlux_SM_Cond",
@@ -395,7 +249,7 @@
130
],
"flags": {},
"order": 11,
"order": 7,
"mode": 0,
"inputs": [
{
@@ -417,7 +271,7 @@
"Node name for S&R": "LucidFlux_SM_Cond"
},
"widgets_values": [
"flux-turbo.safetensors",
"FLUX\\flux-turbo.safetensors",
"none",
1,
1
@@ -435,7 +289,7 @@
214.264036462224
],
"flags": {},
"order": 15,
"order": 12,
"mode": 0,
"inputs": [
{
@@ -468,79 +322,101 @@
},
"widgets_values": [
8,
1267809222,
"fixed",
1995724915,
"randomize",
4,
true
]
},
{
"id": 25,
"type": "SaveImage",
"id": 5,
"type": "LoadImage",
"pos": [
2244.7917362971957,
-71.49938210778471
708.3115065345341,
276.5117775109095
],
"size": [
628.4735705793846,
497.93119292626915
406.7806701660156,
482.0671081542969
],
"flags": {},
"order": 16,
"order": 3,
"mode": 0,
"inputs": [
"inputs": [],
"outputs": [
{
"name": "images",
"name": "IMAGE",
"type": "IMAGE",
"link": 58
"links": [
76
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"outputs": [],
"properties": {},
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"ComfyUI"
"pasted/image (1).png",
"image"
]
},
{
"id": 42,
"type": "PreviewImage",
"id": 28,
"type": "CheckpointLoaderSimple",
"pos": [
1717.5609209670522,
471.19810934308623
923.3578097307902,
-23.686867162620967
],
"size": [
520.6141967773438,
478.749267578125
270,
98
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
"order": 4,
"mode": 4,
"inputs": [],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"link": 87
"name": "MODEL",
"type": "MODEL",
"links": []
},
{
"name": "CLIP",
"type": "CLIP",
"links": []
},
{
"name": "VAE",
"type": "VAE",
"links": []
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": []
"widgets_values": [
"flux1-dev-fp8.safetensors"
]
},
{
"id": 36,
"type": "LucidFlux_SM_Diff_Model",
"pos": [
667.47314453125,
247.47821044921875
899.0830285782389,
139.684935988144
],
"size": [
270,
82
],
"flags": {},
"order": 6,
"order": 5,
"mode": 4,
"inputs": [],
"outputs": [
@@ -561,24 +437,102 @@
]
},
{
"id": 38,
"type": "PreviewImage",
"id": 37,
"type": "LucidFlux_SM_Diffbir",
"pos": [
720.0827438542141,
448.66270642900895
1214.1660410654142,
317.1671812701698
],
"size": [
499.508056640625,
432.97503662109375
270,
102
],
"flags": {},
"order": 14,
"order": 8,
"mode": 4,
"inputs": [
{
"name": "model",
"type": "LucidFlux_SM_diff",
"link": 75
},
{
"name": "image",
"type": "IMAGE",
"link": 76
}
],
"outputs": [
{
"name": "Image",
"type": "IMAGE",
"links": []
}
],
"properties": {
"Node name for S&R": "LucidFlux_SM_Diffbir"
},
"widgets_values": [
1024,
1024
]
},
{
"id": 43,
"type": "LoadImage",
"pos": [
1156.464206693241,
458.59729502911176
],
"size": [
406.7806701660156,
482.0671081542969
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
95
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"pasted/image (2).png",
"image"
]
},
{
"id": 42,
"type": "PreviewImage",
"pos": [
1621.9264199886475,
467.52835860398983
],
"size": [
520.6141967773438,
478.749267578125
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 78
"link": 87
}
],
"outputs": [],
@@ -586,57 +540,6 @@
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 1,
"type": "LucidFlux_SM_Model",
"pos": [
1259.7284617132782,
-73.76634934610702
],
"size": [
270,
222
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "cf_model",
"shape": 7,
"type": "MODEL",
"link": null
}
],
"outputs": [
{
"name": "model",
"type": "LucidFlux_SM",
"links": [
69
]
},
{
"name": "state",
"type": "LucidFlux_SD",
"links": [
90
]
}
],
"properties": {
"Node name for S&R": "LucidFlux_SM_Model"
},
"widgets_values": [
"lucidflux.pth",
"flux1-kj-dev-fp8.safetensors",
false,
false,
true,
false,
2
]
}
],
"links": [
@@ -648,14 +551,6 @@
0,
"IMAGE"
],
[
59,
8,
0,
7,
0,
"CLIP"
],
[
60,
24,
@@ -696,14 +591,6 @@
1,
"IMAGE"
],
[
78,
37,
0,
38,
0,
"IMAGE"
],
[
87,
41,
@@ -712,14 +599,6 @@
0,
"IMAGE"
],
[
88,
43,
0,
41,
0,
"IMAGE"
],
[
90,
1,
@@ -751,6 +630,14 @@
31,
2,
"CONDITIONING"
],
[
95,
43,
0,
41,
0,
"IMAGE"
]
],
"groups": [
@@ -771,18 +658,18 @@
"config": {},
"extra": {
"ds": {
"scale": 0.8390545288824328,
"scale": 0.7627768444385501,
"offset": [
-499.8980837964146,
288.13297006172644
-283.9388023967938,
261.5338197476767
]
},
"frontendVersion": "1.32.9",
"frontendVersion": "1.35.9",
"workflowRendererVersion": "LG",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true,
"workflowRendererVersion": "LG"
"VHS_KeepIntermediate": true
},
"version": 0.4
}
+2 -2
View File
@@ -328,12 +328,12 @@ def infer_diffbir_model(model,input_pli_list,torch_device,):
image_tensor=torch.cat(images,dim=0)
return image_tensor
def load_lucidflux_model(args,ckpt_path,cf_model,use_accelerate,use_quantize,torch_device):
def load_lucidflux_model(args,ckpt_path,cf_model,use_accelerate,use_quantize,block_offload,torch_device):
name =args.name #"flux-dev"
#offload = args.offload
is_schnell = name == "flux-schnell"
model=load_flow_model(name,ckpt_path,cf_model,use_accelerate,use_quantize)
model=load_flow_model(name,ckpt_path,cf_model,use_accelerate,use_quantize,block_offload)
condition_lq=load_single_condition_branch(name, torch_device).to(torch.bfloat16)
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+216 -6
View File
@@ -9,6 +9,155 @@ from .modules.layers import (DoubleStreamBlock, EmbedND, LastLayer,
timestep_embedding)
class BlockGPUManager:
def __init__(self, device="cuda",):
self.device = device
self.managed_modules = []
self.embedder_modules = []
self.output_modules = []
self.fp8_scale=1
# 跟踪哪些blocks当前在GPU上
self.block_types = {}
self.blocks_on_gpu = set()
def get_gpu_memory_usage(self):
"""获取GPU内存使用情况"""
if torch.cuda.is_available():
torch.cuda.synchronize()
allocated = torch.cuda.memory_allocated() / 1024 / 1024 # MB
reserved = torch.cuda.memory_reserved() / 1024 / 1024 # MB
total = torch.cuda.get_device_properties(0).total_memory / 1024 / 1024 # MB
free = total - allocated
return {
'allocated_mb': allocated,
'reserved_mb': reserved,
'total_mb': total,
'free_mb': free
}
return None
def setup_for_inference(self, transformer_model,):
self._collect_managed_modules(transformer_model)
self._initialize_embedder_modules()
self._initialize_output_modules()
return self
def _has_fp8_parameters(self, module: nn.Module) -> bool:
"""检查模块是否包含FP8参数或缓冲区(带缓存)"""
if hasattr(module, '_has_fp8_cached'):
return module._has_fp8_cached
has_fp8 = False
for param in module.parameters(recurse=True):
if param.dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
has_fp8 = True
break
module._has_fp8_cached = has_fp8
return has_fp8
def _deep_convert_fp8_on_cpu(self, module: nn.Module):
"""在CPU上深度转换所有FP8参数"""
if not self._has_fp8_parameters(module):
return
params_to_convert = []
for submodule in module.modules():
for param in submodule.parameters(recurse=False):
if param.dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
params_to_convert.append(param)
# 批量执行转换
with torch.no_grad():
for param in params_to_convert:
param.data = param.to(torch.bfloat16) * self.fp8_scale
def _collect_managed_modules(self, transformer_model):
"""收集所有要管理的模块(简化版本,使用已知的block大小)"""
self.managed_modules = []
self.embedder_modules = []
self.output_modules = []
self.block_types = {}
# 收集dual blocks (double_blocks)
for i, block in enumerate(transformer_model.double_blocks):
self.managed_modules.append(block)
self.block_types[i] = 'dual'
# 收集single blocks (single_blocks)
single_start_idx = len(self.managed_modules)
for i, block in enumerate(transformer_model.single_blocks):
self.managed_modules.append(block)
self.block_types[single_start_idx + i] = 'single'
# 收集embedder和output模块
if hasattr(transformer_model, 'pe_embedder'):
self.embedder_modules.append(transformer_model.pe_embedder)
if hasattr(transformer_model, 'img_in'):
self.embedder_modules.append(transformer_model.img_in)
if hasattr(transformer_model, 'time_in'):
self.embedder_modules.append(transformer_model.time_in)
if hasattr(transformer_model, 'vector_in'):
self.embedder_modules.append(transformer_model.vector_in)
if hasattr(transformer_model, 'guidance_in'):
self.embedder_modules.append(transformer_model.guidance_in)
if hasattr(transformer_model, 'txt_in'):
self.embedder_modules.append(transformer_model.txt_in)
if hasattr(transformer_model, 'final_layer'):
self.output_modules.append(transformer_model.final_layer)
def _initialize_embedder_modules(self):
"""初始化embedder模块,将它们移到GPU"""
for module in self.embedder_modules:
# 先转换FP8参数
if self._has_fp8_parameters(module):
self._deep_convert_fp8_on_cpu(module)
if hasattr(module, 'to'):
module.to(self.device, non_blocking=True)
return self
def _initialize_output_modules(self):
"""初始化output模块,将它们移到GPU"""
for module in self.output_modules:
# 先转换FP8参数
if self._has_fp8_parameters(module):
self._deep_convert_fp8_on_cpu(module)
if hasattr(module, 'to'):
module.to(self.device, non_blocking=True)
return self
def unload_all_blocks_to_cpu(self):
"""卸载所有block到CPU"""
#print(f"[GPU Manager] 卸载所有block到CPU")
# 将所有模块移到CPU
for i, module in enumerate(self.managed_modules):
if hasattr(module, 'to'):
module.to('cpu')
for module in self.embedder_modules:
if hasattr(module, 'to'):
module.to('cpu')
for module in self.output_modules:
if hasattr(module, 'to'):
module.to('cpu')
# 清空GPU缓存
if torch.cuda.is_available():
torch.cuda.empty_cache()
@dataclass
class FluxParams:
in_channels: int
@@ -147,27 +296,60 @@ class Flux(nn.Module):
guidance: Tensor = None,
image_proj: Tensor = None,
ip_scale: float = 1.0,
gpu_manager=None,
) -> Tensor:
if img.ndim != 3 or txt.ndim != 3:
raise ValueError("Input img and txt tensors must have 3 dimensions.")
# running on sequences img
#original_device = img.device
img = self.img_in(img)
vec = self.time_in(timestep_embedding(timesteps, 256))
vec_=timestep_embedding(timesteps, 256)
if gpu_manager is not None:
if vec_.device != gpu_manager.device:
vec_ = vec_.to(gpu_manager.device)
vec = self.time_in(vec_)
if self.params.guidance_embed:
if guidance is None:
raise ValueError("Didn't get guidance strength for guidance distilled model.")
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
vec = vec + self.vector_in(y)
txt = self.txt_in(txt)
guidance=timestep_embedding(guidance, 256)
if gpu_manager is not None:
if guidance.device != gpu_manager.device:
guidance = guidance.to(gpu_manager.device)
vec_ = self.guidance_in(guidance)
vec = vec + vec_
vec = vec + self.vector_in(y)
txt = self.txt_in(txt)
ids = torch.cat((txt_ids, img_ids), dim=1)
pe = self.pe_embedder(ids)
pe = self.pe_embedder(ids)
if block_controlnet_hidden_states is not None:
controlnet_depth = len(block_controlnet_hidden_states)
# Double-stream blocks
# ------ Inside double_blocks loop ------
for index_block, block in enumerate(self.double_blocks):
if gpu_manager is not None:
# 加载当前block到GPU
if index_block < len(self.double_blocks):
module = gpu_manager.managed_modules[index_block]
if hasattr(module, 'to'):
if gpu_manager._has_fp8_parameters(module):
gpu_manager._deep_convert_fp8_on_cpu(module)
module.to(gpu_manager.device)
#print(f"[GPU Manager] 加载dual block {index_block}到GPU")
# 卸载上一个block(如果是第一个block,则不需要卸载)
if index_block > 0 and (index_block - 1) < len(self.double_blocks):
prev_module = gpu_manager.managed_modules[index_block - 1]
if hasattr(prev_module, 'to'):
prev_module.to('cpu')
#print(f"[GPU Manager] 卸载dual block {index_block-1}到CPU")
if self.training and self.gradient_checkpointing:
# Bind _block=block as default arg to avoid late binding
@@ -195,7 +377,35 @@ class Flux(nn.Module):
# ------ Inside single_blocks loop ------
img = torch.cat((txt, img), dim=1)
for block in self.single_blocks:
if gpu_manager is not None and len(self.double_blocks) > 0:
last_dual_idx = len(self.double_blocks) - 1
if last_dual_idx < len(gpu_manager.managed_modules):
module = gpu_manager.managed_modules[last_dual_idx]
if hasattr(module, 'to'):
module.to('cpu')
for block_index,block in enumerate(self.single_blocks):
if gpu_manager is not None:
# 计算当前single block在managed_modules中的索引(从19开始)
single_block_idx = len(self.double_blocks) + block_index
# 加载当前single block到GPU
if single_block_idx < len(gpu_manager.managed_modules):
module = gpu_manager.managed_modules[single_block_idx]
if hasattr(module, 'to'):
if gpu_manager._has_fp8_parameters(module):
gpu_manager._deep_convert_fp8_on_cpu(module)
module.to(gpu_manager.device)
#print(f"[GPU Manager] 加载single block {block_index} (全局索引 {single_block_idx})到GPU")
# 卸载上一个single block(如果是第一个single block,则不需要卸载)
if block_index > 0:
prev_single_idx = len(self.double_blocks) + (block_index - 1)
if prev_single_idx < len(gpu_manager.managed_modules):
prev_module = gpu_manager.managed_modules[prev_single_idx]
if hasattr(prev_module, 'to'):
prev_module.to('cpu')
#print(f"[GPU Manager] 卸载single block {block_index-1} (全局索引 {prev_single_idx})到CPU")
if self.training and self.gradient_checkpointing:
# Same binding trick for _block=block
Binary file not shown.
+7 -1
View File
@@ -57,7 +57,13 @@ class MLPEmbedder(nn.Module):
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
def forward(self, x: Tensor) -> Tensor:
return self.out_layer(self.silu(self.in_layer(x)))
x = self.in_layer(x)
# if x.dtype ==torch.float8_e4m3fn:
# x = x.to(torch.bfloat16)
x= self.silu(x)
# if x.dtype !=self.out_layer.weight.dtype:
# x = x.to(self.out_layer.weight.dtype)
return self.out_layer(x)
class RMSNorm(torch.nn.Module):
+11 -1
View File
@@ -141,6 +141,13 @@ def denoise_lucidflux(
condition_cond_ldr=None,
):
# this is ignored for schnell
offload_mode= model.block_offload if hasattr(model, 'block_offload') else False
if offload_mode:
from .model import BlockGPUManager
gpu_manager = BlockGPUManager(device="cuda")
gpu_manager.setup_for_inference(model)
else:
gpu_manager = None
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
timestep_pairs = zip(timesteps[:-1], timesteps[1:])
timestep_pairs = tqdm(timestep_pairs, total=len(timesteps)-1, desc="Denoising")
@@ -176,11 +183,14 @@ def denoise_lucidflux(
y=vec_in,
timesteps=t_vec,
guidance=guidance_vec,
block_controlnet_hidden_states=[i.to(img.device, dtype) for i in block_res_samples]
block_controlnet_hidden_states=[i.to(img.device, dtype) for i in block_res_samples],
gpu_manager=gpu_manager,
)
img = img + (t_prev - t_curr) * pred
if offload_mode:
gpu_manager.unload_all_blocks_to_cpu()
return img
def unpack(x: Tensor, height: int, width: int) -> Tensor:
+67 -45
View File
@@ -235,66 +235,88 @@ def load_from_repo_id(repo_id, checkpoint_name):
sd = load_sft(ckpt_path, device='cpu')
return sd
def load_flow_model(name,ckpt_path,cf_model,use_accelerate: bool = True,use_quantize: bool = False,device: str='cpu'):
def load_flow_model(name,ckpt_path,cf_model,use_accelerate: bool = True,use_quantize: bool = False,block_offload=True,device: str='cpu'):
# Loading Flux
print("Init model")
#ckpt_path = configs[name].ckpt_path
if use_accelerate:
if block_offload: #keep model on cpu
from contextlib import nullcontext
try:
from accelerate import init_empty_weights,load_checkpoint_and_dispatch
from accelerate import init_empty_weights
is_accelerate_available = True
except:
is_accelerate_available = False
ctx = init_empty_weights if is_accelerate_available else nullcontext
with ctx():
model = Flux(configs[name].params).to(torch.bfloat16)
#model = Flux(configs[name].params).to(torch.bfloat16)
if ckpt_path is not None:
model = load_checkpoint_and_dispatch(
model,
ckpt_path,
device_map="auto",
dtype=torch.bfloat16
)
sd = load_sft(ckpt_path, device='cpu')
else:
original_sd = cf_model.model.diffusion_model.state_dict()
load_checkpoint_and_dispatch_(
model,
original_sd,
device_map="auto",
dtype=torch.bfloat16
)
sd = cf_model.model.diffusion_model.state_dict()
del cf_model
del original_sd
gc.collect()
else:
model = Flux(configs[name].params).to(torch.bfloat16)
if use_quantize:
json_path = os.path.join(folder_paths.base_path, "custom_nodes/ComfyUI_LucidFlux/config.json") #config is for pass block
# Import here to avoid heavy quanto/transformers side effects at module import time
if ckpt_path is not None:
sd = load_sft(ckpt_path, device='cpu')
else:
sd = cf_model.model.diffusion_model.state_dict()
del cf_model
with open(json_path, "r") as f:
quantization_map = json.load(f)
print("Start a quantization process...")
from optimum.quanto import requantize
requantize(model, sd, quantization_map, device=device)
print("Model is quantized!")
else:
if ckpt_path is not None:
sd = load_sft(ckpt_path, device='cpu')
else:
sd = cf_model.model.diffusion_model.state_dict()
del cf_model
missing, unexpected = model.load_state_dict(sd, strict=False, assign=True)
print_load_warning(missing, unexpected)
missing, unexpected = model.load_state_dict(sd, strict=False, assign=True)
del sd
gc.collect()
print_load_warning(missing, unexpected)
else:
if use_accelerate:
from contextlib import nullcontext
try:
from accelerate import init_empty_weights,load_checkpoint_and_dispatch
is_accelerate_available = True
except:
is_accelerate_available = False
ctx = init_empty_weights if is_accelerate_available else nullcontext
with ctx():
model = Flux(configs[name].params).to(torch.bfloat16)
if ckpt_path is not None:
model = load_checkpoint_and_dispatch(
model,
ckpt_path,
device_map="auto",
dtype=torch.bfloat16
)
else:
original_sd = cf_model.model.diffusion_model.state_dict()
load_checkpoint_and_dispatch_(
model,
original_sd,
device_map="auto",
dtype=torch.bfloat16
)
del cf_model
del original_sd
gc.collect()
else:
model = Flux(configs[name].params).to(torch.bfloat16)
if use_quantize:
json_path = os.path.join(folder_paths.base_path, "custom_nodes/ComfyUI_LucidFlux/config.json") #config is for pass block
# Import here to avoid heavy quanto/transformers side effects at module import time
if ckpt_path is not None:
sd = load_sft(ckpt_path, device='cpu')
else:
sd = cf_model.model.diffusion_model.state_dict()
del cf_model
with open(json_path, "r") as f:
quantization_map = json.load(f)
print("Start a quantization process...")
from optimum.quanto import requantize
requantize(model, sd, quantization_map, device=device)
print("Model is quantized!")
else:
if ckpt_path is not None:
sd = load_sft(ckpt_path, device='cpu')
else:
sd = cf_model.model.diffusion_model.state_dict()
del cf_model
missing, unexpected = model.load_state_dict(sd, strict=False, assign=True)
print_load_warning(missing, unexpected)
del sd
gc.collect()
return model