File size: 43,984 Bytes
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d50078f
45bcac4
 
 
 
 
 
 
 
 
 
 
 
d50078f
45bcac4
 
 
d50078f
 
 
 
45bcac4
 
 
 
 
 
 
 
d50078f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45bcac4
b0956bc
 
 
 
45bcac4
 
 
 
 
d50078f
45bcac4
 
 
 
 
 
 
 
 
 
 
d50078f
 
 
 
 
 
 
 
 
45bcac4
 
d50078f
 
 
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d50078f
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
032e1ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d4a8f34
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6f63bf4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45bcac4
b0956bc
 
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b0956bc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45bcac4
 
 
 
 
 
 
 
 
 
6f63bf4
 
 
 
 
 
45bcac4
 
 
 
 
 
 
6f63bf4
 
 
 
 
 
 
 
 
 
45bcac4
 
 
6f63bf4
 
45bcac4
 
 
 
 
 
6f63bf4
 
 
45bcac4
 
 
 
 
b0956bc
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6f63bf4
 
 
 
 
 
 
 
 
 
 
d50078f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45bcac4
032e1ee
9e39d3d
b0956bc
032e1ee
 
 
 
 
9e39d3d
 
 
 
 
 
032e1ee
 
 
 
 
 
45bcac4
 
 
 
 
032e1ee
 
 
6f63bf4
45bcac4
 
 
 
 
 
d50078f
 
45bcac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e39d3d
 
032e1ee
 
 
d4a8f34
 
 
032e1ee
 
6f63bf4
 
d50078f
 
6f63bf4
 
b0956bc
45bcac4
 
 
 
 
 
 
 
 
032e1ee
45bcac4
 
 
 
 
 
 
 
032e1ee
 
 
 
45bcac4
 
 
 
 
 
 
9e39d3d
 
 
 
45bcac4
 
 
 
 
 
032e1ee
45bcac4
 
 
 
 
 
 
 
 
 
032e1ee
 
 
 
 
 
 
 
 
 
 
45bcac4
 
 
 
 
 
 
 
 
 
032e1ee
 
 
45bcac4
 
 
 
032e1ee
 
45bcac4
 
 
 
 
 
 
6f63bf4
 
 
45bcac4
032e1ee
 
 
45bcac4
 
032e1ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6f63bf4
 
032e1ee
 
 
 
 
 
 
 
 
 
 
 
6f63bf4
 
 
 
032e1ee
 
45bcac4
 
 
 
 
 
 
 
 
032e1ee
 
 
 
 
 
45bcac4
 
 
b0956bc
 
 
 
 
 
 
 
45bcac4
 
 
 
6f63bf4
 
 
 
 
 
45bcac4
 
 
 
6f63bf4
 
 
 
 
 
 
032e1ee
 
 
 
 
 
 
 
 
 
45bcac4
 
 
6f63bf4
 
032e1ee
 
 
 
b0956bc
 
032e1ee
 
 
 
 
 
45bcac4
 
 
 
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
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
"""Quantize a BF16 Krea2 checkpoint into a ComfyUI-native quantized checkpoint.



ComfyUI already ships native kernels for these formats via comfy_kitchen, so the output

loads with a plain ``UNETLoader`` -- no custom node, and ordinary LoRA loaders work.



Four targets:



* ``int8``  -> ``int8_tensorwise`` with per-channel scales + convrot (W8A8).

               Natively accelerated on Ampere INT8 tensor cores.

* ``w4a4``  -> ``convrot_w4a4`` (W4A4), smaller and potentially faster, 4-bit quality.

* ``svdq``  -> the same W4A4 residual plus an SVDQuant low-rank bf16 branch, which

               absorbs the outlier directions 4 bits handle worst. Needs this repo's

               loader node; a plain ``UNETLoader`` cannot see the branch.

* ``fp8``   -> ``float8_e4m3fn``, a single per-tensor scale, no convrot, no calibration.

               Lightest touch of the four -- best fidelity, smallest speedup, and the

               only one that needs no rotation or outlier handling because 8-bit float

               already has enough dynamic range for these weights.



None of them need a calibration dataset: int8/w4a4/svdq spread outliers analytically via

the convrot (group-wise Hadamard) rotation, and activations are quantized by the kernel at

run time; fp8 just needs one abs-max scale per tensor.



Only the 224 transformer-block linears are quantized, matching the layer set used by the

reference ``krea2_raw_int8_convrot`` checkpoint. Norms, modulation, the text-fusion stack

and the final layer stay in their original precision -- they are small and quantization

noise there is disproportionately visible.



The layer selection is variant-agnostic: Krea 2 Turbo and the base (non-distilled)

release share the same block naming, so the same command converts either. What differs is

how you sample afterwards, not how you quantize -- pass ``--variant`` so the output is

named accordingly and the file records which one it came from.



    python quantize_krea2.py <bf16-model.safetensors> [--format int8|w4a4|svdq|fp8]

                             [--rank 64] [--variant turbo|base] [--out PATH]

"""

from __future__ import annotations

import argparse
import gc
import importlib.util
import json
import math
import os
import sys
import time

__version__ = "0.2.0"

import torch
from safetensors import safe_open
from safetensors.torch import save_file

def _find_comfyui_root() -> str | None:
    """Locate the ComfyUI root so `comfy.*` imports work regardless of cwd.



    Preference order: $COMFYUI_PATH, then walking up from this file (this script

    lives in ComfyUI/custom_nodes/<pkg>/, so comfy/ is two levels up). None if

    neither resolves -- deliberately *not* cwd. This module is imported by the

    Quantize node inside ComfyUI's own process, and putting cwd at sys.path[0]

    there lets any file in the working directory shadow a real import.

    """
    env = os.environ.get("COMFYUI_PATH")
    if env and os.path.isdir(os.path.join(env, "comfy")):
        return env
    here = os.path.dirname(os.path.abspath(__file__))
    candidate = os.path.abspath(os.path.join(here, "..", ".."))
    if os.path.isdir(os.path.join(candidate, "comfy")):
        return candidate
    return None


# Running inside ComfyUI, `comfy` already imports and the path needs no help.
if importlib.util.find_spec("comfy") is None:
    _root = _find_comfyui_root()
    if _root is not None:
        sys.path.insert(0, _root)

# Guarded because __init__.py imports this module transitively: an unguarded failure
# here would deregister every node in the pack, including Env Check, whose whole job is
# diagnosing a broken quantization backend. Degrade to "no formats available" instead.
QUANT_UNAVAILABLE = ""
try:
    from comfy.quant_ops import QUANT_ALGOS, get_layout_class  # noqa: E402
except Exception as exc:  # pragma: no cover - depends on the ComfyUI build
    QUANT_ALGOS = {}
    get_layout_class = None
    QUANT_UNAVAILABLE = "{}: {}".format(type(exc).__name__, exc)

# Relative improvement the refinement loop must keep making to earn another iteration. See
# `svdquant_split` for the traced numbers this default comes from.
REFINE_TOL = 0.001

_QUANT_SUFFIXES = ("attn.wq", "attn.wk", "attn.wv", "attn.gate", "attn.wo",
                   "mlp.gate", "mlp.up", "mlp.down")


# Checkpoints ship either bare ("blocks.0...") or prefixed ("model.diffusion_model.blocks.0...").
LAYER_PREFIXES = ("model.diffusion_model.", "diffusion_model.", "")

# Formats we know how to reconstruct back to BF16 from disk alone. FP8 storage is one
# byte per element with no packing or pre-rotation, so `qdata * scale` recovers the
# original tensor exactly as it was cast -- no architecture knowledge needed beyond
# what's already in the file. INT8/W4A4 sources pack multiple values per byte and are
# rotated (convrot) before quantization; unpacking those correctly needs the exact
# in/out feature counts from the live nn.Linear, which isn't recoverable from the
# checkpoint alone, so those are rejected instead.
_DEQUANTIZABLE_FORMATS = ("float8_e4m3fn", "float8_e5m2")


def detect_prefix(keys, default: str | None = None) -> str:
    """Return the prefix the transformer blocks live under.



    Raises `RuntimeError` if no prefix matches and no `default` was given. Not

    `SystemExit`: that is a `BaseException`, and `convert()` runs inside ComfyUI's

    executor via the Quantize node, where a BaseException escapes error handling

    entirely. `main()` catches RuntimeError and prints it as a clean CLI line.

    """
    for prefix in LAYER_PREFIXES:
        if any(k.startswith("{}blocks.".format(prefix)) for k in keys):
            return prefix
    if default is not None:
        return default
    raise RuntimeError(
        "Could not find transformer blocks in this checkpoint. Expected keys like\n"
        "  blocks.0.attn.wq.weight  or  model.diffusion_model.blocks.0.attn.wq.weight\n"
        "This does not look like a Krea 2 diffusion model."
    )


def check_requantizable(handle, keys, prefix: str) -> None:
    """Make sure every target layer that's already quantized is something we can

    dequantize back to BF16 (see `_DEQUANTIZABLE_FORMATS`). Anything else -- most

    importantly INT8/W4A4 -- needs the original BF16 (or FP16) release instead.

    """
    for key in keys:
        if not key.endswith(".comfy_quant"):
            continue
        layer = key[: -len(".comfy_quant")]
        if not is_target(layer, prefix):
            continue
        conf = json.loads(bytes(handle.get_tensor(key).tolist()))
        fmt = conf.get("format")
        if fmt not in _DEQUANTIZABLE_FORMATS:
            raise RuntimeError(
                "Layer {} is already quantized as '{}'. Only FP8-quantized layers can be "
                "automatically reconstructed and re-quantized; for INT8/W4A4 sources, use "
                "the original BF16 (or FP16) release of the model instead.".format(layer, fmt)
            )


def dequantize_target_weight(handle, layer: str, device: str) -> torch.Tensor:
    """Load a target layer's weight as BF16, dequantizing it first if it's FP8.



    Scaled FP8 (has a `comfy_quant` marker + `weight_scale`): ``qdata.float() * scale``.

    Unscaled FP8 (a bare ``.to(float8_e4m3fn)`` cast, no marker): ``qdata`` as-is.

    Anything else is already ruled out by `check_requantizable`.

    """
    weight = handle.get_tensor("{}.weight".format(layer)).to(device=device)
    conf_key = "{}.comfy_quant".format(layer)
    if conf_key in handle.keys():
        scale = handle.get_tensor("{}.weight_scale".format(layer)).to(device=device).float()
        return (weight.to(torch.float32) * scale).to(torch.bfloat16)
    return weight.to(torch.bfloat16)


def is_target(layer: str, prefix: str) -> bool:
    """Only the transformer blocks; txtfusion and friends stay high precision."""
    if not layer.startswith("{}blocks.".format(prefix)):
        return False
    return any(layer.endswith(s) for s in _QUANT_SUFFIXES)


