diff --git a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py index 54116a5..f1ad25e 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py @@ -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") diff --git a/extensions_built_in/diffusion_models/z_image/z_image.py b/extensions_built_in/diffusion_models/z_image/z_image.py index 67ffd90..bc2a30e 100644 --- a/extensions_built_in/diffusion_models/z_image/z_image.py +++ b/extensions_built_in/diffusion_models/z_image/z_image.py @@ -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") diff --git a/toolkit/models/v2/PLANNING.md b/toolkit/models/v2/PLANNING.md index b5d7fd1..0048e19 100644 --- a/toolkit/models/v2/PLANNING.md +++ b/toolkit/models/v2/PLANNING.md @@ -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 diff --git a/toolkit/models/v2/_mixin.py b/toolkit/models/v2/_mixin.py index cacc49f..33b9fa9 100644 --- a/toolkit/models/v2/_mixin.py +++ b/toolkit/models/v2/_mixin.py @@ -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) diff --git a/toolkit/models/v2/diffusion_models/qwen_image.py b/toolkit/models/v2/diffusion_models/qwen_image.py index 6d5b7d7..3c4e689 100644 --- a/toolkit/models/v2/diffusion_models/qwen_image.py +++ b/toolkit/models/v2/diffusion_models/qwen_image.py @@ -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, diff --git a/toolkit/models/v2/diffusion_models/wan.py b/toolkit/models/v2/diffusion_models/wan.py index 2ba9e88..98ddd1c 100644 --- a/toolkit/models/v2/diffusion_models/wan.py +++ b/toolkit/models/v2/diffusion_models/wan.py @@ -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)) diff --git a/toolkit/models/v2/diffusion_models/z_image.py b/toolkit/models/v2/diffusion_models/z_image.py index 5ee32c4..794c3ec 100644 --- a/toolkit/models/v2/diffusion_models/z_image.py +++ b/toolkit/models/v2/diffusion_models/z_image.py @@ -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."): diff --git a/toolkit/models/v2/resolver.py b/toolkit/models/v2/resolver.py index fda8ff0..adef8de 100644 --- a/toolkit/models/v2/resolver.py +++ b/toolkit/models/v2/resolver.py @@ -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``.""" diff --git a/toolkit/util/comfy_quant_export.py b/toolkit/util/comfy_quant_export.py new file mode 100644 index 0000000..657c2c6 --- /dev/null +++ b/toolkit/util/comfy_quant_export.py @@ -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 ``.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 diff --git a/toolkit/util/comfy_quant_import.py b/toolkit/util/comfy_quant_import.py index 833d6b2..6cd16b6 100644 --- a/toolkit/util/comfy_quant_import.py +++ b/toolkit/util/comfy_quant_import.py @@ -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 (``.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 diff --git a/toolkit/util/float8_quant.py b/toolkit/util/float8_quant.py new file mode 100644 index 0000000..0ac2ff4 --- /dev/null +++ b/toolkit/util/float8_quant.py @@ -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)) diff --git a/toolkit/util/nvfp4_quant.py b/toolkit/util/nvfp4_quant.py index 72d7155..ef33147 100644 --- a/toolkit/util/nvfp4_quant.py +++ b/toolkit/util/nvfp4_quant.py @@ -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.""" diff --git a/toolkit/util/ostris_quant.py b/toolkit/util/ostris_quant.py index 6384674..e5bd9c6 100644 --- a/toolkit/util/ostris_quant.py +++ b/toolkit/util/ostris_quant.py @@ -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: