comfyui/viggle_turbo.py
5.3 KB · 122 lines · python Raw
1 """Qwen-Image-2.1 viggle-turbo nodes for ComfyUI.
2
3 ViggleTurboSigmas: same schedule as diffusers QwenImage21Pipeline(sigmas=nodes): FlowMatchEulerDiscreteScheduler with
4 dynamic exponential shift, mu = calculate_shift(tokens, 256, 8192, 0.5, 0.9), no shift_terminal, final 0.
5 tokens = (H/16) * (W/16) of the image being sampled, read from the latent that goes into the sampler.
6
7 ViggleTurboLora: the LoRA as a runtime side branch, y = W x + B A x (diffusers/PEFT without fuse_lora()).
8 LoraLoaderModelOnly merges it into the weights instead: on bf16 weights round-to-nearest keeps only ~70% of this
9 LoRA's update, on int8 weights the stochastic requantization keeps all of it but adds noise ~4x its size.
10 """
11
12 import json
13 import math
14
15 import torch
16 import torch.nn.functional as F
17
18 import comfy.patcher_extension
19 import comfy.utils
20 import folder_paths
21
22
23 class ViggleTurboSigmas:
24 @classmethod
25 def INPUT_TYPES(cls):
26 return {
27 "required": {
28 "latent": ("LATENT", {"tooltip": "The latent passed to the sampler; its size sets the resolution shift."}),
29 "nodes": ("STRING", {"default": "1.0, 0.9375, 0.875, 0.75, 0.5, 0.25",
30 "tooltip": "Raw (unshifted) student nodes. v0.3 LoRA: 1.0, 0.9375, 0.875, 0.75, 0.5, 0.25 (6 steps). Add or remove steps at the high-noise end only (5: 1.0, 0.875, ...; 7: 1.0, 0.9583, 0.9167, 0.875, ...); keep 0.875, 0.75, 0.5, 0.25."}),
31 }
32 }
33
34 RETURN_TYPES = ("SIGMAS",)
35 FUNCTION = "get_sigmas"
36 CATEGORY = "sampling/custom_sampling/schedulers"
37
38 def get_sigmas(self, latent, nodes):
39 s = latent["samples"]
40 r = latent.get("downscale_ratio_spacial", 16) / 16 # EmptyLatentImage is /8, the sampler resizes it to /16
41 tokens = round(s.shape[-2] * r) * round(s.shape[-1] * r)
42 mu = 0.5 + (0.9 - 0.5) * (tokens - 256) / (8192 - 256)
43 t = torch.tensor([float(x) for x in nodes.split(",")], dtype=torch.float64)
44 sigmas = math.exp(mu) / (math.exp(mu) + (1 / t - 1))
45 return (torch.cat([sigmas, sigmas.new_zeros(1)]).float(),)
46
47
48 def lora_fwd(x, ab):
49 return F.linear(F.linear(x, ab[0].to(x.dtype)), ab[1].to(x.dtype))
50
51
52 def add_hook(mod, ab):
53 return mod.register_forward_hook(lambda m, inp, out: out + lora_fwd(inp[0], ab))
54
55
56 def add_mlp_hooks(mlp, gate, up, down):
57 # fused SwiGLU: gate_up = [gate_layer; proj], and `out` runs inside an int8/fp16 kernel that bypasses its hooks,
58 # so its branch is added to the MLP output from the (LoRA'd) gate_up output
59 h = {}
60
61 def gate_up_hook(m, inp, out):
62 h["gu"] = out + torch.cat([lora_fwd(inp[0], gate), lora_fwd(inp[0], up)], -1)
63 return h["gu"]
64
65 def mlp_hook(m, inp, out):
66 g, u = h.pop("gu").chunk(2, -1)
67 return out + lora_fwd(F.silu(g) * u, down)
68
69 return [mlp.gate_up.register_forward_hook(gate_up_hook), mlp.register_forward_hook(mlp_hook)]
70
71
72 def run_with_lora(lora, executor, *args, **kwargs):
73 dm = executor.class_obj
74 for ab in lora.values():
75 if ab[0].device != args[0].device:
76 ab[0], ab[1] = ab[0].to(args[0].device), ab[1].to(args[0].device)
77 hooks = []
78 for name, ab in lora.items():
79 parent, _, leaf = name.rpartition(".")
80 if not getattr(dm.get_submodule(parent), "fused", False):
81 hooks.append(add_hook(dm.get_submodule(name), ab))
82 elif leaf == "out":
83 hooks += add_mlp_hooks(dm.get_submodule(parent), lora[parent + ".gate_layer"], lora[parent + ".proj"], ab)
84 try: # the diffusion model is shared with other MODEL outputs, the hooks must not outlive this call
85 return executor(*args, **kwargs)
86 finally:
87 for hk in hooks:
88 hk.remove()
89
90
91 class ViggleTurboLora:
92 @classmethod
93 def INPUT_TYPES(cls):
94 return {
95 "required": {
96 "model": ("MODEL",),
97 "lora_name": (folder_paths.get_filename_list("loras"), {"tooltip": "diffusers-format Qwen-Image-2.1 LoRA."}),
98 "strength": ("FLOAT", {"default": 1.0, "min": -4.0, "max": 4.0, "step": 0.05, "tooltip": "Keep 1.0 for viggle-turbo."}),
99 }
100 }
101
102 RETURN_TYPES = ("MODEL",)
103 FUNCTION = "load"
104 CATEGORY = "loaders"
105 DESCRIPTION = "Applies the LoRA at runtime (y = Wx + BAx) instead of merging it into the weights, which is lossy on bf16 and int8."
106
107 def load(self, model, lora_name, strength):
108 sd, meta = comfy.utils.load_torch_file(folder_paths.get_full_path_or_raise("loras", lora_name), return_metadata=True)
109 cfg = json.loads((meta or {}).get("lora_adapter_metadata", "{}"))
110 scale = strength * cfg.get("transformer.lora_alpha", 1) / cfg.get("transformer.r", 1)
111 lora = {k.removeprefix("transformer.").removesuffix(".lora_A.weight"): [sd[k], sd[k.replace("lora_A", "lora_B")] * scale]
112 for k in sd if k.endswith(".lora_A.weight")}
113 m = model.clone()
114 m.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, "viggle_turbo_lora",
115 lambda executor, *a, **kw: run_with_lora(lora, executor, *a, **kw))
116 return (m,)
117
118
119 NODE_CLASS_MAPPINGS = {"ViggleTurboSigmas": ViggleTurboSigmas, "ViggleTurboLora": ViggleTurboLora}
120 NODE_DISPLAY_NAME_MAPPINGS = {"ViggleTurboSigmas": "Qwen-Image-2.1 Viggle Turbo Sigmas",
121 "ViggleTurboLora": "Qwen-Image-2.1 Viggle Turbo LoRA (unmerged)"}
122