comfyui/viggle_turbo.py
| 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 | |