diff --git a/LucidFlux_node.py b/LucidFlux_node.py index 65a3003..8ab8217 100644 --- a/LucidFlux_node.py +++ b/LucidFlux_node.py @@ -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) diff --git a/__pycache__/LucidFlux_node.cpython-311.pyc b/__pycache__/LucidFlux_node.cpython-311.pyc deleted file mode 100644 index 7e96e2e..0000000 Binary files a/__pycache__/LucidFlux_node.cpython-311.pyc and /dev/null differ diff --git a/__pycache__/__init__.cpython-311.pyc b/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 2be7c51..0000000 Binary files a/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/__pycache__/inference.cpython-311.pyc b/__pycache__/inference.cpython-311.pyc deleted file mode 100644 index ec32921..0000000 Binary files a/__pycache__/inference.cpython-311.pyc and /dev/null differ diff --git a/__pycache__/model_loader_utils.cpython-311.pyc b/__pycache__/model_loader_utils.cpython-311.pyc deleted file mode 100644 index 65ec571..0000000 Binary files a/__pycache__/model_loader_utils.cpython-311.pyc and /dev/null differ diff --git a/example_workflows/example.png b/example_workflows/example.png index bacc993..0d48bee 100644 Binary files a/example_workflows/example.png and b/example_workflows/example.png differ diff --git a/example_workflows/lucidflux1210.json b/example_workflows/lucidflux0112.json similarity index 80% rename from example_workflows/lucidflux1210.json rename to example_workflows/lucidflux0112.json index 305cacc..9db5778 100644 --- a/example_workflows/lucidflux1210.json +++ b/example_workflows/lucidflux0112.json @@ -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 } \ No newline at end of file diff --git a/inference.py b/inference.py index 492e22e..21e114d 100644 --- a/inference.py +++ b/inference.py @@ -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) diff --git a/src/DiffBIR/__pycache__/__init__.cpython-311.pyc b/src/DiffBIR/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 7716e4c..0000000 Binary files a/src/DiffBIR/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/src/DiffBIR/__pycache__/inference.cpython-311.pyc b/src/DiffBIR/__pycache__/inference.cpython-311.pyc deleted file mode 100644 index a589108..0000000 Binary files a/src/DiffBIR/__pycache__/inference.cpython-311.pyc and /dev/null differ diff --git a/src/__pycache__/__init__.cpython-311.pyc b/src/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index 43561c8..0000000 Binary files a/src/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/__init__.cpython-311.pyc b/src/flux/__pycache__/__init__.cpython-311.pyc deleted file mode 100644 index c133c0b..0000000 Binary files a/src/flux/__pycache__/__init__.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/align_color.cpython-311.pyc b/src/flux/__pycache__/align_color.cpython-311.pyc deleted file mode 100644 index ebe0d88..0000000 Binary files a/src/flux/__pycache__/align_color.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/condition.cpython-311.pyc b/src/flux/__pycache__/condition.cpython-311.pyc deleted file mode 100644 index b4cac0b..0000000 Binary files a/src/flux/__pycache__/condition.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/math.cpython-311.pyc b/src/flux/__pycache__/math.cpython-311.pyc deleted file mode 100644 index 196b033..0000000 Binary files a/src/flux/__pycache__/math.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/model.cpython-311.pyc b/src/flux/__pycache__/model.cpython-311.pyc deleted file mode 100644 index 425dab4..0000000 Binary files a/src/flux/__pycache__/model.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/sampling.cpython-311.pyc b/src/flux/__pycache__/sampling.cpython-311.pyc deleted file mode 100644 index 35bc2c4..0000000 Binary files a/src/flux/__pycache__/sampling.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/swinir.cpython-311.pyc b/src/flux/__pycache__/swinir.cpython-311.pyc deleted file mode 100644 index 39b108a..0000000 Binary files a/src/flux/__pycache__/swinir.cpython-311.pyc and /dev/null differ diff --git a/src/flux/__pycache__/util.cpython-311.pyc b/src/flux/__pycache__/util.cpython-311.pyc deleted file mode 100644 index 884bc34..0000000 Binary files a/src/flux/__pycache__/util.cpython-311.pyc and /dev/null differ diff --git a/src/flux/model.py b/src/flux/model.py index e33ddf2..54e4c00 100644 --- a/src/flux/model.py +++ b/src/flux/model.py @@ -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 diff --git a/src/flux/modules/__pycache__/autoencoder.cpython-311.pyc b/src/flux/modules/__pycache__/autoencoder.cpython-311.pyc deleted file mode 100644 index 09bb58c..0000000 Binary files a/src/flux/modules/__pycache__/autoencoder.cpython-311.pyc and /dev/null differ diff --git a/src/flux/modules/__pycache__/conditioner.cpython-311.pyc b/src/flux/modules/__pycache__/conditioner.cpython-311.pyc deleted file mode 100644 index 22f05bd..0000000 Binary files a/src/flux/modules/__pycache__/conditioner.cpython-311.pyc and /dev/null differ diff --git a/src/flux/modules/__pycache__/layers.cpython-311.pyc b/src/flux/modules/__pycache__/layers.cpython-311.pyc deleted file mode 100644 index d765580..0000000 Binary files a/src/flux/modules/__pycache__/layers.cpython-311.pyc and /dev/null differ diff --git a/src/flux/modules/layers.py b/src/flux/modules/layers.py index 567a402..1d326e7 100644 --- a/src/flux/modules/layers.py +++ b/src/flux/modules/layers.py @@ -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): diff --git a/src/flux/sampling.py b/src/flux/sampling.py index 381b122..f28cb4c 100644 --- a/src/flux/sampling.py +++ b/src/flux/sampling.py @@ -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: diff --git a/src/flux/util.py b/src/flux/util.py index b6c878a..1ee8d45 100644 --- a/src/flux/util.py +++ b/src/flux/util.py @@ -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