def leaf_name(layer: str, prefix: str) -> str | None:
    """``blocks.7.attn.wq`` -> ``attn.wq``. None if the layer is not inside a block.



    Only used to make a "nothing matched" failure self-diagnosing: a Krea 2 variant that

    keeps the ``blocks.`` container but renames its leaves would otherwise quantize zero

    layers and still write a full-size file.

    """
    head = "{}blocks.".format(prefix)
    if not layer.startswith(head):
        return None
    rest = layer[len(head):]
    return rest.split(".", 1)[1] if "." in rest else None


# Rank is bought in multiples of this, so a reallocated budget lands on the same grid the
# `rank` widget and the CLI already use.
_RANK_STEP = 8

# Per-leaf rank multipliers, relative to a uniform budget.
#
# Uniform rank is the obvious default and the wrong one. Measured on Krea2 Turbo at rank 64,
# error removed per million branch parameters spans 6.9x across the eight leaves:
#
#     attn.wk  0.0992   (1536x6144, branch 491K params, absorbs 30.1% of the error)
#     attn.wv  0.0844   (1536x6144, 491K, 25.6%)
#     attn.wq  0.0445   (6144x6144, 786K, 21.6%)
#     attn.wo  0.0411   (6144x6144, 786K, 17.6%)
#     attn.gate 0.0384  (6144x6144, 786K, 18.7%)
#     mlp.down 0.0169   (6144x16384, 1.44M, 14.3%)
#     mlp.gate 0.0155   (16384x6144, 1.44M, 13.8%)
#     mlp.up   0.0143   (16384x6144, 1.44M, 12.7%)
#
# The cause is GQA: Krea 2 has 12 kv heads against 48 query heads, so wk/wv are 1536-wide and
# their branch costs a third of an MLP branch -- while absorbing twice as much error. Uniform
# rank therefore spends most of the budget on the layers that reward it least.
#
# These multipliers come from a greedy allocation under a fixed byte budget (branch cost is
# exactly linear in rank, so scaling all of them preserves byte-neutrality at any base rank).
# At base rank 64 they reproduce the solved table: wk 360, wv 256, wq 72, wo 64, gate 56,
# mlp 8, and the checkpoint comes out 0.02% smaller than uniform rank 64.
#
# MEASURED, and it does not deliver. The greedy solve predicted 6% less mean layer weight
# error; the images say otherwise. 10 prompts, seed 987654321, against the same BF16 reference:
#
#     uniform rank 64  LPIPS 0.3403 +-0.1019
#     gqa     rank 64  LPIPS 0.3523 +-0.0473
#
# gqa is closer on 5 of 10 prompts, paired t = +0.55 on 9 df. No mean effect in either
# direction. Two things did not survive scrutiny and are recorded so nobody re-derives them:
# the -6% rests on an absorption ~ sqrt(rank) model fitted to nothing stronger than "returns
# are sublinear", and the tempting pattern that gqa wins on the hard prompts is an artifact --
# re-pairing the gqa scores at random reproduces the same -0.92 correlation, because the
# baseline appears on both sides of the difference.
#
# What is left is real but narrow: the spread across prompts shrinks (variance ratio 4.64,
# F-test p = 0.032) and the worst prompt improves, 0.4975 -> 0.4470. That is one nominally
# significant result among several things tried on the same 10 images, so it is a lead for a
# larger measurement, not a reason to change the default. `uniform` stays the default.
#
# The 6.9x spread itself is solid -- it is measured weight error over real shapes. The lesson
# is that weight error is a poor proxy for image outcome, which is now the third time this
# has held on this model (the refinement objective and the depth hypothesis were the others).
_GQA_ALLOC = {
    "attn.wk": 5.625, "attn.wv": 4.0, "attn.wq": 1.125, "attn.wo": 1.0, "attn.gate": 0.875,
    "mlp.gate": 0.125, "mlp.up": 0.125, "mlp.down": 0.125,
}

RANK_ALLOCATIONS: dict[str, dict[str, float] | None] = {
    "uniform": None,
    "gqa": _GQA_ALLOC,
}


def leaf_ranks(rank: int, rank_alloc: str = "uniform") -> dict[str, int]:
    """Expand a rank budget into a per-leaf rank, one entry per name in `_QUANT_SUFFIXES`.



    `uniform` gives every leaf the same rank -- the historical behaviour. Other allocations

    redistribute the same total branch bytes according to `RANK_ALLOCATIONS`.

    """
    if rank_alloc not in RANK_ALLOCATIONS:
        raise RuntimeError("unknown rank allocation {!r}; expected one of {}".format(
            rank_alloc, ", ".join(sorted(RANK_ALLOCATIONS))))
    mults = RANK_ALLOCATIONS[rank_alloc]
    if mults is None:
        return {leaf: rank for leaf in _QUANT_SUFFIXES}

    # Byte-neutrality is the whole point of a non-uniform allocation, and it only survives
    # while every leaf can be expressed on the step-8 grid. Below that the smallest leaves get
    # rounded *up* to the floor and the file grows: at budget 16 the MLP leaves want rank 2,
    # take 8, and the checkpoint comes out 20% larger than uniform -- which would make any
    # comparison against uniform meaningless. Refuse rather than quietly change the size.
    smallest = min(mults.values())
    minimum = int(math.ceil(_RANK_STEP / smallest / _RANK_STEP)) * _RANK_STEP
    if rank < minimum:
        raise RuntimeError(
            "rank allocation {!r} needs a budget of at least {} (you asked for {}). Its "
            "smallest leaf gets {}x the budget, which below that rounds up to the step-{} "
            "floor and makes the checkpoint *larger* than uniform instead of the same size."
            .format(rank_alloc, minimum, rank, smallest, _RANK_STEP))

    out = {}
    for leaf in _QUANT_SUFFIXES:
        out[leaf] = int(round(rank * mults[leaf] / _RANK_STEP)) * _RANK_STEP
    return out


