File size: 10,863 Bytes
032e1ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d50078f
 
 
032e1ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d50078f
 
 
 
 
 
 
 
 
 
 
 
 
032e1ee
 
d50078f
 
 
 
 
 
 
032e1ee
 
 
 
 
 
 
 
 
 
 
 
 
d50078f
032e1ee
 
 
 
 
 
 
d50078f
 
 
 
 
 
 
 
 
 
 
 
032e1ee
 
 
 
 
 
d50078f
032e1ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d50078f
 
032e1ee
 
d50078f
 
032e1ee
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
"""Run the quantizer from inside ComfyUI instead of a terminal.

Same code path as ``quantize_krea2.py`` on the command line -- this imports `convert` rather
than reimplementing it, so there is one quantizer, not two that drift.

The honest caveats, which the node's DESCRIPTION repeats because people do not read module
docstrings:

* It blocks the queue. 54 s for a single-shot split, ~5.7 min with the refinement loop, on a
  3090. Nothing else runs meanwhile.
* It needs the GPU to itself, so it unloads whatever ComfyUI is holding first. Your next
  generation will pay a reload.
* It writes ~8 GB.
"""

from __future__ import annotations

import logging
import os
import shutil

import comfy.model_management
import comfy.utils
import folder_paths

from .quantize_krea2 import (
    SAMPLER_HINTS,
    convert,
    derive_out_path,
    resolve_format,
)
from .svdquant_diag import _CATEGORY

# A source has to be dequantized to bf16 before it can be requantized, so leave a margin over
# the ~8 GB output rather than exactly it.
_MIN_FREE_BYTES = 12 * 1024 ** 3


def _free_bytes(path: str) -> int:
    return shutil.disk_usage(os.path.dirname(os.path.abspath(path))).free


