Krea-2-SVDQuant-ComfyUI / svdquant_quantize.py
AlperKTS's picture
Sync docs and code: gallery, troubleshooting, cleanup pass
d50078f verified
Raw
History Blame Contribute Delete
10.9 kB
"""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"}