def rank_for_leaf(ranks: dict[str, int], leaf: str | None) -> int:
    """The branch rank for a target layer, given its leaf name. 0 means "no branch".



    Matches by suffix because `is_target` does. An exact-key lookup would disagree with it

    on any variant that nested the leaves one level deeper (`blocks.0.wrapper.attn.wq`):

    such a layer is a quantization target, so it would be quantized, but would find no rank

    entry and be written without a low-rank branch -- an svdq checkpoint silently missing

    branches on part of the network, which is the class of quiet wrongness the zero-layer

    check below exists to prevent.

    """
    if not ranks or leaf is None:
        return 0
    if leaf in ranks:
        return ranks[leaf]
    for suffix, value in ranks.items():
        if leaf.endswith(suffix):
            return value
    return 0


def svd_lowrank(weight: torch.Tensor, rank: int, oversample: int = 16, niter: int = 2):
    """Randomized truncated SVD: return (L1 [out, rank], L2 [rank, in]) with W ~ L1 @ L2.



    Randomized rather than exact because these are 6144x16384 matrices and only the top

    few dozen singular directions matter here; `oversample` extra probe columns plus a

    couple of power iterations recover them to well past the precision 4-bit quantization

    cares about.

    """
    w = weight.float()
    min_dim = min(w.shape)
    rank = max(1, min(int(rank), min_dim))
    q = min(rank + max(0, int(oversample)), min_dim)

    if min_dim <= max(q, 32):
        u, s, vh = torch.linalg.svd(w, full_matrices=False)
        u_r, s_r, vh_r = u[:, :rank], s[:rank], vh[:rank, :]
    else:
        u, s, v = torch.svd_lowrank(w, q=q, niter=max(0, int(niter)))
        u_r, s_r, vh_r = u[:, :rank], s[:rank], v[:, :rank].transpose(-2, -1)

    return (u_r * s_r.unsqueeze(0)).contiguous(), vh_r.contiguous()


def _act_weighting(act_rms: torch.Tensor | None, in_features: int, device, floor: float = 0.05):
    """Turn a captured activation RMS vector into the column weighting `d`, or None.



    Normalised to mean 1 so the weighting only redistributes emphasis and never rescales the

    problem -- otherwise the error metric would no longer be comparable across layers and the

    refinement loop's stopping threshold would mean something different in each one.



    The floor matters more than it looks. `l2` is recovered by dividing by `d`, so a channel

    whose activation is ~0 would divide a finite number by ~0 and put an enormous row into the

    branch -- a branch that then dominates the output for a channel the model never actually

    drives. Clamping to 5% of the mean bounds that amplification at 20x.

    """
    if act_rms is None:
        return None
    d = act_rms.to(device=device, dtype=torch.float32).flatten()
    if d.numel() != in_features:
        raise RuntimeError("activation stats have {} channels but the weight has {}"
                           .format(d.numel(), in_features))
    mean = d.mean()
    if not torch.isfinite(mean) or mean <= 0:
        return None
    d = d / mean
    return d.clamp(min=floor)


def svdquant_split(weight: torch.Tensor, rank: int, fmt: str, groupsize: int,

                   refine_iters: int = 100, act_rms: torch.Tensor | None = None,

                   refine_tol: float = REFINE_TOL):
    """SVDQuant ordering: pull a low-rank bf16 branch out of W, quantize the residual.



    The low-rank branch absorbs the outlier-heavy directions, so the part that has to

    survive 4 bits is better conditioned.



    A single SVD of W is only the first guess: it picks the directions that are largest

    in W, which is not the same as the directions the quantizer handles worst. So this

    iterates the way deepcompressor does -- re-fit the branch to whatever the quantizer

    is *currently* getting wrong, requantize, repeat -- and keeps the best result. Pass

    ``refine_iters=0`` for the plain single-shot split.



    Iteration one is exactly that single-shot split and the best is kept, so refining

    can only match or beat it. Measured on Krea2 Turbo, rank 64: ~10% less

    reconstruction error, every layer improving.



    `refine_tol` is what stops it, and it is *relative* for a reason. It used to be an

    absolute `1e-6` against an error of order 1e-1, so the loop only stopped when an

    iteration failed to improve by one part in a hundred thousand -- which measured at a mean

    of **79.8 of the 100 iterations** across a sample of twelve layers at rank 256. What those

    iterations buy, traced per iteration on the same layers:



        stop at              iterations   reconstruction error

        1e-6 absolute (old)     ~80        baseline

        0.5% relative            ~9        +4.3%

        0.1% relative (now)     ~22        +2.0%



    Conversion goes from ~14 minutes to ~4. The 2% is reconstruction error, which this repo

    has three times measured to be a poor predictor of image outcome (see `--rank-alloc` in

    the README), so read it as a conversion-time change rather than a quality one.



    The objective is weight reconstruction error, which needs no calibration data.

    That is the true output error under the assumption that the input covariance is

    identity -- and spreading outliers with the convrot rotation is what makes that

    assumption a reasonable one. Closing the remaining gap to deepcompressor means

    measuring the real covariance, which is what makes their conversion take hours.



    Returns None for a degenerate weight (all-zero, or one that makes the error metric

    non-finite); the caller quantizes those without a branch rather than aborting.

    """
    w = weight.float()
    d = _act_weighting(act_rms, w.shape[1], w.device)
    # With a weighting, both the SVD and the error metric work in the scaled space: minimising
    # ||(W - W_hat) diag(d)|| approximates the error the layer's real inputs would produce,
    # which is what actually reaches the image. Without one this is the identity and the whole
    # function reduces to exactly the old behaviour.
    w_norm = torch.linalg.matrix_norm(w if d is None else w * d).item()
    if not math.isfinite(w_norm) or w_norm == 0.0:
        return None
    qw = torch.zeros((), device=w.device, dtype=torch.float32)

    best = None
    best_err = float("inf")
    for _ in range(max(1, refine_iters)):
        target = w - qw
        if d is None:
            l1, l2 = svd_lowrank(target, rank, oversample=16, niter=2)
        else:
            # Fit in the scaled space so the retained directions are the ones that matter
            # once real activations hit them, then undo the scale on `l2` so `l1 @ l2` is
            # still an approximation of the *unscaled* residual.
            l1, l2s = svd_lowrank(target * d, rank, oversample=16, niter=2)
            l2 = l2s / d
            del l2s
        l1 = l1.to(torch.bfloat16)
        l2 = l2.to(torch.bfloat16)
        lw = l1.float() @ l2.float()
        # The quantizer must see the unscaled residual: the scale is an analysis device, and
        # baking it into what gets written would change what the kernel reconstructs.
        residual = (w - lw).to(torch.bfloat16)

        qdata, params, layout, _ = _quantize_raw(residual, fmt, groupsize)
        qw = layout.dequantize(qdata, params).float()
        del qdata, params

        left = w - (lw + qw)
        err = (torch.linalg.matrix_norm(left if d is None else left * d) / w_norm).item()
        del lw, left
        # NaN fails every comparison, so without this the loop would neither break nor
        # ever update `best` -- it would burn all `refine_iters` and return the last
        # split rather than the best one.
        if not math.isfinite(err):
            break
        if best is not None and err >= best_err * (1.0 - refine_tol):
            break  # refinement has stopped paying off
        best_err, best = err, (residual, l1, l2)

    return best


def _quantize_raw(weight: torch.Tensor, fmt: str, groupsize: int):
    """Quantize to `fmt`, handing back the layout and params too.



    Split out from `quantize_weight` so the refinement loop can round-trip a weight

    (quantize then dequantize) to see what the quantizer is actually getting wrong.

    """
    layout = get_layout_class(QUANT_ALGOS[fmt]["comfy_tensor_layout"])
    if fmt == "int8_tensorwise":
        qdata, params = layout.quantize(
            weight, is_weight=True, per_channel=True,
            convrot=True, convrot_groupsize=groupsize,
        )
        conf = {"format": "int8_tensorwise", "convrot": True, "convrot_groupsize": groupsize}
    elif fmt == "convrot_w4a4":
        qdata, params = layout.quantize(weight, convrot_groupsize=groupsize)
        conf = {"format": "convrot_w4a4", "convrot_groupsize": groupsize,
                "linear_dtype": getattr(params, "linear_dtype", "int4")}
    elif fmt == "float8_e4m3fn":
        qdata, params = layout.quantize(weight, scale="recalculate")
        conf = {"format": "float8_e4m3fn"}
    else:
        raise ValueError(fmt)
    return qdata, params, layout, conf


def quantize_weight(weight: torch.Tensor, fmt: str, groupsize: int):
    qdata, params, _, conf = _quantize_raw(weight, fmt, groupsize)
    return qdata, {"weight_scale": params.scale}, conf


def conf_tensor(conf: dict) -> torch.Tensor:
    return torch.tensor(list(json.dumps(conf).encode("utf-8")), dtype=torch.uint8)


def load_act_stats(path: str) -> dict[str, torch.Tensor]:
    """Read the file the capture nodes write: {layer name: per-input-channel RMS}."""
    stats = {}
    with safe_open(path, framework="pt", device="cpu") as handle:
        for key in handle.keys():
            stats[key] = handle.get_tensor(key)
    if not stats:
        raise RuntimeError("{} contains no activation statistics".format(path))
    return stats


def check_act_stats_coverage(stats: dict, keys, prefix: str, ranks: dict) -> None:
    """Fail before any GPU work if the statistics don't cover every layer that will branch.



    Partial statistics are the quiet-wrongness case: the covered layers would get an

    activation-aware split and the rest a plain one, producing a checkpoint that is neither,

    and no measurement of it would mean anything.



    Deliberately key-only. Which layers branch depends on `is_target` plus a positive

    `rank_for_leaf`, both of which read the layer *name*; the one weight-dependent step in

    the loop (`min(wanted, min(w.shape))`) can only lower a rank, never raise a zero above

    it, so this pre-pass sees exactly the set the loop will branch. Doing it here instead of

    after the loop turns a ~6 minute mistake into a two-second one.

    """
    covered, missing = 0, []
    for key in keys:
        if not key.endswith(".weight"):
            continue
        layer = key[: -len(".weight")]
        if not is_target(layer, prefix):
            continue
        if rank_for_leaf(ranks, leaf_name(layer, prefix)) <= 0:
            continue
        if stats.get(layer[len(prefix):] if prefix else layer) is None:
            missing.append(layer)
        else:
            covered += 1
    if missing:
        raise RuntimeError(
            "activation stats cover {} of {} branched layers; {} are missing, e.g. {}.\n"
            "The capture run did not see every layer -- usually because sampling was "
            "interrupted, or the statistics came from a different model. Re-capture rather "
            "than building a checkpoint that is activation-aware in some layers and not "
            "others.".format(covered, covered + len(missing), len(missing),
                             ", ".join(missing[:3])))


def convert(src: str, dst: str, fmt: str, groupsize: int, device: str = "cuda", rank: int = 0,

            refine_iters: int = 100, variant: str = "unknown", progress_cb=None,

            rank_alloc: str = "uniform", act_stats: str | None = None,

            weight_patch=None, refine_tol: float = REFINE_TOL):
    """Quantize `src` into `dst`. Returns the summary line it printed.



    `rank` is a budget, `rank_alloc` decides how it is spread across the eight projection

    types -- see `RANK_ALLOCATIONS`. Non-uniform allocations keep the same total branch bytes.



    `weight_patch(key, tensor) -> tensor` runs on every tensor as it is read, before anything

    is quantized. `tools/bake_adapter.py` uses it to add a LoRA delta into the high-precision

    weight, so the low-rank branch is fitted against the *merged* weight rather than the LoRA

    being requantized on top of a finished checkpoint. Streaming it here rather than writing a

    merged source out first matters in practice: the bf16 source is 24 GB.



    `progress_cb(done, total, message)` is called as layers complete, for callers with a

    progress bar to drive (the ComfyUI node). Errors are `RuntimeError`, never `SystemExit`:

    `SystemExit` derives from `BaseException`, so raising it from inside a node would slip

    past ComfyUI's executor instead of being reported as a failed node. `main()` translates

    them back for the CLI.

    """
    out: dict[str, torch.Tensor] = {}
    quantized = 0
    kept = 0
    branched = 0
    observed_leaves: set[str] = set()
    ranks = leaf_ranks(rank, rank_alloc) if rank > 0 else {}
    used_ranks: set[int] = set()
    clamped: dict[str, int] = {}
    stats = load_act_stats(act_stats) if act_stats else {}
    t0 = time.time()

    with safe_open(src, framework="pt", device="cpu") as handle:
        keys = list(handle.keys())
        prefix = detect_prefix(keys)
        check_requantizable(handle, keys, prefix)
        if stats:
            check_act_stats_coverage(stats, keys, prefix, ranks)

        # Companion keys (old scales/markers) for target layers get regenerated fresh
        # when we process the ".weight" key below -- skip the stale copies on disk,
        # otherwise they'd overwrite our new ones later in iteration order.
        stale_companions = set()
        for key in keys:
            if key.endswith(".weight"):
                layer = key[: -len(".weight")]
                if is_target(layer, prefix):
                    for suffix in ("weight_scale", "weight_scale_2", "input_scale", "comfy_quant"):
                        stale_companions.add("{}.{}".format(layer, suffix))

        for i, key in enumerate(keys):
            if key.endswith(".weight"):
                layer = key[: -len(".weight")]
                leaf = leaf_name(layer, prefix)
                if leaf is not None:
                    observed_leaves.add(leaf)
                if is_target(layer, prefix):
                    w = dequantize_target_weight(handle, layer, device)
                    if weight_patch is not None:
                        w = weight_patch(key, w)
                    # A reallocated budget can ask for more rank than the weight has singular
                    # directions (attn.wk is only 1536 wide). `svd_lowrank` clamps internally,
                    # but clamping here too keeps the reported rank honest.
                    wanted = rank_for_leaf(ranks, leaf)
                    leaf_rank = min(wanted, min(w.shape))
                    if leaf and leaf_rank < wanted:
                        clamped[leaf] = leaf_rank
                    if leaf_rank > 0:
                        # Stats are keyed by module path (`blocks.0.attn.wq`); checkpoint
                        # keys may carry a prefix, so strip it to look them up.
                        # Coverage was already proven complete before the loop started, so
                        # a None here cannot happen for a branched layer.
                        act = stats.get(layer[len(prefix):] if prefix else layer)
                        split = svdquant_split(w, leaf_rank, fmt, groupsize, refine_iters,
                                               act_rms=act, refine_tol=refine_tol)
                        if split is None:
                            print("  warning: {} is degenerate (zero or non-finite); "
                                  "quantizing it without a low-rank branch".format(layer),
                                  flush=True)
                        else:
                            w, l1, l2 = split
                            out["{}.svdq_l1".format(layer)] = l1.cpu()
                            out["{}.svdq_l2".format(layer)] = l2.cpu()
                            branched += 1
                            used_ranks.add(int(l1.shape[1]))
                            del l1, l2
                    qdata, scales, conf = quantize_weight(w, fmt, groupsize)
                    out[key] = qdata.cpu()
                    for name, value in scales.items():
                        out["{}.{}".format(layer, name)] = value.cpu()
                    out["{}.comfy_quant".format(layer)] = conf_tensor(conf)
                    del w, qdata, scales
                    quantized += 1
                    if progress_cb is not None:
                        progress_cb(quantized, _EXPECTED_LAYERS,
                                    "quantized {} layers ({:.0f}s)".format(
                                        quantized, time.time() - t0))
                    if quantized % 32 == 0:
                        torch.cuda.empty_cache()
                        print(f"  [{i + 1}/{len(keys)}] quantized {quantized} layers "
                              f"({time.time() - t0:.0f}s)", flush=True)
                    continue
            if key in stale_companions:
                continue
            value = handle.get_tensor(key)
            # The non-quantized layers (txtfusion, embeddings) can carry a LoRA too, and they
            # are the half a bake would otherwise silently drop.
            out[key] = value if weight_patch is None else weight_patch(key, value)
            kept += 1

    gc.collect()
    torch.cuda.empty_cache()

    if quantized == 0:
        raise RuntimeError(
            "Found the transformer blocks but quantized nothing: no layer under "
            "'{}blocks.' ends in one of the expected leaf names.\n"
            "  expected: {}\n"
            "  observed: {}\n"
            "This is a Krea 2 variant with different layer naming; quantize_krea2.py "
            "would otherwise have written a full-size file that does nothing."
            .format(prefix, ", ".join(sorted(_QUANT_SUFFIXES)),
                    ", ".join(sorted(observed_leaves)) or "none")
        )

    # A leaf can only be clamped when the allocation asked for more rank than the weight is
    # wide, which means part of the budget went unspent and the file is *smaller* than the
    # uniform one -- so a comparison against uniform at this budget is no longer like-for-like.
    if clamped:
        print("warning: rank allocation {!r} asked for more rank than these layers are wide, "
              "so they were capped and part of the budget went unspent: {}. The checkpoint is "
              "smaller than uniform rank {}, not the same size -- lower the budget to compare "
              "them fairly.".format(
                  rank_alloc, ", ".join("{}={}".format(k, v) for k, v in sorted(clamped.items())),
                  rank), flush=True)

    created = len(out) - kept - quantized
    factors = "" if not rank else " + {} low-rank factors".format(branched * 2)
    print(f"quantized {quantized} layers; {kept} tensors passed through; "
          f"{created} tensors created (qdata/scale/marker{factors}); "
          f"{len(out)} tensors total; writing {dst} ...", flush=True)

    # Recorded so the loader (and anyone inspecting the file) can tell what produced it.
    # safetensors metadata is str -> str only. The loader treats all of this as optional:
    # checkpoints published before this existed still load, with the rank recovered from
    # the svdq_l1 shape as before.
    # The loader trusts the factor shapes over `krea2_svdquant_rank`, and under a non-uniform
    # allocation no single number is the rank -- so record the budget, the allocation name and
    # the full per-leaf map, and leave `rank` meaning "budget" for older loaders reading it.
    metadata = {
        "krea2_svdquant_tool_version": __version__,
        "krea2_svdquant_format": fmt,
        "krea2_svdquant_rank": str(rank),
        "krea2_svdquant_rank_alloc": rank_alloc,
        "krea2_svdquant_rank_map": json.dumps(ranks, sort_keys=True) if ranks else "",
        "krea2_svdquant_groupsize": str(groupsize),
        "krea2_svdquant_refine_iters": str(refine_iters if rank else 0),
        "krea2_svdquant_variant": variant,
        "krea2_svdquant_quantized_layers": str(quantized),
        "krea2_svdquant_branched_layers": str(branched),
        "krea2_svdquant_source_name": os.path.basename(src),
        "krea2_svdquant_source_bytes": str(os.path.getsize(src)),
        # Two checkpoints with identical rank/refine settings are still different models if
        # one was activation-weighted, so this has to be recoverable from the file.
        "krea2_svdquant_act_stats": os.path.basename(act_stats) if act_stats else "",
    }
    if progress_cb is not None:
        progress_cb(quantized, _EXPECTED_LAYERS, "writing {:.2f} GB ...".format(
            sum(t.numel() * t.element_size() for t in out.values()) / 1024 ** 3))
    save_file(out, dst, metadata=metadata)
    size = os.path.getsize(dst) / 1024 ** 3
    rank_desc = ""
    if used_ranks:
        rank_desc = (" rank {}".format(next(iter(used_ranks))) if len(used_ranks) == 1
                     else " rank {}-{} ({})".format(min(used_ranks), max(used_ranks),
                                                    rank_alloc))
    summary = ("quantized {} layers ({} branched{}) in {:.0f}s -> {}  ({:.2f} GB, source "
               "{:.2f} GB)".format(quantized, branched, rank_desc, time.time() - t0, dst, size,
                                   os.path.getsize(src) / 1024 ** 3))
    print("done in {:.0f}s -> {}  ({:.2f} GB, source {:.2f} GB)".format(
        time.time() - t0, dst, size, os.path.getsize(src) / 1024 ** 3))
    return summary


