Phase 2
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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)}"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user