class Krea2SVDQuantQuantize:
    @classmethod
    def INPUT_TYPES(cls):
        return {
            "required": {
                "source_model": (folder_paths.get_filename_list("diffusion_models"), {
                    "tooltip": "The BF16 Krea 2 checkpoint to quantize (~24 GB). An "
                               "already-quantized file cannot be used as a source, except "
                               "FP8, which is unpacked back to BF16 first.",
                }),
                "format": (["svdq", "w4a4", "int8", "fp8"], {
                    "default": "svdq",
                    "tooltip": "svdq: 4-bit weights and activations plus a low-rank bf16 "
                               "correction branch. w4a4: the same without the branch - "
                               "smaller and ~9%% faster per step. int8: the most faithful "
                               "option and still ~2x fp8 on Ampere. fp8: storage only.",
                }),
                "rank": ("INT", {
                    "default": 64, "min": 8, "max": 1024, "step": 8,
                    "tooltip": "svdq only: size of the low-rank branch. Only pays off with "
                               "refine_iters > 0. Without a LoRA, 64 / 128 / 256 measure the "
                               "same, so 64 is enough. With a LoRA loaded, 256 wins clearly "
                               "and 64 loses most of its advantage - so 256 if you use LoRAs.",
                }),
                "rank_alloc": (["uniform", "gqa"], {
                    "default": "uniform",
                    "tooltip": "svdq only: how the rank budget is spread across the eight "
                               "projection types. Same file size either way. uniform gives "
                               "every layer the same rank. gqa moves the budget to attn.wk / "
                               "attn.wv, which absorb ~2x the quantization error at a third of "
                               "the branch cost because Krea 2 has only 12 kv heads. Measured "
                               "and it does not pay: LPIPS 0.3523 vs 0.3403 for uniform, 5 of "
                               "10 prompts better, no mean effect. It does halve the spread "
                               "across prompts and improve the worst one. Leave on uniform "
                               "unless you are re-testing that.",
                }),
                "refine_iters": ("INT", {
                    "default": 100, "min": 0, "max": 200,
                    "tooltip": "svdq only. 0 is a single-shot SVD split (~54s); 100 refines "
                               "the branch against the quantization error and early-stops "
                               "(~5.7min). Keep this on if rank > 16: refinement is what "
                               "makes rank behave. Without it, raising rank costs file size "
                               "and buys nothing measurable.",
                }),
                "groupsize": ("INT", {
                    "default": 256, "min": 32, "max": 1024, "step": 32,
                    "tooltip": "convrot rotation group size. Unused for fp8.",
                }),
                "variant": (["turbo", "base", "unknown"], {
                    "default": "unknown",
                    "tooltip": "Which Krea 2 release this is. Affects only the output "
                               "filename and the recorded metadata - quantization is "
                               "identical; what differs is the sampler settings afterwards.",
                }),
                "output_name": ("STRING", {
                    "default": "",
                    "tooltip": "Filename inside models/diffusion_models/. Leave empty to "
                               "derive it from the variant and format.",
                }),
                "overwrite": ("BOOLEAN", {
                    "default": False,
                    "tooltip": "Off means an existing file of the same name is an error "
                               "rather than 8 GB written over your last run.",
                }),
            },
            # Optional so workflows saved before this input existed keep validating.
            "optional": {
                "act_stats": ("STRING", {
                    "default": "",
                    "tooltip": "svdq only: an activation-statistics file from the Capture "
                               "nodes (a bare filename is looked up in ComfyUI/output/). "
                               "Fits the low-rank branch against measured per-channel "
                               "activation energy instead of assuming it is uniform. Free at "
                               "inference and the best-measured setting here - LPIPS to BF16 "
                               "0.3378 to 0.2825 with no LoRA. Empty means the plain objective.",
                }),
            },
        }

    @classmethod
    def IS_CHANGED(cls, *args, **kwargs):
        # Writes an ~8 GB file as its side effect, so a cached summary would claim a file
        # exists that the user may have since deleted. Re-queueing is cheap to refuse
        # (overwrite=False still fails fast) and expensive to get wrong.
        return float("nan")

    RETURN_TYPES = ("STRING",)
    RETURN_NAMES = ("summary",)
    OUTPUT_TOOLTIPS = ("Where the checkpoint was written, and what went into it.",)
    OUTPUT_NODE = True
    FUNCTION = "run"
    CATEGORY = _CATEGORY
    TITLE = "Krea2 SVDQuant Quantize"
    DESCRIPTION = ("Builds a quantized Krea 2 checkpoint from a BF16 one, without leaving "
                   "ComfyUI. BLOCKS THE QUEUE while it runs (54s to ~6min), unloads any "
                   "loaded model to free the GPU, and writes ~8 GB. Load the result with the "
                   "Krea2 SVDQuant W4A4 Loader (svdq) or the stock UNETLoader (w4a4/int8/fp8).")

    def run(self, source_model, format, rank, rank_alloc, refine_iters, groupsize, variant,
            output_name, overwrite, act_stats=""):
        src = folder_paths.get_full_path_or_raise("diffusion_models", source_model)

        # `rank` always carries a value from the widget, so "was it set?" cannot be inferred
        # the way argparse does it. Non-svdq formats simply ignore it here rather than
        # erroring, which is the friendlier reading of a dropdown the user cannot un-set.
        fmt, rank = resolve_format(format, rank, rank_was_set=False)

        # Same validation the CLI does for --act-stats: a typed path that silently did nothing
        # would produce a checkpoint indistinguishable from a plain one.
        stats_path = act_stats.strip() or None
        if stats_path is not None:
            if format != "svdq":
                raise RuntimeError("act_stats only applies to format 'svdq': it weights the "
                                   "low-rank branch, and the other formats have no branch.")
            if not os.path.isabs(stats_path):
                stats_path = os.path.join(folder_paths.get_output_directory(), stats_path)
            if not os.path.isfile(stats_path):
                raise RuntimeError("act_stats file not found: {}".format(stats_path))

        if output_name.strip():
            name = output_name.strip()
            if not name.endswith(".safetensors"):
                name += ".safetensors"
            dst = os.path.join(os.path.dirname(src), name)
        else:
            dst, note = derive_out_path(src, format, rank, variant, rank_alloc, stats_path)
            if note:
                logging.info("[krea2-svdquant] %s", note)

        if os.path.exists(dst) and not overwrite:
            raise RuntimeError(
                "{} already exists. Enable 'overwrite', or set a different output_name."
                .format(dst))

        free = _free_bytes(dst)
        if free < _MIN_FREE_BYTES:
            raise RuntimeError(
                "only {:.1f} GB free on the drive holding {}; quantizing needs roughly {:.0f} "
                "GB of headroom.".format(free / 1024 ** 3, os.path.dirname(dst),
                                         _MIN_FREE_BYTES / 1024 ** 3))

        # Quantization wants the card to itself. Without this it competes with whatever the
        # last generation left resident and OOMs on the dequantize-to-bf16 step.
        comfy.model_management.unload_all_models()
        comfy.model_management.soft_empty_cache()

        pbar = comfy.utils.ProgressBar(1)
        state = {"total": 0}

        def progress(done, total, message):
            if total != state["total"]:
                state["total"] = total
                pbar.total = total
            pbar.update_absolute(done, total)

        logging.info("[krea2-svdquant] quantizing %s -> %s (format %s, rank %s/%s, "
                     "refine_iters %s, act_stats %s)", src, dst, fmt, rank, rank_alloc,
                     refine_iters if rank else 0, stats_path or "none")
        summary = convert(src, dst, fmt, groupsize, "cuda", rank, refine_iters,
                          variant=variant, progress_cb=progress,
                          rank_alloc=rank_alloc if rank else "uniform",
                          act_stats=stats_path)

        hint = SAMPLER_HINTS.get(variant)
        loader = ("Krea2 SVDQuant W4A4 Loader" if rank else "the stock UNETLoader")
        text = "\n".join(x for x in (
            summary, "Load it with {}.".format(loader), hint) if x)
        logging.info("[krea2-svdquant] %s", text)
        return {"ui": {"text": [text]}, "result": (text,)}


NODE_CLASS_MAPPINGS = {"Krea2SVDQuantQuantize": Krea2SVDQuantQuantize}
NODE_DISPLAY_NAME_MAPPINGS = {"Krea2SVDQuantQuantize": "Krea2 SVDQuant Quantize"}