# The Krea 2 block count, used only to scale a progress bar. A variant with a different
# depth still quantizes correctly; the bar is just less accurate.
_EXPECTED_LAYERS = 28 * len(_QUANT_SUFFIXES)

_FORMAT_ALIASES = {
    "int8": "int8_tensorwise",
    "w4a4": "convrot_w4a4",
    "svdq": "convrot_w4a4",
    "fp8": "float8_e4m3fn",
}

_GENERIC_STEMS = ("raw", "model", "diffusion_pytorch_model", "turbo")

SAMPLER_HINTS = {
    "base": ("Base (non-distilled) model: start from ~50 steps, cfg 3.5, euler/simple, and a "
             "real negative prompt. See workflows/krea2_base_svdquant_w4a4_t2i.json."),
    "turbo": ("Turbo (distilled) model: 8 steps, cfg 1.0, euler/simple. "
              "See workflows/krea2_turbo_svdquant_w4a4_t2i.json."),
}


def resolve_format(fmt_name: str, rank: int, rank_was_set: bool = True) -> tuple[str, int]:
    """Map a CLI/node format name to a comfy_kitchen format, validating the rank with it.



    Shared by `main()` and the ComfyUI node so the two cannot drift apart on which

    combinations are legal.

    """
    fmt = _FORMAT_ALIASES[fmt_name]
    if fmt not in QUANT_ALGOS:
        raise RuntimeError("{} is not available in this ComfyUI build".format(fmt))
    # Silently zeroing the rank here used to make `--format w4a4 --rank 128` look like it had
    # done something it had not.
    if fmt_name != "svdq" and rank_was_set:
        raise RuntimeError(
            "rank only applies to format 'svdq' (you asked for '{}'). The low-rank branch is "
            "what distinguishes svdq from plain w4a4.".format(fmt_name))
    if fmt_name == "svdq" and rank <= 0:
        raise RuntimeError("format 'svdq' needs rank > 0; use format 'w4a4' for no branch")
    return fmt, (rank if fmt_name == "svdq" else 0)


