This commit is contained in:
Jaret Burkett
2026-08-27 15:37:13 -06:00
parent 9ed2e0b8e7
commit 9113420b61
13 changed files with 799 additions and 76 deletions

View File

@@ -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")

View File

@@ -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")

View File

@@ -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

View File

@@ -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)

View File

@@ -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,

View File

@@ -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))

View File

@@ -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."):

View File

@@ -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``."""

View 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

View File

@@ -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

View 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))

View File

@@ -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."""

View File

@@ -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: