Phase 2
This commit is contained in:
@@ -17,6 +17,7 @@ from toolkit.accelerator import get_accelerator, unwrap_model
|
||||
from toolkit.util.quantize import quantize_model
|
||||
import torch.nn.functional as F
|
||||
from toolkit.memory_management import MemoryManager
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from diffusers import (
|
||||
@@ -101,9 +102,17 @@ class QwenImageModel(QwenImageVAEHolderMixin, BaseModel):
|
||||
if os.path.exists(te_folder_path):
|
||||
base_model_path = model_path
|
||||
|
||||
transformer = QwenImageTransformer2DModel.load_model(model_path, dtype=model_dtype)
|
||||
transformer = QwenImageTransformer2DModel.load_model(
|
||||
model_path,
|
||||
dtype=model_dtype,
|
||||
use_comfy_weights=self.model_config.model_kwargs.get(
|
||||
"use_comfy_weights", True
|
||||
),
|
||||
)
|
||||
|
||||
if self.model_config.quantize:
|
||||
if self.model_config.quantize and not getattr(
|
||||
transformer, "aitk_is_quantized", False
|
||||
):
|
||||
self.print_and_status_update("Quantizing Transformer")
|
||||
quantize_model(self, transformer)
|
||||
flush()
|
||||
@@ -350,17 +359,17 @@ class QwenImageModel(QwenImageVAEHolderMixin, BaseModel):
|
||||
return False
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
# comfy-format single-file save (diffusers keys ARE the comfy layout
|
||||
# for qwen image); prequantized layers keep their quantized storage
|
||||
transformer: QwenImageTransformer2DModel = unwrap_model(self.model)
|
||||
transformer.save_pretrained(
|
||||
save_directory=os.path.join(output_path, "transformer"),
|
||||
safe_serialization=True,
|
||||
if not output_path.endswith(".safetensors"):
|
||||
output_path += ".safetensors"
|
||||
transformer.save_model(
|
||||
output_path,
|
||||
dtype=save_dtype,
|
||||
metadata=get_meta_for_safetensors(meta, name=self.arch),
|
||||
)
|
||||
|
||||
meta_path = os.path.join(output_path, "aitk_meta.yaml")
|
||||
with open(meta_path, "w") as f:
|
||||
yaml.dump(meta, f)
|
||||
|
||||
def get_loss_target(self, *args, **kwargs):
|
||||
noise = kwargs.get("noise")
|
||||
batch = kwargs.get("batch")
|
||||
|
||||
@@ -3,7 +3,6 @@ from typing import List, Optional
|
||||
|
||||
import huggingface_hub
|
||||
import torch
|
||||
import yaml
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig, NetworkConfig
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
from toolkit.models.base_model import BaseModel
|
||||
@@ -13,15 +12,12 @@ from toolkit.samplers.custom_flowmatch_sampler import (
|
||||
CustomFlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from toolkit.util.quantize import (
|
||||
quantize_model,
|
||||
dequantize_if_quantized,
|
||||
)
|
||||
from toolkit.util.quantize import quantize_model
|
||||
from toolkit.memory_management import MemoryManager
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.models.v2.text_encoders.qwen3 import Qwen3TextEncoder
|
||||
from toolkit.models.v2.vae.autoencoder_kl import KLVAE
|
||||
from safetensors.torch import load_file, save_file
|
||||
from safetensors.torch import load_file
|
||||
|
||||
|
||||
try:
|
||||
@@ -200,6 +196,9 @@ class ZImageModel(BaseModel):
|
||||
qtype=qtype,
|
||||
quantize_device=self.device_torch,
|
||||
config_path=base_model_path if self.is_single_file else None,
|
||||
use_comfy_weights=self.model_config.model_kwargs.get(
|
||||
"use_comfy_weights", True
|
||||
),
|
||||
)
|
||||
flush()
|
||||
|
||||
@@ -400,31 +399,17 @@ class ZImageModel(BaseModel):
|
||||
return ZImageTransformer2DModel.get_quantization_exclude_modules()
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# comfy-format single-file save (the standard save format regardless
|
||||
# of how the model was loaded); prequantized layers keep their
|
||||
# quantized storage
|
||||
transformer: ZImageTransformer2DModel = unwrap_model(self.model)
|
||||
if self.is_single_file:
|
||||
# loaded from a single-file checkpoint, save back in that format
|
||||
sd = transformer.state_dict()
|
||||
save_dict = {}
|
||||
for key, value in sd.items():
|
||||
# dequantize any quantized (e.g. torchao) weights so we save plain tensors
|
||||
save_dict[key] = (
|
||||
dequantize_if_quantized(value).clone().to("cpu", dtype=save_dtype)
|
||||
)
|
||||
save_dict = transformer.convert_state_dict_on_save(save_dict)
|
||||
|
||||
if not output_path.endswith(".safetensors"):
|
||||
output_path += ".safetensors"
|
||||
meta = get_meta_for_safetensors(meta, name=self.arch)
|
||||
save_file(save_dict, output_path, metadata=meta)
|
||||
else:
|
||||
transformer.save_pretrained(
|
||||
save_directory=os.path.join(output_path, "transformer"),
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
meta_path = os.path.join(output_path, "aitk_meta.yaml")
|
||||
with open(meta_path, "w") as f:
|
||||
yaml.dump(meta, f)
|
||||
if not output_path.endswith(".safetensors"):
|
||||
output_path += ".safetensors"
|
||||
transformer.save_model(
|
||||
output_path,
|
||||
dtype=save_dtype,
|
||||
metadata=get_meta_for_safetensors(meta, name=self.arch),
|
||||
)
|
||||
|
||||
def get_loss_target(self, *args, **kwargs):
|
||||
noise = kwargs.get("noise")
|
||||
|
||||
@@ -252,10 +252,89 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord
|
||||
(`dequantize_if_quantized` everywhere), raw-`quantize()` → `quantize_model()`
|
||||
|
||||
### Phase 2 — comfy weights become the preferred source
|
||||
- [ ] Wire `comfy_weight_names` per model; standard-repo `name_or_path` + existing
|
||||
comfy weights → load comfy
|
||||
- [ ] Comfy-format save (bf16 + convrot8/nvfp4 quantized) as the default
|
||||
full-weight save
|
||||
|
||||
Decisions:
|
||||
- Comfy weights come from the Comfy-Org hub repos (per-model repos, comfy
|
||||
layout nested under `split_files/` — stripped when placing files into
|
||||
MODELS_PATH). Repos ship several precision variants of each component.
|
||||
- **Selection preference: convrot8 > float8 mixed > float8 > bf16 > fp16**
|
||||
(`resolver.comfy_precision_rank`; nvfp4/unmarked rank last and are only
|
||||
used when explicitly listed). **Local-first**: the best-ranked LOCAL
|
||||
candidate wins; only when no candidate is local is the best-ranked one
|
||||
downloaded.
|
||||
- Per-model candidate lists (`aitk_comfy_weight_names`) hold only the
|
||||
variants the class can actually digest. Constraint discovered: comfy
|
||||
convrot/quantized files carry markers on the ORIGINAL module layout (e.g.
|
||||
z_image's fused attention.qkv) — diffusers-layout classes with split
|
||||
modules can't attach them until they grow fused-layout support; until then
|
||||
those models list bf16/fp8 variants only. (Vendored comfy-layout classes —
|
||||
minimax/ltx pattern — take convrot directly.)
|
||||
- The standard repo still supplies the config; local dirs and unregistered
|
||||
repos load as-is; `model_kwargs.use_comfy_weights: false` opts out.
|
||||
|
||||
- [x] Mechanism: `resolver.comfy_precision_rank` / `comfy_local_rel` /
|
||||
`resolve_comfy_candidates` + `OstrisModelMixin.resolve_comfy_weights`,
|
||||
integrated into `load_model` (comfy file preferred for registered
|
||||
standard repos, downloaded into the shared comfy layout)
|
||||
- [x] First wired model: z_image (Comfy-Org/z_image_turbo, bf16 candidate).
|
||||
Verified end-to-end against the real shared ComfyUI folder
|
||||
(MODELS_PATH=/mnt/Models/comfy_models): standard-repo name_or_path
|
||||
resolved to the locally-present comfy file and produced the identical
|
||||
generation to the diffusers-shards load
|
||||
- [x] qwen_image wired: its comfy files use the diffusers key layout
|
||||
directly — `fp8mixed` (float8-mixed, rank 1) attaches its 839
|
||||
`float8_e4m3fn` markers straight onto the class; candidates fp8mixed →
|
||||
fp8_e4m3fn (raw cast) → bf16. Verified with real weights: loaded the
|
||||
shared folder's local fp8 file (local-first, no download) and generated
|
||||
correctly. Holder skips re-quantization for prequantized checkpoints.
|
||||
- [x] New `float8_e4m3fn` Ostris backend (toolkit/util/float8_quant.py):
|
||||
ComfyUI's fp8 + per-tensor-scale storage with dequantized matmul, in
|
||||
get_ostris_quantizer + comfy import/export. Round-trip verified.
|
||||
- [x] wan family wired: comfy wan files (original key layout) convert via
|
||||
diffusers' own `convert_wan_transformer_to_diffusers` (rename-only, so
|
||||
quantized weight/scale keys ride along with their modules). Candidate
|
||||
keys support `(repo, subfolder)` tuples for wan2.2 A14B's dual DiTs
|
||||
(transformer = high noise, transformer_2 = low noise) and per-entry
|
||||
comfy-repo overrides ({"repo": ..., "files": [...]}) since wan2.1 and
|
||||
2.2 files live in different Comfy-Org repos. Wired: 2.2 TI2V-5B,
|
||||
T2V/I2V-A14B (fp8_scaled), 2.1 T2V 1.3B/14B, I2V 480P/720P. Verified
|
||||
with real weights: wan21 1.3B (downloaded comfy bf16) and wan22 5B
|
||||
(local comfy fp16) both load through the converter and generate video.
|
||||
Fix along the way: the mixin's meta build now uses accelerate
|
||||
init_empty_weights (params meta, buffers real) so init-computed
|
||||
non-persistent buffers like wan's rope tables materialize.
|
||||
- [x] Legacy ComfyUI scaled-fp8 support (`scaled_fp8` marker + per-layer
|
||||
fp8 weight / scalar scale_weight, e.g. every wan *_fp8_scaled file):
|
||||
imports onto the float8 backend; scale_input (activation quant) is
|
||||
dropped, matmuls run dequantized.
|
||||
- [ ] Wire remaining archs' candidate lists (chroma/others as their key
|
||||
conversions are verified per file); wan comfy-format save needs the
|
||||
inverse key mapping (defer with the other save flips)
|
||||
- [x] Fused-layout quantized attach for diffusers-split classes:
|
||||
`split_fused_quantized_keys` / `fuse_split_quantized_keys`
|
||||
(comfy_quant_import) do exact out-dim row surgery on quantized comfy
|
||||
entries for all three formats (int8 rows+scales slice; fp8 scalar and
|
||||
nvfp4 per-tensor scales shared; nvfp4 block scales
|
||||
unswizzle→split→reswizzle). z_image's load/save converters use them, so
|
||||
its convrot8 candidate is live and top-ranked. Unit-verified exact both
|
||||
directions.
|
||||
- [x] Comfy-format save: `save_model` auto-keeps quantized storage
|
||||
(comfy_quant markers) for convrot8 / nvfp4 / convrotcomfyw4a4 layers via
|
||||
`toolkit/util/comfy_quant_export.py` (inverse of comfy_quant_import;
|
||||
nvfp4 nibbles re-swapped + scales re-swizzled to the cuBLAS tile
|
||||
layout), plain layers save at bf16; partially-exportable models fall
|
||||
back to dequantized. Round-trip verified: save → mixin reload → outputs
|
||||
match for convrot8, nvfp4, and plain layers.
|
||||
- [x] Save unification started: z_image and qwen_image holders now save
|
||||
comfy-format single files via the mixin regardless of how they loaded
|
||||
(z_image's dual-style branch deleted). Real round trip verified: the
|
||||
published z_image int8_convrot file loads (270 quantized linears,
|
||||
split-attach), resaves to the IDENTICAL 857-key comfy layout with
|
||||
bit-exact fused qkv weights/scales/markers, and the reload's quantized
|
||||
forward is bit-identical — toolkit saves are byte-compatible with
|
||||
ComfyUI.
|
||||
- [ ] Flip the remaining per-arch `save_model` overrides as each arch's
|
||||
save-side key conversion is in place
|
||||
- [ ] Publish/verify comfy repacks per model as they flip
|
||||
|
||||
### Phase 3 — live server
|
||||
@@ -291,9 +370,10 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord
|
||||
device); ideogram4's fp8 release renders its own "blocked by safety
|
||||
filter" card for a plain cat prompt (model behavior, not a bug —
|
||||
investigate its trigger).
|
||||
- [ ] Round-trip test per model: load → save comfy format → reload from the save →
|
||||
outputs match (bf16) / load cleanly (quantized saves). Lands with the
|
||||
Phase 2 comfy save path.
|
||||
- [x] Round-trip verified for the first comfy-save arch: z_image convrot
|
||||
load → comfy save → identical key set + bit-exact quantized entries vs
|
||||
the published file → reload → bit-identical quantized forward. Extend
|
||||
per arch as saves flip.
|
||||
- [ ] Each newly migrated model adds its test in the same PR as its migration.
|
||||
|
||||
## TODO / look at later
|
||||
|
||||
@@ -51,12 +51,15 @@ class OstrisModelMixin:
|
||||
# .safetensors file and no config_path is given
|
||||
aitk_config_repo: Optional[str] = None
|
||||
|
||||
# ---- comfy weight sources (Phase 2 flips the loading default to these) ----
|
||||
# hub repo the comfy-format weight files are published in
|
||||
# ---- comfy weight sources ----
|
||||
# hub repo the comfy-format weight files are published in (Comfy-Org/...)
|
||||
aitk_comfy_repo: Optional[str] = None
|
||||
# standard name_or_path (hub repo id) -> repo-relative comfy weight file
|
||||
# (ComfyUI folder layout: diffusion_models/, text_encoders/, vae/, ...)
|
||||
aitk_comfy_weight_names: Dict[str, str] = {}
|
||||
# candidates for this module: every precision variant the class can digest.
|
||||
# Selection ranks them convrot8 > float8 mixed > float8 > bf16 > fp16
|
||||
# (resolver.comfy_precision_rank); a locally-present candidate always wins
|
||||
# over a download.
|
||||
aitk_comfy_weight_names: Dict[str, List[str]] = {}
|
||||
|
||||
# ---- tokenizer/processor source, for text-encoder modules ----
|
||||
aitk_tokenizer_repo: Optional[str] = None
|
||||
@@ -110,7 +113,12 @@ class OstrisModelMixin:
|
||||
|
||||
@classmethod
|
||||
def aitk_from_config(cls, config):
|
||||
with torch.device("meta"):
|
||||
# params on meta (materialized by the state-dict assign), buffers real:
|
||||
# non-persistent buffers (e.g. rope tables) are computed at init and
|
||||
# never appear in checkpoints
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
with init_empty_weights(include_buffers=False):
|
||||
return cls.from_config(config)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -128,6 +136,7 @@ class OstrisModelMixin:
|
||||
exclude_quant_modules: Optional[List[str]] = None,
|
||||
config_path: Optional[str] = None,
|
||||
subfolder: Optional[str] = None,
|
||||
use_comfy_weights: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
"""Load a model universally from a given name or path.
|
||||
@@ -136,6 +145,13 @@ class OstrisModelMixin:
|
||||
single .safetensors file, or a remote single file
|
||||
("org/repo/file.safetensors").
|
||||
|
||||
When name_or_path is a standard hub repo this class has comfy weight
|
||||
candidates registered for (aitk_comfy_weight_names), the comfy file is
|
||||
the preferred source: the best-ranked local candidate, or the
|
||||
top-preference candidate downloaded into the comfy layout under
|
||||
MODELS_PATH. The standard repo still supplies the config. Local dirs
|
||||
and unregistered repos load as-is; use_comfy_weights=False opts out.
|
||||
|
||||
qtype: quantize the weights after loading. quantize_device: where to run the
|
||||
quantization math; blocks are moved there one at a time and returned to where
|
||||
they were.
|
||||
@@ -149,6 +165,18 @@ class OstrisModelMixin:
|
||||
elif subfolder == "":
|
||||
subfolder = None
|
||||
|
||||
if (
|
||||
use_comfy_weights
|
||||
and not name_or_path.endswith(".safetensors")
|
||||
and not os.path.exists(name_or_path)
|
||||
):
|
||||
comfy_path = cls.resolve_comfy_weights(name_or_path, subfolder=subfolder)
|
||||
if comfy_path is not None:
|
||||
if config_path is None:
|
||||
# the standard repo supplies the config for the comfy file
|
||||
config_path = name_or_path
|
||||
name_or_path = comfy_path
|
||||
|
||||
if name_or_path.endswith(".safetensors"):
|
||||
file_path = cls._resolve_single_file(name_or_path)
|
||||
model = cls._load_single_file(
|
||||
@@ -241,7 +269,9 @@ class OstrisModelMixin:
|
||||
(the tail of the single-file path; also callable directly for
|
||||
checkpoints read from non-safetensors sources)."""
|
||||
state_dict = cls.convert_state_dict_on_load(state_dict)
|
||||
has_quant_markers = any(k.endswith(".comfy_quant") for k in state_dict)
|
||||
has_quant_markers = "scaled_fp8" in state_dict or any(
|
||||
k.endswith(".comfy_quant") for k in state_dict
|
||||
)
|
||||
config = cls.aitk_config_from_state_dict(state_dict)
|
||||
if config is None:
|
||||
config = cls._load_single_file_config(config_path, subfolder)
|
||||
@@ -299,17 +329,42 @@ class OstrisModelMixin:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@classmethod
|
||||
def find_comfy_weights(cls, name_or_path: str) -> Optional[str]:
|
||||
"""Local comfy-format weight file registered for a standard
|
||||
``name_or_path``, or None. Never downloads — Phase 2 flips the loading
|
||||
default to comfy sources; until then callers opt in explicitly."""
|
||||
from toolkit.models.v2.resolver import resolve_comfy_file
|
||||
def resolve_comfy_weights(
|
||||
cls,
|
||||
name_or_path: str,
|
||||
subfolder: Optional[str] = None,
|
||||
local_only: bool = False,
|
||||
hf_token: Optional[str] = None,
|
||||
status_fn: Optional[callable] = None,
|
||||
) -> Optional[str]:
|
||||
"""The comfy-format weight file replacing a standard ``name_or_path``,
|
||||
or None when this class has none registered for it. Best-ranked local
|
||||
candidate wins; otherwise the top-preference candidate is downloaded
|
||||
to the comfy layout under MODELS_PATH (unless local_only).
|
||||
|
||||
rel_path = cls.aitk_comfy_weight_names.get(name_or_path)
|
||||
if rel_path is None:
|
||||
Candidate keys may be plain repo ids or ``(repo_id, subfolder)``
|
||||
tuples for checkpoints holding several of this component (e.g.
|
||||
wan2.2's transformer / transformer_2)."""
|
||||
from toolkit.models.v2.resolver import resolve_comfy_candidates
|
||||
|
||||
candidates = None
|
||||
if subfolder is not None:
|
||||
candidates = cls.aitk_comfy_weight_names.get((name_or_path, subfolder))
|
||||
if candidates is None:
|
||||
candidates = cls.aitk_comfy_weight_names.get(name_or_path)
|
||||
repo_id = cls.aitk_comfy_repo
|
||||
if isinstance(candidates, dict):
|
||||
# entries may carry their own comfy repo ({"repo": ..., "files": [...]})
|
||||
repo_id = candidates.get("repo", repo_id)
|
||||
candidates = candidates.get("files")
|
||||
if not candidates or repo_id is None:
|
||||
return None
|
||||
return resolve_comfy_file(
|
||||
rel_path, repo_id=cls.aitk_comfy_repo, local_only=True
|
||||
return resolve_comfy_candidates(
|
||||
candidates,
|
||||
repo_id=repo_id,
|
||||
hf_token=hf_token,
|
||||
status_fn=status_fn,
|
||||
local_only=local_only,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -374,21 +429,47 @@ class OstrisModelMixin:
|
||||
output_path: str,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
quantized: Optional[bool] = None,
|
||||
):
|
||||
"""Save as a single-file .safetensors in the model's original (comfy)
|
||||
key layout, via convert_state_dict_on_save. Quantized weights are
|
||||
dequantized to full precision; dtype, when given, casts the floating
|
||||
point tensors. (Quantized-storage saves — comfy_quant markers — come
|
||||
with Phase 2.)"""
|
||||
key layout, via convert_state_dict_on_save.
|
||||
|
||||
quantized: None (auto) keeps quantized storage — comfy_quant markers,
|
||||
loadable by ComfyUI and by this class — when the model holds quantized
|
||||
layers whose backend has a comfy format (convrot8 / nvfp4 /
|
||||
convrotcomfyw4a4), and saves dequantized full precision otherwise.
|
||||
True forces quantized storage (raises if any quantized layer's backend
|
||||
has no comfy format); False forces a dequantized save. dtype, when
|
||||
given, casts the floating point (non-quantized) tensors."""
|
||||
from safetensors.torch import save_file
|
||||
from toolkit.util.quantize import dequantize_if_quantized
|
||||
|
||||
q_entries: Dict[str, torch.Tensor] = {}
|
||||
exported: List[str] = []
|
||||
if quantized is None or quantized:
|
||||
from toolkit.util.comfy_quant_export import export_comfy_quantized_layers
|
||||
|
||||
q_entries, exported, unexportable = export_comfy_quantized_layers(self)
|
||||
if unexportable:
|
||||
if quantized:
|
||||
raise ValueError(
|
||||
f"{type(self).__name__}: quantized layers without a comfy "
|
||||
f"storage format: {unexportable[:8]}"
|
||||
)
|
||||
# auto mode: a partially-quantized file would be inconsistent —
|
||||
# fall back to a fully dequantized save
|
||||
q_entries, exported = {}, []
|
||||
|
||||
skip_keys = {f"{name}.weight" for name in exported}
|
||||
state_dict = {}
|
||||
for key, value in self.state_dict().items():
|
||||
if key in skip_keys:
|
||||
continue
|
||||
value = dequantize_if_quantized(value)
|
||||
if dtype is not None and value.is_floating_point():
|
||||
value = value.to(dtype=dtype)
|
||||
state_dict[key] = value.detach().to("cpu").contiguous()
|
||||
state_dict.update(q_entries)
|
||||
state_dict = self.convert_state_dict_on_save(state_dict)
|
||||
|
||||
parent = os.path.dirname(output_path)
|
||||
@@ -479,5 +560,7 @@ class OstrisTransformersMixin(OstrisModelMixin):
|
||||
|
||||
@classmethod
|
||||
def aitk_from_config(cls, config):
|
||||
with torch.device("meta"):
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
with init_empty_weights(include_buffers=False):
|
||||
return cls(config)
|
||||
|
||||
@@ -11,13 +11,36 @@ class QwenImageTransformer2DModel(
|
||||
aitk_subfolder = "transformer"
|
||||
aitk_config_repo = "Qwen/Qwen-Image"
|
||||
|
||||
aitk_comfy_repo = "Comfy-Org/Qwen-Image_ComfyUI"
|
||||
# the comfy files use the diffusers key layout directly (no conversion);
|
||||
# fp8mixed carries float8_e4m3fn comfy_quant markers that attach straight
|
||||
# onto this class's modules
|
||||
aitk_comfy_weight_names = {
|
||||
"Qwen/Qwen-Image": [
|
||||
"split_files/diffusion_models/qwen_image_fp8mixed.safetensors",
|
||||
# raw fp8 cast (no markers, diffusers keys) — loads via from_single_file
|
||||
"split_files/diffusion_models/qwen_image_fp8_e4m3fn.safetensors",
|
||||
"split_files/diffusion_models/qwen_image_bf16.safetensors",
|
||||
],
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["transformer_blocks"]
|
||||
|
||||
@classmethod
|
||||
def _load_single_file(cls, file_path, dtype, config_path=None, subfolder=None):
|
||||
# single-file checkpoints in the wild carry diffusers or original key
|
||||
from safetensors import safe_open
|
||||
|
||||
with safe_open(file_path, framework="pt") as f:
|
||||
has_markers = any(k.endswith(".comfy_quant") for k in f.keys())
|
||||
if has_markers:
|
||||
# comfy prequantized checkpoint (diffusers key layout): the mixin
|
||||
# path attaches the quantized layers
|
||||
return super()._load_single_file(
|
||||
file_path, dtype, config_path=config_path, subfolder=subfolder
|
||||
)
|
||||
# other single-file checkpoints carry diffusers or original key
|
||||
# layouts; diffusers' single-file machinery owns that conversion
|
||||
model = cls.from_single_file(
|
||||
file_path,
|
||||
|
||||
@@ -8,6 +8,79 @@ class WanTransformer3DModel(DiffusersWanTransformer3DModel, OstrisModelMixin):
|
||||
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
aitk_comfy_repo = "Comfy-Org/Wan_2.2_ComfyUI_Repackaged"
|
||||
# comfy wan files use the original key layout; convert_state_dict_on_load
|
||||
# renames them (pure substring renames, so the legacy scaled_fp8 weight /
|
||||
# scale keys ride along with their modules). wan2.2 A14B checkpoints hold
|
||||
# two DiTs, keyed by (repo, subfolder): transformer = high noise,
|
||||
# transformer_2 = low noise. wan2.1 entries override the comfy repo.
|
||||
aitk_comfy_weight_names = {
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": [
|
||||
"split_files/diffusion_models/wan2.2_ti2v_5B_fp16.safetensors",
|
||||
],
|
||||
("Wan-AI/Wan2.2-T2V-A14B-Diffusers", "transformer"): [
|
||||
"split_files/diffusion_models/wan2.2_t2v_high_noise_14B_fp8_scaled.safetensors",
|
||||
],
|
||||
("Wan-AI/Wan2.2-T2V-A14B-Diffusers", "transformer_2"): [
|
||||
"split_files/diffusion_models/wan2.2_t2v_low_noise_14B_fp8_scaled.safetensors",
|
||||
],
|
||||
("Wan-AI/Wan2.2-I2V-A14B-Diffusers", "transformer"): [
|
||||
"split_files/diffusion_models/wan2.2_i2v_high_noise_14B_fp8_scaled.safetensors",
|
||||
],
|
||||
("Wan-AI/Wan2.2-I2V-A14B-Diffusers", "transformer_2"): [
|
||||
"split_files/diffusion_models/wan2.2_i2v_low_noise_14B_fp8_scaled.safetensors",
|
||||
],
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": {
|
||||
"repo": "Comfy-Org/Wan_2.1_ComfyUI_repackaged",
|
||||
"files": [
|
||||
"split_files/diffusion_models/wan2.1_t2v_1.3B_bf16.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_t2v_1.3B_fp16.safetensors",
|
||||
],
|
||||
},
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": {
|
||||
"repo": "Comfy-Org/Wan_2.1_ComfyUI_repackaged",
|
||||
"files": [
|
||||
"split_files/diffusion_models/wan2.1_t2v_14B_fp8_scaled.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_t2v_14B_bf16.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_t2v_14B_fp16.safetensors",
|
||||
],
|
||||
},
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": {
|
||||
"repo": "Comfy-Org/Wan_2.1_ComfyUI_repackaged",
|
||||
"files": [
|
||||
"split_files/diffusion_models/wan2.1_i2v_480p_14B_fp8_scaled.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_i2v_480p_14B_bf16.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_i2v_480p_14B_fp16.safetensors",
|
||||
],
|
||||
},
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": {
|
||||
"repo": "Comfy-Org/Wan_2.1_ComfyUI_repackaged",
|
||||
"files": [
|
||||
"split_files/diffusion_models/wan2.1_i2v_720p_14B_fp8_scaled.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_i2v_720p_14B_bf16.safetensors",
|
||||
"split_files/diffusion_models/wan2.1_i2v_720p_14B_fp16.safetensors",
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["blocks"]
|
||||
|
||||
@classmethod
|
||||
def convert_state_dict_on_load(cls, state_dict):
|
||||
# original/comfy wan keys -> diffusers layout via diffusers' own
|
||||
# converter (rename-only for the base/i2v/vace variants)
|
||||
is_original = any(
|
||||
".self_attn." in k
|
||||
or ".cross_attn." in k
|
||||
or k.startswith("model.diffusion_model.")
|
||||
for k in state_dict
|
||||
)
|
||||
if not is_original:
|
||||
return state_dict
|
||||
from diffusers.loaders.single_file_utils import (
|
||||
convert_wan_transformer_to_diffusers,
|
||||
)
|
||||
|
||||
return convert_wan_transformer_to_diffusers(dict(state_dict))
|
||||
|
||||
@@ -11,6 +11,17 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix
|
||||
# repo to pull the config from when loading a single-file checkpoint
|
||||
aitk_config_repo = "Tongyi-MAI/Z-Image-Turbo"
|
||||
|
||||
aitk_comfy_repo = "Comfy-Org/z_image_turbo"
|
||||
# the int8_convrot file marks the FUSED attention.qkv modules; the load
|
||||
# converter splits those entries exactly into to_q/to_k/to_v (row-sliced
|
||||
# qdata + scales), so it attaches onto this split layout directly
|
||||
aitk_comfy_weight_names = {
|
||||
"Tongyi-MAI/Z-Image-Turbo": [
|
||||
"split_files/diffusion_models/z_image_turbo_int8_convrot.safetensors",
|
||||
"split_files/diffusion_models/z_image_turbo_bf16.safetensors",
|
||||
],
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["layers"]
|
||||
@@ -36,6 +47,22 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix
|
||||
@classmethod
|
||||
def convert_state_dict_on_load(cls, state_dict):
|
||||
"""Convert a single-file Z-Image checkpoint to diffusers transformer keys."""
|
||||
from toolkit.util.comfy_quant_import import split_fused_quantized_keys
|
||||
|
||||
state_dict = dict(state_dict)
|
||||
# quantized fused qkv entries (comfy convrot/fp8/nvfp4 files) split
|
||||
# exactly into the three projections: row-consecutive weights/scales
|
||||
for marker_key in [
|
||||
k for k in list(state_dict) if k.endswith(".attention.qkv.comfy_quant")
|
||||
]:
|
||||
prefix = marker_key[: -len(".comfy_quant")]
|
||||
base = prefix[: -len(".qkv")]
|
||||
split_fused_quantized_keys(
|
||||
state_dict,
|
||||
prefix,
|
||||
[f"{base}.to_q", f"{base}.to_k", f"{base}.to_v"],
|
||||
)
|
||||
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
k = key
|
||||
@@ -47,9 +74,10 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix
|
||||
new_sd[prefix + ".attention.to_k.weight"] = k_proj
|
||||
new_sd[prefix + ".attention.to_v.weight"] = v
|
||||
continue
|
||||
k = k.replace(".attention.out.weight", ".attention.to_out.0.weight")
|
||||
k = k.replace(".attention.q_norm.weight", ".attention.norm_q.weight")
|
||||
k = k.replace(".attention.k_norm.weight", ".attention.norm_k.weight")
|
||||
# module-prefix renames so weight/bias/scale/marker keys all map
|
||||
k = k.replace(".attention.out.", ".attention.to_out.0.")
|
||||
k = k.replace(".attention.q_norm.", ".attention.norm_q.")
|
||||
k = k.replace(".attention.k_norm.", ".attention.norm_k.")
|
||||
if k.startswith("x_embedder."):
|
||||
k = "all_x_embedder.2-1." + k[len("x_embedder.") :]
|
||||
elif k.startswith("final_layer."):
|
||||
@@ -60,6 +88,20 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix
|
||||
@classmethod
|
||||
def convert_state_dict_on_save(cls, state_dict):
|
||||
"""Convert a diffusers transformer state dict back to the single-file layout."""
|
||||
from toolkit.util.comfy_quant_import import fuse_split_quantized_keys
|
||||
|
||||
state_dict = dict(state_dict)
|
||||
# quantized split projections fuse back into the single-file qkv entry
|
||||
for marker_key in [
|
||||
k for k in list(state_dict) if k.endswith(".attention.to_q.comfy_quant")
|
||||
]:
|
||||
base = marker_key[: -len(".to_q.comfy_quant")]
|
||||
fuse_split_quantized_keys(
|
||||
state_dict,
|
||||
[f"{base}.to_q", f"{base}.to_k", f"{base}.to_v"],
|
||||
f"{base}.qkv",
|
||||
)
|
||||
|
||||
new_sd = {}
|
||||
qkv_cache = {}
|
||||
for key, value in state_dict.items():
|
||||
@@ -90,9 +132,9 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix
|
||||
break
|
||||
if matched:
|
||||
continue
|
||||
k = k.replace(".attention.to_out.0.weight", ".attention.out.weight")
|
||||
k = k.replace(".attention.norm_q.weight", ".attention.q_norm.weight")
|
||||
k = k.replace(".attention.norm_k.weight", ".attention.k_norm.weight")
|
||||
k = k.replace(".attention.to_out.0.", ".attention.out.")
|
||||
k = k.replace(".attention.norm_q.", ".attention.q_norm.")
|
||||
k = k.replace(".attention.norm_k.", ".attention.k_norm.")
|
||||
if k.startswith("all_x_embedder.2-1."):
|
||||
k = "x_embedder." + k[len("all_x_embedder.2-1.") :]
|
||||
elif k.startswith("all_final_layer.2-1."):
|
||||
|
||||
@@ -16,6 +16,78 @@ from typing import Callable, Iterable, Optional
|
||||
from toolkit.paths import MODELS_PATH
|
||||
|
||||
|
||||
def comfy_precision_rank(filename: str) -> int:
|
||||
"""Load-preference rank for a comfy weight filename:
|
||||
convrot8 (0) > float8 mixed (1) > float8 (2) > bf16 (3) > fp16 (4) >
|
||||
anything else, e.g. nvfp4 or unmarked (5)."""
|
||||
name = os.path.basename(filename).lower()
|
||||
if "convrot" in name:
|
||||
return 0
|
||||
is_fp8 = "fp8" in name or "float8" in name or "e4m3" in name
|
||||
if is_fp8 and "mixed" in name:
|
||||
return 1
|
||||
if is_fp8:
|
||||
return 2
|
||||
if "bf16" in name:
|
||||
return 3
|
||||
if "fp16" in name:
|
||||
return 4
|
||||
return 5
|
||||
|
||||
|
||||
def comfy_local_rel(repo_rel: str) -> str:
|
||||
"""Repo file path -> ComfyUI models-folder path. Comfy-Org repos nest the
|
||||
comfy layout under a packaging prefix (split_files/, non_official/) that is
|
||||
not part of the shared models folder layout."""
|
||||
for prefix in ("split_files/", "non_official/"):
|
||||
if repo_rel.startswith(prefix):
|
||||
return repo_rel[len(prefix):]
|
||||
return repo_rel
|
||||
|
||||
|
||||
def resolve_comfy_candidates(
|
||||
candidates: Iterable[str],
|
||||
repo_id: str,
|
||||
hf_token: Optional[str] = None,
|
||||
status_fn: Optional[Callable[[str], None]] = None,
|
||||
local_only: bool = False,
|
||||
) -> Optional[str]:
|
||||
"""Pick the best comfy weight file among precision variants of one
|
||||
component (repo-relative paths, ranked by comfy_precision_rank then list
|
||||
order). The best-ranked LOCAL candidate wins; only when no candidate is
|
||||
local is the best-ranked one downloaded to its comfy-layout location
|
||||
under MODELS_PATH."""
|
||||
candidates = list(candidates)
|
||||
ordered = sorted(
|
||||
candidates, key=lambda c: (comfy_precision_rank(c), candidates.index(c))
|
||||
)
|
||||
for repo_rel in ordered:
|
||||
found = resolve_comfy_file(
|
||||
comfy_local_rel(repo_rel), repo_id, local_only=True
|
||||
)
|
||||
if found is not None:
|
||||
return found
|
||||
if local_only:
|
||||
return None
|
||||
|
||||
import huggingface_hub
|
||||
|
||||
best = ordered[0]
|
||||
local_rel = comfy_local_rel(best)
|
||||
if status_fn is not None:
|
||||
status_fn(f"Downloading {best} from {repo_id} into {MODELS_PATH}")
|
||||
path = huggingface_hub.hf_hub_download(
|
||||
repo_id=repo_id, filename=best, token=hf_token, local_dir=MODELS_PATH
|
||||
)
|
||||
target = os.path.join(MODELS_PATH, local_rel)
|
||||
if os.path.abspath(path) != os.path.abspath(target):
|
||||
# move out of the packaging prefix into the shared comfy layout
|
||||
os.makedirs(os.path.dirname(target), exist_ok=True)
|
||||
os.replace(path, target)
|
||||
return target
|
||||
return path
|
||||
|
||||
|
||||
def find_file_recursive(root_dir: str, filename: str) -> Optional[str]:
|
||||
"""First (breadth-stable, sorted) match of ``filename`` anywhere under
|
||||
``root_dir``."""
|
||||
|
||||
104
toolkit/util/comfy_quant_export.py
Normal file
104
toolkit/util/comfy_quant_export.py
Normal file
@@ -0,0 +1,104 @@
|
||||
"""Export toolkit-quantized modules into ComfyUI ``comfy_quant`` checkpoints —
|
||||
the inverse of toolkit/util/comfy_quant_import.py.
|
||||
|
||||
Every quantized OstrisLinear whose backend has a comfy storage format emits
|
||||
its quantized tensors plus the ``<prefix>.comfy_quant`` uint8 JSON marker:
|
||||
|
||||
- convrot8 (int8_tensorwise + convrot): weight int8 [out, in], fp32
|
||||
weight_scale [out, 1] (comfy_kitchen's per-channel convention)
|
||||
- nvfp4: high-nibble-first packed fp4 pairs, e4m3 block scales re-swizzled
|
||||
to the cuBLAS 128x4 tile layout, fp32 weight_scale_2 per-tensor scale and
|
||||
optional AWQ pre_quant_scale
|
||||
- float8_e4m3fn: fp8_e4m3 weight + one fp32 per-tensor weight_scale
|
||||
- convrotcomfyw4a4: via convrot_quant.export_comfy_convrot_w4a4
|
||||
|
||||
Biases are NOT emitted here — they are ordinary parameters and flow through
|
||||
the regular state_dict path.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from toolkit.util.ostris_quant import OstrisLinear
|
||||
|
||||
|
||||
def comfy_quant_marker(conf: dict) -> torch.Tensor:
|
||||
return torch.tensor(list(json.dumps(conf).encode("utf-8")), dtype=torch.uint8)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def export_comfy_quantized_layers(
|
||||
root: torch.nn.Module,
|
||||
) -> Tuple[Dict[str, torch.Tensor], List[str], List[str]]:
|
||||
"""Comfy-format state-dict entries for every quantized OstrisLinear in
|
||||
``root``. Returns ``(entries, exported_names, unexportable_names)`` —
|
||||
entries are keyed by module path in root's layout (run
|
||||
convert_state_dict_on_save afterwards for the checkpoint layout);
|
||||
unexportable_names lists quantized modules whose backend has no comfy
|
||||
storage format (the caller decides whether to dequantize instead)."""
|
||||
from toolkit.util.convrot_quant import (
|
||||
ConvRotComfyW4A4Quantizer,
|
||||
export_comfy_convrot_w4a4,
|
||||
)
|
||||
from toolkit.util.nvfp4_quant import swap_nvfp4_nibbles, swizzle_nvfp4_scales
|
||||
|
||||
entries: Dict[str, torch.Tensor] = {}
|
||||
exported: List[str] = []
|
||||
unexportable: List[str] = []
|
||||
|
||||
for name, module in root.named_modules():
|
||||
if not isinstance(module, OstrisLinear):
|
||||
continue
|
||||
|
||||
if hasattr(module, "cr8_qdata"):
|
||||
rot = int(getattr(module, "cr8_rot_size", 1) or 1)
|
||||
conf = {"format": "int8_tensorwise"}
|
||||
if rot > 1:
|
||||
conf.update({"convrot": True, "convrot_groupsize": rot})
|
||||
entries[f"{name}.weight"] = module.cr8_qdata.detach().cpu().contiguous()
|
||||
entries[f"{name}.weight_scale"] = (
|
||||
module.cr8_scales.view(torch.float32)
|
||||
.detach()
|
||||
.cpu()
|
||||
.reshape(module.out_features, 1)
|
||||
.contiguous()
|
||||
)
|
||||
entries[f"{name}.comfy_quant"] = comfy_quant_marker(conf)
|
||||
elif hasattr(module, "nv4_qdata"):
|
||||
entries[f"{name}.weight"] = swap_nvfp4_nibbles(
|
||||
module.nv4_qdata.detach().cpu()
|
||||
)
|
||||
scales = module.nv4_scales.view(torch.float8_e4m3fn).detach().cpu()
|
||||
entries[f"{name}.weight_scale"] = swizzle_nvfp4_scales(
|
||||
scales.reshape(module.out_features, module.in_features // 16)
|
||||
).view(torch.float8_e4m3fn)
|
||||
entries[f"{name}.weight_scale_2"] = (
|
||||
module.nv4_pts.view(torch.float32).detach().cpu().reshape(())
|
||||
)
|
||||
if hasattr(module, "nv4_pre_scale"):
|
||||
entries[f"{name}.pre_quant_scale"] = (
|
||||
module.nv4_pre_scale.view(torch.float32).detach().cpu().contiguous()
|
||||
)
|
||||
entries[f"{name}.comfy_quant"] = comfy_quant_marker({"format": "nvfp4"})
|
||||
elif hasattr(module, "f8_qdata"):
|
||||
entries[f"{name}.weight"] = module.f8_qdata.detach().cpu().contiguous()
|
||||
entries[f"{name}.weight_scale"] = (
|
||||
module.f8_scale.view(torch.float32).detach().cpu().reshape(())
|
||||
)
|
||||
entries[f"{name}.comfy_quant"] = comfy_quant_marker(
|
||||
{"format": "float8_e4m3fn", "full_precision_matrix_mult": True}
|
||||
)
|
||||
elif isinstance(module.ostris_quantizer, ConvRotComfyW4A4Quantizer):
|
||||
layer_entries = export_comfy_convrot_w4a4(module, f"{name}.")
|
||||
layer_entries.pop(f"{name}.bias", None)
|
||||
entries.update(
|
||||
{k: v.detach().cpu() if torch.is_tensor(v) else v for k, v in layer_entries.items()}
|
||||
)
|
||||
else:
|
||||
unexportable.append(name)
|
||||
continue
|
||||
exported.append(name)
|
||||
|
||||
return entries, exported, unexportable
|
||||
@@ -71,6 +71,135 @@ class Int8Embedding(torch.nn.Module):
|
||||
return out.to(input_ids.device).reshape(*input_ids.shape, self.embedding_dim)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def split_fused_quantized_keys(
|
||||
state_dict: Dict[str, torch.Tensor],
|
||||
prefix: str,
|
||||
dst_prefixes,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Split one fused quantized comfy entry (``<prefix>.weight`` /
|
||||
``.weight_scale`` / ``.comfy_quant`` / ...) into equal row ranges under
|
||||
``dst_prefixes`` (out-dim concat order). Exact for every supported format:
|
||||
int8 rows and their per-row scales slice; fp8's per-tensor scale and
|
||||
nvfp4's weight_scale_2 / pre_quant_scale are shared by every split; nvfp4
|
||||
block scales are unswizzled, row-split, and re-swizzled. Mutates and
|
||||
returns state_dict. Used by classes whose module layout splits a fused
|
||||
checkpoint projection (e.g. qkv -> to_q/to_k/to_v)."""
|
||||
from toolkit.util.nvfp4_quant import swizzle_nvfp4_scales
|
||||
|
||||
marker = state_dict.pop(f"{prefix}.comfy_quant")
|
||||
conf = parse_comfy_quant_blob(marker)
|
||||
fmt = conf.get("format")
|
||||
|
||||
weight = state_dict.pop(f"{prefix}.weight")
|
||||
scale = state_dict.pop(f"{prefix}.weight_scale", None)
|
||||
pts = state_dict.pop(f"{prefix}.weight_scale_2", None)
|
||||
pre = state_dict.pop(f"{prefix}.pre_quant_scale", None)
|
||||
bias = state_dict.pop(f"{prefix}.bias", None)
|
||||
state_dict.pop(f"{prefix}.input_scale", None)
|
||||
|
||||
n = len(dst_prefixes)
|
||||
if weight.shape[0] % n != 0:
|
||||
raise ValueError(
|
||||
f"{prefix}: fused out dim {weight.shape[0]} does not split into {n}"
|
||||
)
|
||||
rows = weight.shape[0] // n
|
||||
|
||||
scale_parts = None
|
||||
if scale is not None:
|
||||
if fmt == "nvfp4":
|
||||
in_features = weight.shape[1] * 2 # packed fp4 pairs
|
||||
full = unswizzle_nvfp4_scales(
|
||||
scale.view(torch.float8_e4m3fn), weight.shape[0], in_features // 16
|
||||
)
|
||||
scale_parts = [
|
||||
swizzle_nvfp4_scales(p).view(torch.float8_e4m3fn)
|
||||
for p in full.split(rows, dim=0)
|
||||
]
|
||||
elif scale.ndim == 0 or scale.numel() == 1:
|
||||
scale_parts = [scale.clone() for _ in range(n)]
|
||||
else:
|
||||
scale_parts = list(scale.reshape(weight.shape[0], -1).split(rows, dim=0))
|
||||
|
||||
for i, dst in enumerate(dst_prefixes):
|
||||
state_dict[f"{dst}.comfy_quant"] = marker.clone()
|
||||
state_dict[f"{dst}.weight"] = weight[i * rows : (i + 1) * rows].contiguous()
|
||||
if scale_parts is not None:
|
||||
state_dict[f"{dst}.weight_scale"] = scale_parts[i].contiguous()
|
||||
if pts is not None:
|
||||
state_dict[f"{dst}.weight_scale_2"] = pts.clone()
|
||||
if pre is not None:
|
||||
state_dict[f"{dst}.pre_quant_scale"] = pre.clone()
|
||||
if bias is not None:
|
||||
state_dict[f"{dst}.bias"] = bias[i * rows : (i + 1) * rows].contiguous()
|
||||
return state_dict
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def fuse_split_quantized_keys(
|
||||
state_dict: Dict[str, torch.Tensor],
|
||||
src_prefixes,
|
||||
prefix: str,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Inverse of split_fused_quantized_keys: concatenate N split quantized
|
||||
comfy entries back into one fused entry (out-dim concat in src order).
|
||||
All parts must share the same format config; fp8 parts must share the same
|
||||
per-tensor scale (true for entries produced by the splitter). Mutates and
|
||||
returns state_dict."""
|
||||
from toolkit.util.nvfp4_quant import swizzle_nvfp4_scales
|
||||
|
||||
markers = [state_dict.pop(f"{p}.comfy_quant") for p in src_prefixes]
|
||||
confs = [parse_comfy_quant_blob(m) for m in markers]
|
||||
if any(c != confs[0] for c in confs[1:]):
|
||||
raise ValueError(f"{prefix}: split parts carry different quant configs")
|
||||
fmt = confs[0].get("format")
|
||||
|
||||
weights = [state_dict.pop(f"{p}.weight") for p in src_prefixes]
|
||||
scales = [state_dict.pop(f"{p}.weight_scale", None) for p in src_prefixes]
|
||||
ptss = [state_dict.pop(f"{p}.weight_scale_2", None) for p in src_prefixes]
|
||||
pres = [state_dict.pop(f"{p}.pre_quant_scale", None) for p in src_prefixes]
|
||||
biases = [state_dict.pop(f"{p}.bias", None) for p in src_prefixes]
|
||||
|
||||
state_dict[f"{prefix}.comfy_quant"] = markers[0]
|
||||
weight = torch.cat(weights, dim=0).contiguous()
|
||||
state_dict[f"{prefix}.weight"] = weight
|
||||
if scales[0] is not None:
|
||||
if fmt == "nvfp4":
|
||||
in_features = weight.shape[1] * 2
|
||||
rows = [w.shape[0] for w in weights]
|
||||
full = torch.cat(
|
||||
[
|
||||
unswizzle_nvfp4_scales(
|
||||
s.view(torch.float8_e4m3fn), r, in_features // 16
|
||||
)
|
||||
for s, r in zip(scales, rows)
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
state_dict[f"{prefix}.weight_scale"] = swizzle_nvfp4_scales(full).view(
|
||||
torch.float8_e4m3fn
|
||||
)
|
||||
elif scales[0].ndim == 0 or scales[0].numel() == 1:
|
||||
if any(
|
||||
not torch.equal(s.reshape(-1), scales[0].reshape(-1)) for s in scales[1:]
|
||||
):
|
||||
raise ValueError(
|
||||
f"{prefix}: per-tensor scales differ across split parts"
|
||||
)
|
||||
state_dict[f"{prefix}.weight_scale"] = scales[0]
|
||||
else:
|
||||
state_dict[f"{prefix}.weight_scale"] = torch.cat(
|
||||
[s.reshape(w.shape[0], -1) for s, w in zip(scales, weights)], dim=0
|
||||
).contiguous()
|
||||
if ptss[0] is not None:
|
||||
state_dict[f"{prefix}.weight_scale_2"] = ptss[0]
|
||||
if pres[0] is not None:
|
||||
state_dict[f"{prefix}.pre_quant_scale"] = pres[0]
|
||||
if biases[0] is not None:
|
||||
state_dict[f"{prefix}.bias"] = torch.cat(biases, dim=0).contiguous()
|
||||
return state_dict
|
||||
|
||||
|
||||
def _to_ostris(module: torch.nn.Linear, quantizer, orig_dtype: torch.dtype) -> OstrisLinear:
|
||||
if "weight" in module._parameters:
|
||||
del module._parameters["weight"]
|
||||
@@ -128,7 +257,19 @@ def import_comfy_quantized_layers(
|
||||
"expected nn.Linear or nn.Embedding"
|
||||
)
|
||||
|
||||
if fmt == "int8_tensorwise":
|
||||
if fmt == "float8_e4m3fn":
|
||||
# fp8_e4m3 weight + fp32 per-tensor scale, dequantized matmul
|
||||
from toolkit.util.float8_quant import Float8Quantizer
|
||||
|
||||
quantizer = get_ostris_quantizer("float8_e4m3fn")
|
||||
Float8Quantizer.attach_(
|
||||
module,
|
||||
weight.view(torch.float8_e4m3fn)
|
||||
if weight.dtype != torch.float8_e4m3fn
|
||||
else weight,
|
||||
weight_scale,
|
||||
)
|
||||
elif fmt == "int8_tensorwise":
|
||||
rot = int(conf.get("convrot_groupsize", 256)) if conf.get("convrot") else 1
|
||||
quantizer = get_ostris_quantizer("convrot8")
|
||||
module.register_buffer("cr8_qdata", weight.contiguous(), persistent=False)
|
||||
@@ -158,7 +299,7 @@ def import_comfy_quantized_layers(
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported comfy quant format {fmt!r} on {prefix} "
|
||||
"(supported: int8_tensorwise, nvfp4)"
|
||||
"(supported: int8_tensorwise, nvfp4, float8_e4m3fn)"
|
||||
)
|
||||
|
||||
# drop unused calibration extras if present
|
||||
@@ -174,4 +315,40 @@ def import_comfy_quantized_layers(
|
||||
)
|
||||
converted += 1
|
||||
|
||||
# legacy ComfyUI scaled-fp8 checkpoints (e.g. the wan *_fp8_scaled files):
|
||||
# a top-level ``scaled_fp8`` marker tensor plus per-layer fp8 ``weight``
|
||||
# and scalar fp32 ``scale_weight`` — the float8 backend's exact storage.
|
||||
# ``scale_input`` (activation quant) is dropped; matmuls run dequantized.
|
||||
if "scaled_fp8" in state_dict:
|
||||
from toolkit.util.float8_quant import Float8Quantizer
|
||||
|
||||
state_dict.pop("scaled_fp8")
|
||||
for scale_key in [k for k in state_dict if k.endswith(".scale_weight")]:
|
||||
prefix = scale_key[: -len(".scale_weight")]
|
||||
module_path = key_map(prefix) if key_map is not None else prefix
|
||||
module = root.get_submodule(module_path)
|
||||
if not isinstance(module, torch.nn.Linear):
|
||||
raise ValueError(
|
||||
f"scaled_fp8 entry {prefix} points at {type(module).__name__}, "
|
||||
"expected nn.Linear"
|
||||
)
|
||||
weight = state_dict.pop(f"{prefix}.weight")
|
||||
scale = state_dict.pop(scale_key)
|
||||
state_dict.pop(f"{prefix}.scale_input", None)
|
||||
quantizer = get_ostris_quantizer("float8_e4m3fn")
|
||||
Float8Quantizer.attach_(
|
||||
module,
|
||||
weight
|
||||
if weight.dtype == torch.float8_e4m3fn
|
||||
else weight.view(torch.float8_e4m3fn),
|
||||
scale,
|
||||
)
|
||||
_to_ostris(module, quantizer, orig_dtype)
|
||||
bias = state_dict.pop(f"{prefix}.bias", None)
|
||||
if bias is not None and module.bias is not None:
|
||||
module._parameters["bias"] = torch.nn.Parameter(
|
||||
bias.detach().clone(), requires_grad=False
|
||||
)
|
||||
converted += 1
|
||||
|
||||
return state_dict, converted
|
||||
|
||||
53
toolkit/util/float8_quant.py
Normal file
53
toolkit/util/float8_quant.py
Normal file
@@ -0,0 +1,53 @@
|
||||
"""ComfyUI-style float8 weight storage as an Ostris backend.
|
||||
|
||||
Matches the comfy_quant ``{"format": "float8_e4m3fn",
|
||||
"full_precision_matrix_mult": true}`` layout: the weight stored as
|
||||
torch.float8_e4m3fn plus one fp32 per-tensor scale, matmuls running on the
|
||||
dequantized weight (W8A16 numerics). Used both to import comfy fp8/fp8-mixed
|
||||
checkpoints and to quantize/export in that format.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from toolkit.util.ostris_quant import OstrisLinear, OstrisQuantizer
|
||||
|
||||
FLOAT8_QTYPES = ["float8_e4m3fn"]
|
||||
|
||||
F8_MAX = torch.finfo(torch.float8_e4m3fn).max
|
||||
|
||||
|
||||
class Float8Quantizer(OstrisQuantizer):
|
||||
"""fp8_e4m3 weight + fp32 per-tensor scale, dequantized matmul."""
|
||||
|
||||
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None:
|
||||
scale = (weight_fp32.abs().max() / F8_MAX).clamp(min=1e-12)
|
||||
q = (weight_fp32 / scale).clamp(-F8_MAX, F8_MAX).to(torch.float8_e4m3fn)
|
||||
self.attach_(module, q, scale)
|
||||
|
||||
@staticmethod
|
||||
def attach_(
|
||||
module: torch.nn.Module,
|
||||
qweight: torch.Tensor, # float8_e4m3fn (out, in)
|
||||
scale: torch.Tensor, # fp32 scalar
|
||||
) -> None:
|
||||
"""Register the quantized representation on the module. Used both by
|
||||
quantize_ and by importers of pre-quantized checkpoints."""
|
||||
module.register_buffer("f8_qdata", qweight.contiguous(), persistent=False)
|
||||
module.register_buffer(
|
||||
"f8_scale",
|
||||
scale.detach().float().clone().reshape(1).view(torch.uint8),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
def dequantize(self, module: "OstrisLinear") -> torch.Tensor:
|
||||
scale = module.f8_scale.view(torch.float32)[0]
|
||||
return module.f8_qdata.to(torch.float32) * scale
|
||||
|
||||
@torch.no_grad()
|
||||
def requantize_(self, module: "OstrisLinear", fp_weight: torch.Tensor) -> None:
|
||||
w = fp_weight.to(torch.float32)
|
||||
scale = (w.abs().max() / F8_MAX).clamp(min=1e-12)
|
||||
module.f8_qdata.copy_((w / scale).clamp(-F8_MAX, F8_MAX).to(torch.float8_e4m3fn))
|
||||
module.f8_scale.copy_(scale.reshape(1).view(torch.uint8))
|
||||
@@ -59,6 +59,25 @@ def swap_nvfp4_nibbles(packed: torch.Tensor) -> torch.Tensor:
|
||||
return ((packed << 4) | (packed >> 4)).contiguous()
|
||||
|
||||
|
||||
def swizzle_nvfp4_scales(scales: torch.Tensor) -> torch.Tensor:
|
||||
"""Inverse of unswizzle_nvfp4_scales: row-major (rows, cols) block scales
|
||||
into the cuBLAS 128x4-tile layout ComfyUI checkpoints store (comfy_kitchen's
|
||||
``to_blocked``). Pads to tile boundaries when needed."""
|
||||
rows, cols = scales.shape
|
||||
n_row_blocks = (rows + 127) // 128
|
||||
n_col_blocks = (cols + 3) // 4
|
||||
padded_rows = n_row_blocks * 128
|
||||
padded_cols = n_col_blocks * 4
|
||||
if (padded_rows, padded_cols) != (rows, cols):
|
||||
padded = scales.new_zeros(padded_rows, padded_cols)
|
||||
padded[:rows, :cols] = scales
|
||||
scales = padded
|
||||
x = scales.reshape(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
|
||||
x = x.reshape(n_row_blocks, n_col_blocks, 4, 32, 4)
|
||||
x = x.transpose(2, 3).reshape(-1, 32, 16)
|
||||
return x.reshape(-1).contiguous()
|
||||
|
||||
|
||||
class Nvfp4Quantizer(OstrisQuantizer):
|
||||
"""Block-16 nvfp4 weights, full-precision activations. One instance is
|
||||
shareable across modules."""
|
||||
|
||||
@@ -186,6 +186,7 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]:
|
||||
"""Resolve a qtype string to a quantizer backend instance, or None if the qtype
|
||||
does not belong to a custom backend. Add new backends here."""
|
||||
from toolkit.util.convrot_quant import CONVROT_QTYPES, get_convrot_quantizer
|
||||
from toolkit.util.float8_quant import FLOAT8_QTYPES, Float8Quantizer
|
||||
from toolkit.util.nvfp4_quant import NVFP4_QTYPES, Nvfp4Quantizer
|
||||
from toolkit.util.orbit_quant import ORBIT_QTYPES, OrbitQuantizer
|
||||
from toolkit.util.orbit_vq_quant import ORBIT_VQ_QTYPES, OrbitVQQuantizer
|
||||
@@ -200,6 +201,8 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]:
|
||||
quantizer = get_convrot_quantizer(qtype)
|
||||
elif qtype in NVFP4_QTYPES:
|
||||
quantizer = Nvfp4Quantizer()
|
||||
elif qtype in FLOAT8_QTYPES:
|
||||
quantizer = Float8Quantizer()
|
||||
elif qtype in UINTX_QTYPES:
|
||||
quantizer = UIntXQuantizer(UINTX_QTYPES[qtype])
|
||||
if quantizer is not None:
|
||||
|
||||
Reference in New Issue
Block a user