it will hurt your feelings... and GPU

This commit is contained in:
cubiq
2024-09-20 17:25:29 +02:00
parent 508b893e1e
commit 4723f523e6
5 changed files with 807 additions and 0 deletions
+42
View File
@@ -0,0 +1,42 @@
# Flux blocks patcher sampler
This is an (very) advanced and (very) experimental custom node for the ComfyUI. It allows you to iteratively change the blocks weights of Flux models and check the difference each value makes.
## Usage
The `blocks` parameter accepts a list of blocks you want to iterate over, one per line. Each line should have the format `regex=weight`. For example the following lines will iterate over all the weights in the double_blocks first and then the single_blocks:
```
double_blocks\.([0-9]+)\.(img|txt)_(mod|attn|mlp\.[02])\.(lin|qkv|proj)\.(weight|bias)=1.1
single_blocks\.([0-9]+)\.(linear[12]|modulation\.lin)\.(weight|bias)=1.1
```
The regex above shows all the block options you have, but it can be as complex or simple as you want. For example if you want to target the whole `double_blocks 0`, you can use something like:
```
double_blocks\.0\.=1.1
```
To patch the `img` weights only of all double blocks you can use the following:
```
double_blocks\.([0-9]+)\.img_=1.1
```
To patch all the single blocks you can use:
```
single_blocks=1.1
```
The `Block Params Plot` node will then take the output of this iterator and plot the parameters directly onto the images.
## TODO
If there will be interest, I might add the following features:
- [ ] Speed up the patching process
- [ ] Support other models than Flux
- [ ] Make the patcher more user friendly, with a UI to edit the regex
- [ ] Better plotting of the blocks
+3
View File
@@ -0,0 +1,3 @@
from .blockpatcher import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+174
View File
@@ -0,0 +1,174 @@
from comfy_extras.nodes_custom_sampler import Noise_RandomNoise, BasicScheduler, BasicGuider, SamplerCustomAdvanced
from comfy_extras.nodes_latent import LatentBatch
from comfy_extras.nodes_model_advanced import ModelSamplingFlux, ModelSamplingAuraFlow
from node_helpers import conditioning_set_values
import comfy.samplers
import re
import os
from pathlib import Path
import torch
import torch.nn.functional as F
import torchvision.transforms.v2 as T
#import folder_paths
FONTS_DIR = os.path.join(os.path.dirname(os.path.realpath(__file__)), "fonts")
class FluxBlockPatcherSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL", ),
"conditioning": ("CONDITIONING", ),
"latent_image": ("LATENT", ),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 24, "min": 1, "max": 10000}),
"sampler": (comfy.samplers.KSampler.SAMPLERS, ),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
"guidance": ("FLOAT", {"default": 3.5, "min": -10.0, "max": 10.0, "step": 0.1}),
"blocks": ("STRING", { "multiline": True, "dynamicPrompts": True, "default": "double_blocks\.([0-9]+)\.(img|txt)_(mod|attn|mlp\.[02])\.(lin|qkv|proj)\.(weight|bias)=1.1\nsingle_blocks\.([0-9]+)\.(linear[12]|modulation\.lin)\.(weight|bias)=1.1" }),
}
}
RETURN_TYPES = ("LATENT", "SAMPLER_PARAMS", "STRING",)
RETURN_NAMES = ("latent", "sampler_params", "patched_blocks",)
FUNCTION = "apply_style"
def apply_style(self, model, conditioning, latent_image, noise_seed, steps, sampler, scheduler, guidance, blocks):
#is_schnell = model.model.model_type == comfy.model_base.ModelType.FLOW
sd = model.model_state_dict()
blocks = blocks.split("\n")
blocks = [b.strip() for b in blocks if b.strip()]
patched_blocks = []
fbi_params = []
out_latent = None
noise = Noise_RandomNoise(noise_seed)
sigmas = BasicScheduler().get_sigmas(model, scheduler, steps, 1.0)[0]
cond = conditioning_set_values(conditioning, {"guidance": guidance})
sca = SamplerCustomAdvanced()
latentbatch = LatentBatch()
samplerobject = comfy.samplers.sampler_object(sampler)
for b in blocks:
b = b.split("=")
block = b[0].strip()
value = float(b[1].strip())
m = model.clone()
out = { "regex": block, "value": value, "blocks": [] }
for k in sd:
if re.search(block, k):
m.add_patches({k: (None,)}, 0.0, value)
patched_blocks.append(f"{k}: {value}")
out["blocks"].append(k)
guider = BasicGuider().get_guider(m, cond)[0]
latent = sca.sample(noise, guider, samplerobject, sigmas, latent_image)[1]
fbi_params.append(out)
if out_latent is None:
out_latent = latent
else:
out_latent = latentbatch.batch(out_latent, latent)[0]
#m = None
#del m
patched_blocks = "\n".join(patched_blocks)
return (out_latent, fbi_params, patched_blocks)
class PlotBlockParams:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"images": ("IMAGE", ),
"params": ("SAMPLER_PARAMS", ),
"cols_num": ("INT", {"default": -1, "min": -1, "max": 1024 }),
"add_params": (["false", "true"], {"default": "true"}),
}}
RETURN_TYPES = ("IMAGE", )
FUNCTION = "execute"
CATEGORY = "essentials/sampling"
def execute(self, images, params, cols_num, add_params):
from PIL import Image, ImageDraw, ImageFont
import math
#import textwrap
if images.shape[0] != len(params):
raise ValueError("Number of images and number of parameters do not match.")
_params = params.copy()
if cols_num == 0:
cols_num = int(math.sqrt(images.shape[0]))
cols_num = max(1, min(cols_num, 1024))
width = images.shape[2]
out_image = []
font = ImageFont.truetype(os.path.join(FONTS_DIR, 'ShareTechMono-Regular.ttf'), min(32, int(20*(width/1024))))
text_padding = 3
line_height = font.getmask('Q').getbbox()[3] + font.getmetrics()[1] + text_padding*2
#char_width = font.getbbox('M')[2]+1 # using monospace font
for (image, param) in zip(images, _params):
image = image.permute(2, 0, 1)
if add_params != "false":
text = f"{param['regex']}: {param['value']}"
lines = text.split("\n")
text_height = line_height * len(lines)
text_image = Image.new('RGB', (width, text_height), color=(0, 0, 0))
for i, line in enumerate(lines):
draw = ImageDraw.Draw(text_image)
draw.text((text_padding, i * line_height + text_padding), line, font=font, fill=(255, 255, 255))
text_image = T.ToTensor()(text_image).to(image.device)
image = torch.cat([image, text_image], 1)
# a little cleanup
image = torch.nan_to_num(image, nan=0.0).clamp(0.0, 1.0)
out_image.append(image)
out_image = torch.stack(out_image, 0).permute(0, 2, 3, 1)
# merge images
if cols_num > -1:
cols = min(cols_num, out_image.shape[0])
b, h, w, c = out_image.shape
rows = math.ceil(b / cols)
# Pad the tensor if necessary
if b % cols != 0:
padding = cols - (b % cols)
out_image = F.pad(out_image, (0, 0, 0, 0, 0, 0, 0, padding))
b = out_image.shape[0]
# Reshape and transpose
out_image = out_image.reshape(rows, cols, h, w, c)
out_image = out_image.permute(0, 2, 1, 3, 4)
out_image = out_image.reshape(rows * h, cols * w, c).unsqueeze(0)
return (out_image, )
NODE_CLASS_MAPPINGS = {
"FluxBlockPatcherSampler": FluxBlockPatcherSampler,
"PlotBlockParams": PlotBlockParams,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FluxBlockPatcherSampler": "Flux Block Patcher Sampler",
"PlotBlockParams": "Plot Block Params",
}
Binary file not shown.
+588
View File
@@ -0,0 +1,588 @@
{
"last_node_id": 65,
"last_link_id": 137,
"nodes": [
{
"id": 11,
"type": "DualCLIPLoader",
"pos": {
"0": 138,
"1": 35
},
"size": {
"0": 210,
"1": 106
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "CLIP",
"type": "CLIP",
"links": [
10
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "DualCLIPLoader"
},
"widgets_values": [
"t5xxl_fp16.safetensors",
"clip_l.safetensors",
"flux"
]
},
{
"id": 12,
"type": "UNETLoader",
"pos": {
"0": 553,
"1": -111
},
"size": {
"0": 229.8605194091797,
"1": 82
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
131
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "UNETLoader"
},
"widgets_values": [
"flux1-dev.sft",
"fp8_e4m3fn"
]
},
{
"id": 6,
"type": "CLIPTextEncode",
"pos": {
"0": 415,
"1": 35
},
"size": {
"0": 366.7709045410156,
"1": 201.41677856445312
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 10
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
132
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"black cat sniffing a succulent plant"
]
},
{
"id": 59,
"type": "EmptySD3LatentImage",
"pos": {
"0": 468,
"1": 303
},
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
133
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "EmptySD3LatentImage"
},
"widgets_values": [
1024,
1024,
1
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": {
"0": 1377,
"1": -31
},
"size": {
"0": 140,
"1": 46
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 134
},
{
"name": "vae",
"type": "VAE",
"link": 12
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
135
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 65,
"type": "PlotBlockParams",
"pos": {
"0": 1584,
"1": 33
},
"size": [
210,
102
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 135
},
{
"name": "params",
"type": "SAMPLER_PARAMS",
"link": 136
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
137
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PlotBlockParams"
},
"widgets_values": [
0,
"true"
]
},
{
"id": 64,
"type": "FluxBlockPatcherSampler",
"pos": {
"0": 891,
"1": -32
},
"size": [
339.0323991530754,
299.821227993353
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 131
},
{
"name": "conditioning",
"type": "CONDITIONING",
"link": 132
},
{
"name": "latent_image",
"type": "LATENT",
"link": 133
}
],
"outputs": [
{
"name": "latent",
"type": "LATENT",
"links": [
134
],
"shape": 3,
"slot_index": 0
},
{
"name": "sampler_params",
"type": "SAMPLER_PARAMS",
"links": [
136
],
"shape": 3,
"slot_index": 1
},
{
"name": "patched_blocks",
"type": "STRING",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "FluxBlockPatcherSampler"
},
"widgets_values": [
0,
"fixed",
24,
"deis",
"beta",
3,
"double_blocks\\.0\\.=1.1\ndouble_blocks\\.1\\.=1.1\ndouble_blocks\\.2\\.=1.1\ndouble_blocks\\.3\\.=1.1"
]
},
{
"id": 10,
"type": "VAELoader",
"pos": {
"0": 1020,
"1": 319
},
"size": {
"0": 210,
"1": 58
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "VAE",
"type": "VAE",
"links": [
12
],
"slot_index": 0,
"shape": 3
}
],
"properties": {
"Node name for S&R": "VAELoader"
},
"widgets_values": [
"ae.sft"
]
},
{
"id": 26,
"type": "PreviewImage",
"pos": {
"0": 1851,
"1": 30
},
"size": [
865.0372553163784,
918.2602634191438
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 137
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
}
}
],
"links": [
[
10,
11,
0,
6,
0,
"CLIP"
],
[
12,
10,
0,
8,
1,
"VAE"
],
[
131,
12,
0,
64,
0,
"MODEL"
],
[
132,
6,
0,
64,
1,
"CONDITIONING"
],
[
133,
59,
0,
64,
2,
"LATENT"
],
[
134,
64,
0,
8,
0,
"LATENT"
],
[
135,
8,
0,
65,
0,
"IMAGE"
],
[
136,
64,
1,
65,
1,
"SAMPLER_PARAMS"
],
[
137,
65,
0,
26,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.630394086312919,
"offset": [
-175.11789190414686,
613.9746410020917
]
},
"groupNodes": {
"test": {
"nodes": [
{
"type": "PrimitiveNode",
"pos": [
88,
617
],
"size": {
"0": 210,
"1": 82
},
"flags": {},
"order": 5,
"mode": 0,
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [],
"widget": {
"name": "width"
}
}
],
"title": "width",
"properties": {
"Run widget replace on values": false
},
"index": 0
},
{
"type": "PrimitiveNode",
"pos": [
87,
745
],
"size": {
"0": 210,
"1": 82
},
"flags": {},
"order": 6,
"mode": 0,
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [],
"widget": {
"name": "height"
}
}
],
"title": "height",
"properties": {
"Run widget replace on values": false
},
"index": 1
},
{
"type": "EmptyLatentImage",
"pos": [
402,
596
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "width",
"type": "INT",
"link": null,
"widget": {
"name": "width"
},
"slot_index": 0
},
{
"name": "height",
"type": "INT",
"link": null,
"widget": {
"name": "height"
},
"slot_index": 1
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
1024,
1024,
1
],
"index": 2
}
],
"links": [
[
0,
0,
2,
0,
32,
"INT"
],
[
1,
0,
2,
1,
33,
"INT"
]
],
"external": [
[
2,
0,
"LATENT"
]
]
}
}
},
"version": 0.4
}