This commit is contained in:
Jaret Burkett
2026-08-27 18:18:31 -06:00
parent 45886f01b2
commit 520d96aac3
27 changed files with 333 additions and 253 deletions

View File

@@ -207,8 +207,11 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord
variants inherit
- [x] nucleus_image — `v2/diffusion_models/nucleus_image.py`, TE stanza
collapsed to prepare_text_encoder
- [ ] krea2, ideogram4, mageflow — TE/VAE migrated; their custom local DiT
classes still to be rebased onto the mixin
- [x] krea2, ideogram4, mageflow — DiTs on the mixin via the `config=`
passthrough (holder builds config from model_kwargs/holder state, mixin
does build/markers/whitelist/casting; plain nn.Module classes now work
with the default builder). krea2 + ideogram4 verified by harness;
mageflow untestable while its repo 404s
- [x] chroma, chroma_radiance — both vendored Chroma classes now carry
`OstrisModelMixin` with the block-count sniff moved into a new
`aitk_config_from_state_dict` hook (mixin now supports checkpoint-derived
@@ -217,14 +220,23 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord
depends on holder state (patch_size), not the checkpoint
- [x] flux_kontext — `v2/diffusion_models/flux.py` (FluxTransformer2DModel);
whole model now loads through v2 (transformer, T5, CLIP, KLVAE)
- [ ] flux2 — TE/VAE partially migrated (flux2_kl); custom DiT still local.
krea2/mageflow/ideogram4/zeta_chroma DiTs stay model-specific: their
configs come from model_kwargs / holder state, so the mixin adds nothing
until the comfy-weights flip (Phase 2)
- [ ] minimax_h3 (+ ref2va), ltx2 family — already on the shared resolver +
comfy_quant_import; the full mixin port waits for Phase 2, when the
mixin's single-file precision policy (stored-precision loading, fp32-key
protection) is settled to match their deliberate behavior
- [x] flux2 + zeta_chroma DiTs on the mixin via `config=` passthrough
(flux2_klein_4b verified by harness). Every arch's DiT now loads
through the mixin except: ltx2 family (one-file→two-modules split,
helper delegated), anima (diffusers modular pipeline), ace_step
(bundled single-file loader), z_image_l2p's local subclass, and the
grandfathered legacy stable_diffusion_model archs
- [x] minimax_h3 (+ ref2va) transformer ported to the mixin: config sniffed
from the checkpoint via aitk_config_from_state_dict (adaln_t_table),
marker attach + stored-precision load via the new
`aitk_cast_on_load = False` knob. Verified on the real pruned convrot
file: 200 ConvRot linears, pruned table detected, fp32/fp16/bf16 mix
preserved, no meta leftovers. Its TE stays custom (50-layer truncation
+ key_map). ltx2.5's `_load_quantized_module` now delegates to the
mixin's whitelist/meta helper (~30 lines deleted); its full port is
blocked on the one-comfy-file → transformer+connectors split, which
doesn't fit the per-class single-file shape — revisit with the live
server's component model
- [x] wan21 / wan22 family — `v2/diffusion_models/wan.py`
(WanTransformer3DModel, both wan22 dual loads included) +
`v2/text_encoders/umt5.py` (UMT5TextEncoder + PatchedT5Tokenizer;
@@ -307,9 +319,16 @@ Decisions:
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.
- [x] wan comfy-format saves: `convert_state_dict_on_save` inverts diffusers'
rename table (base/t2v/i2v; vace/animate excluded — their reverse
mappings collide). Round-trip verified on both real comfy files
(exact; the 2.1 file's legacy model.diffusion_model. prefix drops per
the modern convention) and with real weights (load → save → 825-key
original-layout file → reload bit-equal). wan21 + wan22_5b save one
comfy file; wan22_14b saves the comfy-standard _high_noise/_low_noise
pair instead of two diffusers folders.
- [ ] 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)
conversions are verified per file)
- [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
@@ -333,6 +352,14 @@ Decisions:
bit-exact fused qkv weights/scales/markers, and the reload's quantized
forward is bit-identical — toolkit saves are byte-compatible with
ComfyUI.
- [x] chroma + chroma_radiance saves flipped to the mixin (their class keys
ARE the original layout) — also fixes their quanto-only dequant bug
(torchao/Ostris weights now dequantize on save). Tiny-model round trip
verified. Save flips so far: z_image, qwen_image, wan21, wan22_5b,
wan22_14b (dual files), chroma ×2.
- [ ] flux_kontext comfy wiring deferred: its Comfy-Org repo ships a single
legacy-fp8 file in fused BFL layout — needs the flux fused-split
conversion (split_fused_quantized_keys pattern + BFL↔diffusers maps)
- [ ] 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
@@ -354,6 +381,13 @@ Decisions:
- [x] Missing weights skip rather than fail: default is HF_HUB_OFFLINE=1 and
hub/file errors classify as SKIP; `--allow-download` opts into fetching.
(GPU + local-weights test, not CI-portable.)
- [x] Final certification sweep (2026-08-27, post-polish): 14/14 runnable
archs PASS — comfy-source loads (zimage convrot8, qwen fp8, wan ×2 +
fp8 umt5 TE), all ported holder-config DiTs, and the migrated
quantize_model paths (chroma, flux_kontext, f_light block-streamed) in
one run; mageflow remains the upstream 404 skip. One regression caught
and fixed: qwen's _load_single_file override needed the new config
kwarg.
- [x] Full sweep run 2026-08-27: 14/15 PASS (zimage, qwen_image, krea2,
boogu_image, ernie_image, ideogram4, hidream_o1, anima, wan21, wan22_5b,
chroma, flux_kontext, flux2_klein_4b, ltx2.3 — the quantized 22B ltx
@@ -378,20 +412,28 @@ Decisions:
## TODO / look at later
- [ ] Quantize-path consolidation quirks: `quantize_kwargs` is honored only by the
raw `quantize()` call sites and silently dropped by `quantize_model()`; the
ARA path inside `quantize_model` hardcodes `uint8`. Decide the unified
behavior when consolidating.
- [x] Quantize consolidation: `quantize_model` now honors `quantize_kwargs`
(blocks + extras) and tolerates missing block names; the chroma ×2,
flux_kontext, f_light, omnigen2 raw-quantize sites migrated onto it
(gaining block streaming, excludes, ARA, dequant-on-save patching) with
holder block names added. Remaining raw sites are legacy/extension
(flex2, cogview4, stable_diffusion_model). The ARA uint8 hardcode
stands — revisit if a non-uint8 ARA base is ever wanted.
- [x] Last known `qtype_te` bugs fixed (flux2's Mistral TE, anima's
text_conditioner) — 9/9 sites from the survey now correct outside the
grandfathered legacy monolith (cogview4/legacy SD remain as-is).
- [x] wan comfy-TE resolved for real: UMT5TextEncoder carries comfy
candidates (fp8_e4m3fn_scaled via the legacy importer, fp16), files
already in transformers key layout (spiece blob dropped, tied
embed_tokens materialized). Verified: wan21 samples with the local
comfy fp8 TE. The loaders/umt5.py `comfy_files` param stays as a
no-op shim for old callers.
- [ ] Registry hardening: error (don't fall back to SD1) on unknown arch; lazy
- [x] Registry hardening: unknown archs now raise with the known-arch list
(legacy monolith archs whitelisted via LEGACY_ARCHS). Still open: lazy
per-arch imports; single source of truth shared with the UI's
`options.tsx` model list.
- [ ] Fake/stub components: consolidate on `toolkit/models/FakeVAE.py` /
`toolkit/unloader.py`, delete local copies.
- [x] Stub dedup where identical: chroma_radiance imports FakeCLIP/FakeConfig
from chroma_model. The other Fake* copies (hidream_o1, flux2, zeta)
carry model-specific values — left in place.
- [ ] Vendored upstream code (hidream/src, omnigen2/src, ltx2 converter's private
comfy-quant parser): dedupe against toolkit utils where practical.

View File

@@ -67,6 +67,12 @@ class OstrisModelMixin:
aitk_processor_repo: Optional[str] = None
aitk_processor_subfolder: Optional[str] = None
# single-file loads without quant markers cast tensors to the requested
# dtype; classes whose checkpoints carry a deliberate precision mix (e.g.
# fp32 norms next to bf16 weights) set this False to load at stored
# precision instead
aitk_cast_on_load: bool = True
# ---- state set by the loader / quantizer ----
aitk_is_quantized: bool = False
aitk_qtype: Optional[str] = None
@@ -119,7 +125,10 @@ class OstrisModelMixin:
from accelerate import init_empty_weights
with init_empty_weights(include_buffers=False):
return cls.from_config(config)
if hasattr(cls, "from_config"):
return cls.from_config(config)
# plain nn.Module classes take their params/config object directly
return cls(config)
# ------------------------------------------------------------------
# loading
@@ -135,6 +144,7 @@ class OstrisModelMixin:
quantize_device: Optional[torch.device] = None,
exclude_quant_modules: Optional[List[str]] = None,
config_path: Optional[str] = None,
config=None,
subfolder: Optional[str] = None,
use_comfy_weights: bool = True,
**kwargs,
@@ -180,7 +190,11 @@ class OstrisModelMixin:
if name_or_path.endswith(".safetensors"):
file_path = cls._resolve_single_file(name_or_path)
model = cls._load_single_file(
file_path, dtype=dtype, config_path=config_path, subfolder=subfolder
file_path,
dtype=dtype,
config_path=config_path,
config=config,
subfolder=subfolder,
)
else:
if os.path.isdir(name_or_path):
@@ -243,11 +257,16 @@ class OstrisModelMixin:
file_path: str,
dtype: torch.dtype,
config_path: Optional[str] = None,
config=None,
subfolder: Optional[str] = None,
):
state_dict = load_file(file_path)
return cls.load_from_state_dict(
state_dict, dtype, config_path=config_path, subfolder=subfolder
state_dict,
dtype,
config_path=config_path,
config=config,
subfolder=subfolder,
)
@classmethod
@@ -263,16 +282,20 @@ class OstrisModelMixin:
state_dict: Dict[str, torch.Tensor],
dtype: torch.dtype,
config_path: Optional[str] = None,
config=None,
subfolder: Optional[str] = None,
):
"""Build the model and load an already-read single-file state dict
(the tail of the single-file path; also callable directly for
checkpoints read from non-safetensors sources)."""
checkpoints read from non-safetensors sources). ``config``, when
given, is used directly (for models whose config comes from the
holder, e.g. model_kwargs-driven archs)."""
state_dict = cls.convert_state_dict_on_load(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.aitk_config_from_state_dict(state_dict)
if config is None:
config = cls._load_single_file_config(config_path, subfolder)
model = cls.aitk_from_config(config)
@@ -288,11 +311,15 @@ class OstrisModelMixin:
)
cls._load_state_dict_with_quantized(model, state_dict)
model.aitk_is_quantized = True
else:
elif cls.aitk_cast_on_load:
for key, value in state_dict.items():
state_dict[key] = value.to(dtype=dtype)
if value.is_floating_point():
state_dict[key] = value.to(dtype=dtype)
model.load_state_dict(state_dict, assign=True)
model.to(dtype=dtype)
else:
# stored-precision load (the checkpoint's dtype mix is deliberate)
model.load_state_dict(state_dict, assign=True)
del state_dict
flush()
return model

View File

@@ -29,7 +29,9 @@ class QwenImageTransformer2DModel(
return ["transformer_blocks"]
@classmethod
def _load_single_file(cls, file_path, dtype, config_path=None, subfolder=None):
def _load_single_file(
cls, file_path, dtype, config_path=None, config=None, subfolder=None
):
from safetensors import safe_open
with safe_open(file_path, framework="pt") as f:
@@ -38,7 +40,11 @@ class QwenImageTransformer2DModel(
# 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
file_path,
dtype,
config_path=config_path,
config=config,
subfolder=subfolder,
)
# other single-file checkpoints carry diffusers or original key
# layouts; diffusers' single-file machinery owns that conversion

View File

@@ -84,3 +84,54 @@ class WanTransformer3DModel(DiffusersWanTransformer3DModel, OstrisModelMixin):
)
return convert_wan_transformer_to_diffusers(dict(state_dict))
# inverse of diffusers' rename table, for the base/t2v/i2v variants (no
# vace/animate: their extra mappings collide in reverse). Order matters:
# longest/most-specific first, and the norm2/norm3 swap uses a placeholder.
_SAVE_RENAMES = [
("condition_embedder.time_embedder.linear_1", "time_embedding.0"),
("condition_embedder.time_embedder.linear_2", "time_embedding.2"),
("condition_embedder.text_embedder.linear_1", "text_embedding.0"),
("condition_embedder.text_embedder.linear_2", "text_embedding.2"),
("condition_embedder.time_proj", "time_projection.1"),
("condition_embedder.image_embedder.norm1", "img_emb.proj.0"),
("condition_embedder.image_embedder.ff.net.0.proj", "img_emb.proj.1"),
("condition_embedder.image_embedder.ff.net.2", "img_emb.proj.3"),
("condition_embedder.image_embedder.norm2", "img_emb.proj.4"),
("ffn.net.0.proj", "ffn.0"),
("ffn.net.2", "ffn.2"),
(".norm_added_k.", ".norm_k_img."),
(".add_k_proj.", ".k_img."),
(".add_v_proj.", ".v_img."),
(".to_out.0.", ".o."),
(".to_q.", ".q."),
(".to_k.", ".k."),
(".to_v.", ".v."),
("attn2", "cross_attn"),
("attn1", "self_attn"),
# norm2 <-> norm3 swap back
("norm3", "norm__placeholder"),
("norm2", "norm3"),
("norm__placeholder", "norm2"),
("proj_out", "head.head"),
]
@classmethod
def convert_state_dict_on_save(cls, state_dict):
"""Diffusers layout back to the original/comfy wan key layout
(rename-only, so quantized weight/scale/marker keys ride along)."""
if any(".self_attn." in k or ".cross_attn." in k for k in state_dict):
return state_dict # already original layout
new_sd = {}
for key, value in state_dict.items():
k = key
# scale_shift_table: blocks keep the name as `modulation`, the
# top-level one belongs to the output head
if k == "scale_shift_table":
k = "head.modulation"
elif k.endswith(".scale_shift_table"):
k = k[: -len("scale_shift_table")] + "modulation"
for src, dst in cls._SAVE_RENAMES:
k = k.replace(src, dst)
new_sd[k] = value
return new_sd

View File

@@ -44,6 +44,7 @@ from typing import Any, Callable, Dict, List, Optional, Union
from toolkit.models.wan21.wan_lora_convert import convert_to_diffusers, convert_to_original
from toolkit.util.quantize import quantize_model
from toolkit.models.v2.text_encoders.umt5 import UMT5TextEncoder
from toolkit.metadata import get_meta_for_safetensors
# for generation only?
scheduler_configUniPC = {
@@ -680,17 +681,16 @@ class Wan21(BaseModel):
return False
def save_model(self, output_path, meta, save_dtype):
# only save the unet
transformer: Wan21 = unwrap_model(self.model)
transformer.save_pretrained(
save_directory=os.path.join(output_path, 'transformer'),
safe_serialization=True,
# comfy-format single-file save (original wan key layout)
transformer = unwrap_model(self.model)
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

@@ -41,10 +41,34 @@ def get_all_models() -> List[BaseModel]:
return all_model_classes
# archs the legacy StableDiffusion monolith still serves (see the arch
# normalization in toolkit/config_modules.py)
LEGACY_ARCHS = {
"sd1",
"sd2",
"sd3",
"sdxl",
"pixart",
"pixart_sigma",
"auraflow",
"flux",
"lumina2",
"vega",
"ssd",
}
def get_model_class(config: ModelConfig):
all_models = get_all_models()
for ModelClass in all_models:
if ModelClass.arch == config.arch:
return ModelClass
# default to the legacy model
return StableDiffusion
if config.arch in LEGACY_ARCHS:
return StableDiffusion
# a typo'd or unregistered arch used to silently fall back to SD1; error
# instead (a broken extension import also lands here — its error was
# printed during get_all_models)
known = sorted({m.arch for m in all_models if m.arch} | LEGACY_ARCHS)
raise ValueError(
f"Unknown model arch {config.arch!r}. Known archs: {', '.join(known)}"
)

View File

@@ -429,9 +429,10 @@ def quantize_model(
# quantize model the original way without an accuracy recovery adapter
# move and quantize only certain pieces at a time.
quantization_type = get_qtype(base_model.model_config.qtype)
quantize_kwargs = base_model.model_config.quantize_kwargs or {}
# all_blocks = list(model_to_quantize.transformer_blocks)
all_blocks: List[torch.nn.Module] = []
transformer_block_names = base_model.get_transformer_block_names()
transformer_block_names = base_model.get_transformer_block_names() or []
for name in transformer_block_names:
# name may be a dotted path for models that nest their blocks
# (e.g. hidream_o1's "model.language_model.layers").
@@ -456,7 +457,12 @@ def quantize_model(
block.to(base_model.device_torch, dtype=base_model.torch_dtype, non_blocking=True)
# exclude patterns with a leading wildcard (e.g. "*adaln_proj*")
# also apply inside blocks, where names are block-relative
quantize(block, weights=quantization_type, exclude=exclude_modules)
quantize(
block,
weights=quantization_type,
exclude=exclude_modules,
**quantize_kwargs,
)
freeze(block)
# NOT non_blocking: an async D2H allocates the cpu destination in pinned
# memory, which the caching host allocator keeps forever (with power-of-2
@@ -472,5 +478,10 @@ def quantize_model(
# device without having to move the transformer blocks to the device first
base_model.print_and_status_update(" - quantizing extras")
# model_to_quantize.to(base_model.device_torch, dtype=base_model.torch_dtype)
quantize(model_to_quantize, weights=quantization_type, exclude=exclude_modules)
quantize(
model_to_quantize,
weights=quantization_type,
exclude=exclude_modules,
**quantize_kwargs,
)
freeze(model_to_quantize)