def derive_out_path(src: str, fmt_name: str, rank: int, variant: str,

                    rank_alloc: str = "uniform", act_stats: str | None = None

                    ) -> tuple[str, str | None]:
    """The default output path, plus a note to show when the source name is uninformative."""
    stem = os.path.splitext(os.path.basename(src))[0]
    note = None
    if variant != "unknown":
        stem = "Krea2-{}".format(variant.capitalize())
    elif stem.lower() in _GENERIC_STEMS:
        note = ("note: '{}' is a generic filename. Set the variant (or an explicit output "
                "name) to get a checkpoint name you will still recognise next month."
                .format(stem))
    # The allocation goes in the filename because two checkpoints with the same budget and
    # different allocations are the same size -- nothing else would tell them apart.
    alloc_tag = "" if rank_alloc == "uniform" else "-{}".format(rank_alloc)
    # Same reasoning as alloc_tag: an activation-weighted build is byte-identical in shape to
    # a plain one, so without a tag the two are indistinguishable on disk.
    act_tag = "-actaware" if act_stats else ""
    suffix = ("SVDQuant-W4A4-rank{}{}{}".format(rank, alloc_tag, act_tag) if rank else
              ("{}-convrot".format(fmt_name.upper()) if fmt_name != "fp8" else "FP8"))
    return os.path.join(os.path.dirname(src), "{}-{}.safetensors".format(stem, suffix)), note


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("src")
    ap.add_argument("--format", choices=["int8", "w4a4", "svdq", "fp8"], default="int8",
                    help="svdq = w4a4 residual + SVDQuant low-rank bf16 branch; "
                         "fp8 = float8_e4m3fn, no convrot, no low-rank branch")
    ap.add_argument("--groupsize", type=int, default=256, help="unused for fp8")
    ap.add_argument("--rank", type=int, default=64, help="low-rank branch budget, svdq only")
    ap.add_argument("--rank-alloc", choices=sorted(RANK_ALLOCATIONS), default="uniform",
                    help="how to spread the rank budget across the eight projection types. "
                         "uniform = same rank everywhere. gqa = byte-neutral reallocation "
                         "towards the GQA kv projections, which absorb ~2x the error at a "
                         "third of the branch cost (see RANK_ALLOCATIONS)")
    ap.add_argument("--refine-iters", type=int, default=100,
                    help="svdq only: refine the low-rank branch against the quantization "
                         "error, keeping the best (0 = plain single-shot SVD, much faster "
                         "but ~10%% more reconstruction error). This is a cap; --refine-tol "
                         "is what usually stops the loop first")
    ap.add_argument("--refine-tol", type=float, default=REFINE_TOL, metavar="FRACTION",
                    help="svdq only: stop refining a layer once an iteration improves its "
                         "reconstruction error by less than this fraction (default %(default)s "
                         "= 0.1%%). Lower means more iterations for less return: 0.001 takes "
                         "~22 iterations per layer, 0.005 takes ~9 for 2%% more error, and 0 "
                         "restores the old behaviour of running nearly all --refine-iters")
    ap.add_argument("--variant", choices=["turbo", "base", "unknown"], default="unknown",
                    help="which Krea 2 release this is. Only affects the output filename "
                         "and the recorded metadata -- quantization is identical for both, "
                         "the difference is the sampler settings you run afterwards")
    ap.add_argument("--act-stats", default=None, metavar="PATH",
                    help="svdq only: activation statistics from the Krea2 SVDQuant Capture "
                         "nodes. Weights the low-rank split by per-input-channel activation "
                         "RMS, so the branch spends its capacity on the directions the model "
                         "actually drives instead of on the largest weights. Costs nothing "
                         "at runtime -- same rank, format, size and kernel")
    ap.add_argument("--out", default=None)
    ap.add_argument("--device", default="cuda")
    args = ap.parse_args()

    if args.act_stats:
        if args.format != "svdq":
            raise SystemExit("--act-stats only applies to format 'svdq': it weights the "
                             "low-rank split, and the other formats have no branch")
        if not os.path.exists(args.act_stats):
            raise SystemExit("--act-stats file not found: {}".format(args.act_stats))

    # RuntimeError is the shared failure type (see `convert`); the CLI wants SystemExit so it
    # prints one clean line instead of a traceback.
    try:
        fmt, rank = resolve_format(args.format, args.rank,
                                   rank_was_set=args.rank != ap.get_default("rank"))
    except RuntimeError as exc:
        raise SystemExit(str(exc)) from None

    if args.rank_alloc != "uniform" and not rank:
        raise SystemExit("--rank-alloc only applies to format 'svdq'")

    out = args.out
    if out is None:
        out, note = derive_out_path(args.src, args.format, rank, args.variant, args.rank_alloc,
                                    args.act_stats)
        if note:
            print(note, flush=True)
    try:
        convert(args.src, out, fmt, args.groupsize, args.device, rank, args.refine_iters,
                variant=args.variant, rank_alloc=args.rank_alloc, act_stats=args.act_stats,
                refine_tol=args.refine_tol)
    except RuntimeError as exc:
        raise SystemExit(str(exc)) from None

    hint = SAMPLER_HINTS.get(args.variant)
    if hint:
        print("\n" + hint)


if __name__ == "__main__":
    main()