This commit is contained in:
Jaret Burkett
2026-08-27 11:53:08 -06:00
parent e8d9cf6d35
commit 8db198ec0a
51 changed files with 985 additions and 1273 deletions

View File

@@ -1627,7 +1627,39 @@ class BaseModel:
encoder.to(*args, **kwargs)
else:
self.text_encoder.to(*args, **kwargs)
def prepare_text_encoder(self, text_encoder, dtype=None):
"""Standard post-load text-encoder policy: layer offloading, device
placement, then quantize_te. Skips quantization when the checkpoint
loaded pre-quantized."""
from optimum.quanto import freeze
from toolkit.memory_management import MemoryManager
from toolkit.util.quantize import get_qtype, quantize
dtype = dtype if dtype is not None else self.torch_dtype
if (
self.model_config.layer_offloading
and self.model_config.layer_offloading_text_encoder_percent > 0
):
MemoryManager.attach(
text_encoder,
self.device_torch,
offload_percent=self.model_config.layer_offloading_text_encoder_percent,
)
text_encoder.to(self.device_torch, dtype=dtype)
flush()
if self.model_config.quantize_te and not getattr(
text_encoder, "aitk_is_quantized", False
):
self.print_and_status_update("Quantizing Text Encoder")
quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te))
freeze(text_encoder)
flush()
return text_encoder
def convert_lora_weights_before_save(self, state_dict):
# can be overridden in child classes to convert weights before saving
if self.lora_keys_use_comfy_prefix:

View File

@@ -1,46 +1,29 @@
from typing import List
import torch
from transformers import T5Tokenizer, UMT5EncoderModel
class PatchedT5Tokenizer(T5Tokenizer):
def __init__(
self,
vocab: str | list[tuple[str, float]] | None = None,
eos_token="</s>",
unk_token="<unk>",
pad_token="<pad>",
_spm_precompiled_charsmap=None,
extra_ids=100,
additional_special_tokens=None,
**kwargs,
):
super().__init__(
vocab=vocab,
eos_token=eos_token,
unk_token=unk_token,
pad_token=pad_token,
_spm_precompiled_charsmap=None, # this is passing a empty byte string for some reason now.
extra_ids=extra_ids,
additional_special_tokens=additional_special_tokens,
**kwargs,
)
import torch
from toolkit.models.v2.text_encoders.umt5 import (
PatchedT5Tokenizer,
UMT5TextEncoder,
)
def get_umt5_encoder(
model_path: str,
tokenizer_subfolder: str = None,
encoder_subfolder: str = None,
torch_dtype: str = torch.bfloat16,
comfy_files: List[str] = [
"text_encoders/umt5_xxl_fp16.safetensors",
"text_encoders/umt5_xxl_fp8_e4m3fn_scaled.safetensors",
],
) -> UMT5EncoderModel:
"""
Load the UMT5 encoder model from the specified path.
"""
tokenizer = PatchedT5Tokenizer.from_pretrained(model_path, subfolder=tokenizer_subfolder)
# reserved for the comfy-weights flip (Phase 2); accepted for
# signature compatibility, not consulted yet
comfy_files: List[str] = None,
):
"""Load the UMT5 tokenizer + encoder. Thin compatibility wrapper around
toolkit/models/v2/text_encoders/umt5.py."""
tokenizer = UMT5TextEncoder.load_tokenizer(
model_path, subfolder=tokenizer_subfolder or ""
)
print(f"Using {model_path} for UMT5 encoder.")
text_encoder = UMT5EncoderModel.from_pretrained(
model_path, subfolder=encoder_subfolder, torch_dtype=torch_dtype
text_encoder = UMT5TextEncoder.load_model(
model_path, dtype=torch_dtype, subfolder=encoder_subfolder or ""
)
return tokenizer, text_encoder

View File

@@ -166,21 +166,88 @@ Every arch's components become v2 classes; if `name_or_path` is diffusers, it st
loads via diffusers. Nothing about sources or outputs changes yet. Suggested order
(worst duplication first), each including its loading test (see Testing):
- [ ] Shared TEs: `text_encoders/qwen3_vl.py`, `text_encoders/qwen3.py`,
`text_encoders/t5.py`, `text_encoders/clip.py`
- [ ] Shared VAEs: `vae/flux_kl.py` (kills flux2 + ideogram4 copies and the 4
scattered `AutoencoderKL` loads), `vae/qwen_image.py` (mean/std handling
built in)
- [ ] z_image / z_image_l2p (already half-migrated)
- [ ] qwen_image family (qwen_image, qwen_image_edit, qwen_image_edit_plus)
- [ ] nucleus_image, krea2, ideogram4, mageflow
- [ ] chroma, chroma_radiance, zeta_chroma
- [ ] flux2, flux_kontext
- [ ] minimax_h3 (+ ref2va), ltx2 family
- [ ] wan21 / wan22 family
- [ ] hidream family, omnigen2
- [ ] anima, boogu_image, ernie_image, f_light, prx_pixel_t2i
- [ ] audio_models (ace_step)
- [x] `text_encoders/qwen3.py` — Qwen3TextEncoder + `OstrisTransformersMixin`
backend + `BaseModel.prepare_text_encoder` policy helper; the 3 verbatim
TE stanzas (z_image, z_image_l2p, zeta_chroma) replaced. Verified with
real Z-Image weights (load + encode on GPU).
- [x] `text_encoders/qwen3_vl.py` — Qwen3VLTextEncoder with
`drop_vision_tower` / `patch_vision_patch_embed`; the 4 identical
`patch_qwen_vl_patch_embed` copies (krea2, mageflow, boogu_image,
Qwen3VLCaptioner) consolidated; TE loads migrated in krea2, mageflow,
nucleus_image. Still on their own paths: ideogram4 (loads via AutoModel),
minimax_h3 (custom truncated/prequantized comfy load — port later),
qwen_image (Qwen2.5-VL, needs its own class)
- [x] `text_encoders/t5.py`, `text_encoders/clip.py` — T5TextEncoder,
CLIPTextEncoder, CLIPTextEncoderWithProjection; migrated chroma ×2,
flux_kontext, f_light (T5 stanzas → `prepare_text_encoder`, fixing their
`qtype` → `qtype_te` bug) and hidream (CLIP ×2 + T5 with subfolder
overrides; slow-tokenizer classes preserved via `use_fast=False`)
- [x] `vae/qwen_image.py` — QwenImageVAE + QwenImageVAEHolderMixin (frame-dim +
latents mean/std handling built in, tiling opt-in via
`vae_decode_tiled_on_low_vram`); the triplicated encode/decode deleted
from qwen_image, nucleus_image, krea2 and all three VAE loads routed
through the v2 loader
- [x] `vae/autoencoder_kl.py` — KLVAE (diffusers AutoencoderKL through the
universal loader); migrated the scattered loads in chroma, flux_kontext,
f_light, hidream, z_image
- [x] `vae/flux2_kl.py` — the BFL-style Flux2 KL autoencoder unified from the
flux2 + ideogram4 copies (both files deleted; flux2's
encode/decode/small-decoder superset + ideogram4's diffusers key
converter). Verified bit-identical to both originals (weights, encode/
decode outputs, converter mapping) and round-tripped real ae.safetensors
weights on GPU. Packing/normalization stays per-model — flux2 packs
`(c pi pj)` with BatchNorm running stats, ideogram4 packs `(ph pw c)`
with its latent_norm tables; the conventions are incompatible.
- [x] z_image — transformer, TE (qwen3), and VAE (KLVAE) all on v2 modules.
z_image_l2p still has its local progressive-transformer subclass
(rebasing it onto the v2 class deferred; its TE is migrated)
- [x] qwen_image family — `v2/diffusion_models/qwen_image.py` (single-file
loads stay on diffusers' from_single_file until the comfy flip) +
`v2/text_encoders/qwen25_vl.py` (slow tokenizer preserved); edit
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] 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
configs + `load_from_state_dict` for non-safetensors sources, used by
radiance's .pth path). zeta_chroma transformer left as-is: its config
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] wan21 / wan22 family — `v2/diffusion_models/wan.py`
(WanTransformer3DModel, both wan22 dual loads included) +
`v2/text_encoders/umt5.py` (UMT5TextEncoder + PatchedT5Tokenizer;
`loaders/umt5.py` is now a thin compat shim, `comfy_files` still
reserved for Phase 2 — no local comfy umt5 file to verify the key
conversion against). wan21's TE `qtype` → `qtype_te` bug fixed via
prepare_text_encoder
- [x] hidream family — vendored transformer carries the mixin;
`v2/diffusion_models/hidream.py` wraps the diffusers class for
hidream_e1; both load via the switchable `hidream_transformer_class`
through `load_model`
- [x] omnigen2 — vendored transformer carries the mixin, load migrated
- [x] boogu_image, ernie_image, prx_pixel_t2i — their vendored diffusers-style
DiT classes now carry OstrisModelMixin (subfolder + block names on the
class) and the holders load via `load_model`
- [x] f_light — DiT class carries the mixin (`aitk_subfolder="dit_model"`),
load migrated
- [ ] anima — loads through diffusers modular pipelines (AnimaModularPipeline);
not a mixin fit, revisit at Phase 2
- [ ] flux2 DiT — holder-config params classes (Flux2/Klein variants), defer
like krea2/mageflow
- [ ] ace_step — one bundled safetensors holds model+TE+VAE+tokenizer via its
own load_models; decomposing into v2 components is its own task
- [ ] Per-model fixes folded in as each migrates: `qtype_te` bug, dequant-on-save
(`dequantize_if_quantized` everywhere), raw-`quantize()` → `quantize_model()`

View File

@@ -141,9 +141,13 @@ class OstrisModelMixin:
they were.
config_path: config source for single-file loads, overriding aitk_config_repo.
device: move the finished model there before returning.
subfolder: overrides the class's aitk_subfolder; pass "" to explicitly
load from the checkpoint root (e.g. a raw hub repo).
"""
if subfolder is None:
subfolder = cls.aitk_subfolder
elif subfolder == "":
subfolder = None
if name_or_path.endswith(".safetensors"):
file_path = cls._resolve_single_file(name_or_path)
@@ -213,11 +217,34 @@ class OstrisModelMixin:
config_path: Optional[str] = None,
subfolder: Optional[str] = None,
):
config = cls._load_single_file_config(config_path, subfolder)
state_dict = load_file(file_path)
return cls.load_from_state_dict(
state_dict, dtype, config_path=config_path, subfolder=subfolder
)
@classmethod
def aitk_config_from_state_dict(cls, state_dict: Dict[str, torch.Tensor]):
"""Derive the model config from the checkpoint itself (e.g. sniffing
block counts from key indices). Return None (the default) to load the
config from config_path / aitk_config_repo instead."""
return None
@classmethod
def load_from_state_dict(
cls,
state_dict: Dict[str, torch.Tensor],
dtype: torch.dtype,
config_path: Optional[str] = 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)."""
state_dict = cls.convert_state_dict_on_load(state_dict)
has_quant_markers = 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)
model = cls.aitk_from_config(config)
if has_quant_markers:
@@ -290,23 +317,51 @@ class OstrisModelMixin:
# ------------------------------------------------------------------
@classmethod
def load_tokenizer(cls, **kwargs):
def load_tokenizer(
cls,
name_or_path: Optional[str] = None,
subfolder: Optional[str] = None,
**kwargs,
):
"""Load the tokenizer from ``name_or_path`` (a checkpoint dir or repo
holding it at aitk_tokenizer_subfolder), falling back to the class's
aitk_tokenizer_repo. subfolder overrides the class default; "" loads
from the checkpoint root."""
from transformers import AutoTokenizer
if cls.aitk_tokenizer_repo is None:
source = name_or_path if name_or_path is not None else cls.aitk_tokenizer_repo
if source is None:
raise ValueError(f"{cls.__name__} does not declare aitk_tokenizer_repo")
if subfolder is None:
subfolder = cls.aitk_tokenizer_subfolder
if subfolder and os.path.isdir(source) and not os.path.isdir(
os.path.join(source, subfolder)
):
subfolder = None
return AutoTokenizer.from_pretrained(
cls.aitk_tokenizer_repo, subfolder=cls.aitk_tokenizer_subfolder, **kwargs
source, subfolder=subfolder or "", **kwargs
)
@classmethod
def load_processor(cls, **kwargs):
def load_processor(
cls,
name_or_path: Optional[str] = None,
subfolder: Optional[str] = None,
**kwargs,
):
from transformers import AutoProcessor
if cls.aitk_processor_repo is None:
source = name_or_path if name_or_path is not None else cls.aitk_processor_repo
if source is None:
raise ValueError(f"{cls.__name__} does not declare aitk_processor_repo")
if subfolder is None:
subfolder = cls.aitk_processor_subfolder
if subfolder and os.path.isdir(source) and not os.path.isdir(
os.path.join(source, subfolder)
):
subfolder = None
return AutoProcessor.from_pretrained(
cls.aitk_processor_repo, subfolder=cls.aitk_processor_subfolder, **kwargs
source, subfolder=subfolder or "", **kwargs
)
# ------------------------------------------------------------------
@@ -403,3 +458,26 @@ class OstrisModelMixin:
self.aitk_qtype = qtype
flush()
return self
class OstrisTransformersMixin(OstrisModelMixin):
"""OstrisModelMixin with the backend hooks speaking the transformers-lib
API (PreTrainedModel / AutoConfig) instead of diffusers ModelMixin. Base
for text-encoder and vision-encoder modules."""
@classmethod
def aitk_from_pretrained(cls, path, subfolder=None, dtype=None, **kwargs):
return cls.from_pretrained(
path, subfolder=subfolder or "", torch_dtype=dtype, **kwargs
)
@classmethod
def aitk_load_config(cls, path, subfolder=None):
from transformers import AutoConfig
return AutoConfig.from_pretrained(path, subfolder=subfolder or "")
@classmethod
def aitk_from_config(cls, config):
with torch.device("meta"):
return cls(config)

View File

@@ -0,0 +1,14 @@
from diffusers import FluxTransformer2DModel as DiffusersFluxTransformer2DModel
from .._mixin import OstrisModelMixin
class FluxTransformer2DModel(DiffusersFluxTransformer2DModel, OstrisModelMixin):
"""Flux1-family DiT (flux, flux_kontext, chroma-adjacent finetunes in
diffusers layout)."""
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks", "single_transformer_blocks"]

View File

@@ -0,0 +1,18 @@
from diffusers.models import (
HiDreamImageTransformer2DModel as DiffusersHiDreamImageTransformer2DModel,
)
from .._mixin import OstrisModelMixin
class HiDreamImageTransformer2DModel(
DiffusersHiDreamImageTransformer2DModel, OstrisModelMixin
):
"""The diffusers HiDream DiT (hidream_e1; the base hidream arch uses the
vendored copy in the hidream extension)."""
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["double_stream_blocks", "single_stream_blocks"]

View File

@@ -0,0 +1,15 @@
from diffusers import (
NucleusMoEImageTransformer2DModel as DiffusersNucleusMoEImageTransformer2DModel,
)
from .._mixin import OstrisModelMixin
class NucleusMoEImageTransformer2DModel(
DiffusersNucleusMoEImageTransformer2DModel, OstrisModelMixin
):
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks"]

View File

@@ -0,0 +1,29 @@
from diffusers import (
QwenImageTransformer2DModel as DiffusersQwenImageTransformer2DModel,
)
from .._mixin import OstrisModelMixin
class QwenImageTransformer2DModel(
DiffusersQwenImageTransformer2DModel, OstrisModelMixin
):
aitk_subfolder = "transformer"
aitk_config_repo = "Qwen/Qwen-Image"
@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
# layouts; diffusers' single-file machinery owns that conversion
model = cls.from_single_file(
file_path,
config=config_path if config_path is not None else cls.aitk_config_repo,
subfolder="transformer",
torch_dtype=dtype,
)
model.to(dtype)
return model

View File

@@ -0,0 +1,13 @@
from diffusers import WanTransformer3DModel as DiffusersWanTransformer3DModel
from .._mixin import OstrisModelMixin
class WanTransformer3DModel(DiffusersWanTransformer3DModel, OstrisModelMixin):
"""Wan 2.1/2.2 video DiT (wan22 loads two of these into its dual wrapper)."""
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["blocks"]

View File

@@ -0,0 +1,27 @@
from transformers import CLIPTextModel, CLIPTextModelWithProjection
from .._mixin import OstrisTransformersMixin
class CLIPTextEncoder(CLIPTextModel, OstrisTransformersMixin):
"""CLIP-L text encoder (flux-style checkpoints: text_encoder/ +
tokenizer/)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["text_model.encoder.layers"]
class CLIPTextEncoderWithProjection(CLIPTextModelWithProjection, OstrisTransformersMixin):
"""CLIP text encoder with the projection head (SDXL / SD3 / HiDream style
checkpoints; the second encoder lives at text_encoder_2/ + tokenizer_2/)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["text_model.encoder.layers"]

View File

@@ -0,0 +1,22 @@
from transformers import Qwen2_5_VLForConditionalGeneration
from .._mixin import OstrisTransformersMixin
class Qwen25VLTextEncoder(Qwen2_5_VLForConditionalGeneration, OstrisTransformersMixin):
"""Qwen2.5-VL conditioning stack (qwen_image family). Loads from a
checkpoint's text_encoder/ subfolder; the edit variants keep the vision
tower, plain t2i drops it."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["model.language_model.layers"]
def drop_vision_tower(self):
"""Text-only conditioning: the vision tower is dead weight."""
if getattr(self.model, "visual", None) is not None:
self.model.visual = None
return self

View File

@@ -0,0 +1,17 @@
from transformers import Qwen3ForCausalLM
from .._mixin import OstrisTransformersMixin
class Qwen3TextEncoder(Qwen3ForCausalLM, OstrisTransformersMixin):
"""Qwen3 causal-LM text encoder (Z-Image family, Zeta-Chroma, ...). Loads
from a checkpoint's text_encoder/ subfolder, a hub repo, or a single
.safetensors file; the tokenizer rides in the checkpoint's tokenizer/
subfolder."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["model.layers"]

View File

@@ -0,0 +1,53 @@
import torch
import torch.nn.functional as F
from transformers import Qwen3VLForConditionalGeneration
from .._mixin import OstrisTransformersMixin
def patch_qwen_vl_patch_embed(model) -> int:
"""Qwen-VL's vision patch_embed is a Conv3d whose kernel == stride, i.e. a plain
linear projection of each flattened patch. bf16 Conv3d has no fast cuDNN kernel and
falls back to a slow, GPU-underutilizing path. Swap it for the equivalent F.linear
(a GEMM). The weight is read lazily so this survives later .to(device)/dtype moves.
Returns the number of patch_embed modules patched."""
patched = 0
for module in model.modules():
proj = getattr(module, "proj", None)
if isinstance(proj, torch.nn.Conv3d) and tuple(proj.kernel_size) == tuple(
proj.stride
):
def fast_forward(hidden_states, _proj=proj):
w = _proj.weight.reshape(_proj.weight.shape[0], -1)
x = hidden_states.view(-1, w.shape[1]).to(w.dtype)
return F.linear(x, w, _proj.bias)
module.forward = fast_forward
patched += 1
return patched
class Qwen3VLTextEncoder(Qwen3VLForConditionalGeneration, OstrisTransformersMixin):
"""Qwen3-VL conditioning stack (krea2, mageflow, nucleus_image, ideogram4,
minimax_h3, ...). Loads from a checkpoint's text_encoder/ subfolder or a
raw Qwen repo (pass subfolder=\"\" for the latter)."""
aitk_subfolder = "text_encoder"
aitk_processor_subfolder = "processor"
@classmethod
def get_transformer_block_names(cls):
return ["model.language_model.layers"]
def drop_vision_tower(self):
"""Text-only conditioning: the vision tower is dead weight — drop it to
free VRAM and skip loading its (bf16-slow) Conv3d patch_embed."""
if getattr(self.model, "visual", None) is not None:
self.model.visual = None
return self
def patch_vision_patch_embed(self) -> int:
"""Keep the vision tower (reference images ride into the embeddings)
but swap its Conv3d patch_embed for an equivalent GEMM."""
return patch_qwen_vl_patch_embed(self)

View File

@@ -0,0 +1,16 @@
from transformers import T5EncoderModel
from .._mixin import OstrisTransformersMixin
class T5TextEncoder(T5EncoderModel, OstrisTransformersMixin):
"""T5-XXL text encoder. Defaults to the flux-style checkpoint layout
(text_encoder_2/ + tokenizer_2/); pass subfolder overrides for checkpoints
that keep it at text_encoder/ + tokenizer/ (e.g. f-lite)."""
aitk_subfolder = "text_encoder_2"
aitk_tokenizer_subfolder = "tokenizer_2"
@classmethod
def get_transformer_block_names(cls):
return ["encoder.block"]

View File

@@ -0,0 +1,49 @@
import torch
from transformers import T5Tokenizer, UMT5EncoderModel
from .._mixin import OstrisTransformersMixin
class PatchedT5Tokenizer(T5Tokenizer):
def __init__(
self,
vocab=None,
eos_token="</s>",
unk_token="<unk>",
pad_token="<pad>",
_spm_precompiled_charsmap=None,
extra_ids=100,
additional_special_tokens=None,
**kwargs,
):
super().__init__(
vocab=vocab,
eos_token=eos_token,
unk_token=unk_token,
pad_token=pad_token,
_spm_precompiled_charsmap=None, # this is passing a empty byte string for some reason now.
extra_ids=extra_ids,
additional_special_tokens=additional_special_tokens,
**kwargs,
)
class UMT5TextEncoder(UMT5EncoderModel, OstrisTransformersMixin):
"""UMT5-XXL text encoder (wan family)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["encoder.block"]
@classmethod
def load_tokenizer(cls, name_or_path=None, subfolder=None, **kwargs):
# T5's tokenizer needs the _spm_precompiled_charsmap patch
source = name_or_path if name_or_path is not None else cls.aitk_tokenizer_repo
if subfolder is None:
subfolder = cls.aitk_tokenizer_subfolder
return PatchedT5Tokenizer.from_pretrained(
source, subfolder=subfolder or "", **kwargs
)

View File

@@ -0,0 +1,10 @@
from diffusers import AutoencoderKL
from .._mixin import OstrisModelMixin
class KLVAE(AutoencoderKL, OstrisModelMixin):
"""The diffusers AutoencoderKL (SD/SDXL/Flux1/Z-Image image VAEs), loaded
from a checkpoint's vae/ subfolder through the universal loader."""
aitk_subfolder = "vae"

View File

@@ -0,0 +1,539 @@
"""The BFL-style Flux2 KL autoencoder (32ch latents, 2x2 pixel-shuffle
packing to 128ch), shared by the flux2 family and ideogram4.
flux2 loads the raw BFL ae.safetensors layout and uses encode/decode (with the
BatchNorm running-stats latent normalization); ideogram4 loads diffusers-format
checkpoints via convert_diffusers_state_dict and drives encoder/decoder
directly with its own patchify + latent-norm tables. The two models pack the
128 latent channels in different orders — the packing/normalization stays
per-model.
"""
from __future__ import annotations
import math
import re
from dataclasses import dataclass, field
import torch
import torch.utils.checkpoint as ckpt
from einops import rearrange
from torch import Tensor, nn
@dataclass
class AutoEncoderParams:
resolution: int = 256
in_channels: int = 3
ch: int = 128
out_ch: int = 3
ch_mult: list[int] = field(default_factory=lambda: [1, 2, 4, 4])
num_res_blocks: int = 2
z_channels: int = 32
@dataclass
class AutoEncoderSmallDecoderParams:
resolution: int = 256
in_channels: int = 3
ch: int = 128
ch_encoder: int = 96
out_ch: int = 3
ch_mult: list[int] = field(default_factory=lambda: [1, 2, 4, 4])
num_res_blocks: int = 2
z_channels: int = 32
def swish(x: Tensor) -> Tensor:
return x * torch.sigmoid(x)
class AttnBlock(nn.Module):
def __init__(self, in_channels: int):
super().__init__()
self.in_channels = in_channels
self.norm = nn.GroupNorm(
num_groups=32, num_channels=in_channels, eps=1e-6, affine=True
)
self.q = nn.Conv2d(in_channels, in_channels, kernel_size=1)
self.k = nn.Conv2d(in_channels, in_channels, kernel_size=1)
self.v = nn.Conv2d(in_channels, in_channels, kernel_size=1)
self.proj_out = nn.Conv2d(in_channels, in_channels, kernel_size=1)
def attention(self, h_: Tensor) -> Tensor:
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
b, c, h, w = q.shape
q = rearrange(q, "b c h w -> b 1 (h w) c").contiguous()
k = rearrange(k, "b c h w -> b 1 (h w) c").contiguous()
v = rearrange(v, "b c h w -> b 1 (h w) c").contiguous()
h_ = nn.functional.scaled_dot_product_attention(q, k, v)
return rearrange(h_, "b 1 (h w) c -> b c h w", h=h, w=w, c=c, b=b)
def forward(self, x: Tensor) -> Tensor:
return x + self.proj_out(self.attention(x))
class ResnetBlock(nn.Module):
def __init__(self, in_channels: int, out_channels: int):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.norm1 = nn.GroupNorm(
num_groups=32, num_channels=in_channels, eps=1e-6, affine=True
)
self.conv1 = nn.Conv2d(
in_channels, out_channels, kernel_size=3, stride=1, padding=1
)
self.norm2 = nn.GroupNorm(
num_groups=32, num_channels=out_channels, eps=1e-6, affine=True
)
self.conv2 = nn.Conv2d(
out_channels, out_channels, kernel_size=3, stride=1, padding=1
)
if self.in_channels != self.out_channels:
self.nin_shortcut = nn.Conv2d(
in_channels, out_channels, kernel_size=1, stride=1, padding=0
)
def forward(self, x):
h = x
h = self.norm1(h)
h = swish(h)
h = self.conv1(h)
h = self.norm2(h)
h = swish(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
x = self.nin_shortcut(x)
return x + h
class Downsample(nn.Module):
def __init__(self, in_channels: int):
super().__init__()
# no asymmetric padding in torch conv, must do it ourselves
self.conv = nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=2, padding=0
)
def forward(self, x: Tensor):
pad = (0, 1, 0, 1)
x = nn.functional.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
return x
class Upsample(nn.Module):
def __init__(self, in_channels: int):
super().__init__()
self.conv = nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1
)
def forward(self, x: Tensor):
x = nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
x = self.conv(x)
return x
class Encoder(nn.Module):
def __init__(
self,
resolution: int,
in_channels: int,
ch: int,
ch_mult: list[int],
num_res_blocks: int,
z_channels: int,
):
super().__init__()
self.quant_conv = torch.nn.Conv2d(2 * z_channels, 2 * z_channels, 1)
self.ch = ch
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
# downsampling
self.conv_in = nn.Conv2d(
in_channels, self.ch, kernel_size=3, stride=1, padding=1
)
curr_res = resolution
in_ch_mult = (1,) + tuple(ch_mult)
self.in_ch_mult = in_ch_mult
self.down = nn.ModuleList()
block_in = self.ch
for i_level in range(self.num_resolutions):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = ch * in_ch_mult[i_level]
block_out = ch * ch_mult[i_level]
for _ in range(self.num_res_blocks):
block.append(ResnetBlock(in_channels=block_in, out_channels=block_out))
block_in = block_out
down = nn.Module()
down.block = block
down.attn = attn
if i_level != self.num_resolutions - 1:
down.downsample = Downsample(block_in)
curr_res = curr_res // 2
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in, out_channels=block_in)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in, out_channels=block_in)
# end
self.norm_out = nn.GroupNorm(
num_groups=32, num_channels=block_in, eps=1e-6, affine=True
)
self.conv_out = nn.Conv2d(
block_in, 2 * z_channels, kernel_size=3, stride=1, padding=1
)
self.gradient_checkpointing = False
def enable_gradient_checkpointing(self):
self.gradient_checkpointing = True
def forward(self, x: Tensor) -> Tensor:
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_resolutions):
for i_block in range(self.num_res_blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.down[i_level].block[i_block], hs[-1])
if len(self.down[i_level].attn) > 0:
h = ckpt.checkpoint(self.down[i_level].attn[i_block], h)
else:
h = self.down[i_level].block[i_block](hs[-1])
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
hs.append(h)
if i_level != self.num_resolutions - 1:
if torch.is_grad_enabled() and self.gradient_checkpointing:
hs.append(ckpt.checkpoint(self.down[i_level].downsample, hs[-1]))
else:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.mid.block_1, h)
h = ckpt.checkpoint(self.mid.attn_1, h)
h = ckpt.checkpoint(self.mid.block_2, h)
else:
h = self.mid.block_1(h)
h = self.mid.attn_1(h)
h = self.mid.block_2(h)
# end
h = self.norm_out(h)
h = swish(h)
h = self.conv_out(h)
h = self.quant_conv(h)
return h
class Decoder(nn.Module):
def __init__(
self,
ch: int,
out_ch: int,
ch_mult: list[int],
num_res_blocks: int,
in_channels: int,
resolution: int,
z_channels: int,
):
super().__init__()
self.post_quant_conv = torch.nn.Conv2d(z_channels, z_channels, 1)
self.ch = ch
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.ffactor = 2 ** (self.num_resolutions - 1)
# compute in_ch_mult, block_in and curr_res at lowest res
block_in = ch * ch_mult[self.num_resolutions - 1]
curr_res = resolution // 2 ** (self.num_resolutions - 1)
self.z_shape = (1, z_channels, curr_res, curr_res)
# z to block_in
self.conv_in = nn.Conv2d(
z_channels, block_in, kernel_size=3, stride=1, padding=1
)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in, out_channels=block_in)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in, out_channels=block_in)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch * ch_mult[i_level]
for _ in range(self.num_res_blocks + 1):
block.append(ResnetBlock(in_channels=block_in, out_channels=block_out))
block_in = block_out
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
up.upsample = Upsample(block_in)
curr_res = curr_res * 2
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = nn.GroupNorm(
num_groups=32, num_channels=block_in, eps=1e-6, affine=True
)
self.conv_out = nn.Conv2d(block_in, out_ch, kernel_size=3, stride=1, padding=1)
self.gradient_checkpointing = False
def enable_gradient_checkpointing(self):
self.gradient_checkpointing = True
def forward(self, z: Tensor) -> Tensor:
z = self.post_quant_conv(z)
# get dtype for proper tracing
upscale_dtype = next(self.up.parameters()).dtype
# z to block_in
h = self.conv_in(z)
# middle
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.mid.block_1, h)
h = ckpt.checkpoint(self.mid.attn_1, h)
h = ckpt.checkpoint(self.mid.block_2, h)
else:
h = self.mid.block_1(h)
h = self.mid.attn_1(h)
h = self.mid.block_2(h)
# cast to proper dtype
h = h.to(upscale_dtype)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks + 1):
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.up[i_level].block[i_block], h)
if len(self.up[i_level].attn) > 0:
h = ckpt.checkpoint(self.up[i_level].attn[i_block], h)
else:
h = self.up[i_level].block[i_block](h)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.up[i_level].upsample, h)
else:
h = self.up[i_level].upsample(h)
# end
h = self.norm_out(h)
h = swish(h)
h = self.conv_out(h)
return h
class AutoEncoder(nn.Module):
def __init__(self, params: AutoEncoderParams):
super().__init__()
self.params = params
self.encoder = Encoder(
resolution=params.resolution,
in_channels=params.in_channels,
ch=params.ch,
ch_mult=params.ch_mult,
num_res_blocks=params.num_res_blocks,
z_channels=params.z_channels,
)
decoder_ch = params.ch
if hasattr(params, "ch_encoder"):
decoder_ch = params.ch_encoder
self.decoder = Decoder(
resolution=params.resolution,
in_channels=params.in_channels,
ch=decoder_ch,
out_ch=params.out_ch,
ch_mult=params.ch_mult,
num_res_blocks=params.num_res_blocks,
z_channels=params.z_channels,
)
self.bn_eps = 1e-4
self.bn_momentum = 0.1
self.ps = [2, 2]
self.bn = torch.nn.BatchNorm2d(
math.prod(self.ps) * params.z_channels,
eps=self.bn_eps,
momentum=self.bn_momentum,
affine=False,
track_running_stats=True,
)
self._gradient_checkpointing = False
@property
def gradient_checkpointing(self):
return self._gradient_checkpointing
@gradient_checkpointing.setter
def gradient_checkpointing(self, value: bool):
self._gradient_checkpointing = value
self.encoder.gradient_checkpointing = value
self.decoder.gradient_checkpointing = value
@property
def device(self):
return next(self.parameters()).device
@property
def dtype(self):
return next(self.parameters()).dtype
def enable_gradient_checkpointing(self):
self.gradient_checkpointing = True
self.encoder.enable_gradient_checkpointing()
self.decoder.enable_gradient_checkpointing()
def normalize(self, z):
self.bn.eval()
return self.bn(z)
def inv_normalize(self, z):
self.bn.eval()
s = torch.sqrt(self.bn.running_var.view(1, -1, 1, 1) + self.bn_eps)
m = self.bn.running_mean.view(1, -1, 1, 1)
return z * s + m
def encode(self, x: Tensor) -> Tensor:
moments = self.encoder(x)
mean = torch.chunk(moments, 2, dim=1)[0]
z = rearrange(
mean,
"... c (i pi) (j pj) -> ... (c pi pj) i j",
pi=self.ps[0],
pj=self.ps[1],
)
z = self.normalize(z)
return z
def decode(self, z: Tensor) -> Tensor:
z = self.inv_normalize(z)
z = rearrange(
z,
"... (c pi pj) i j -> ... c (i pi) (j pj)",
pi=self.ps[0],
pj=self.ps[1],
)
dec = self.decoder(z)
return dec
_NUM_RESOLUTIONS = 4
def convert_diffusers_state_dict(src: dict[str, Tensor]) -> dict[str, Tensor]:
out: dict[str, Tensor] = {}
attn_substrings = (".mid.attn_1.",)
for src_key, tensor in src.items():
dst_key = _rewrite_diffusers_key(src_key)
if dst_key is None:
raise KeyError(f"Unrecognized diffusers VAE state-dict key: {src_key}")
if (
any(s in dst_key for s in attn_substrings)
and dst_key.endswith(".weight")
and tensor.ndim == 2
):
tensor = tensor.unsqueeze(-1).unsqueeze(-1)
out[dst_key] = tensor
return out
def _rewrite_diffusers_key(key: str) -> str | None:
if key.startswith("bn."):
return key
if key.startswith("quant_conv."):
return key.replace("quant_conv.", "encoder.quant_conv.", 1)
if key.startswith("post_quant_conv."):
return key.replace("post_quant_conv.", "decoder.post_quant_conv.", 1)
if key == "encoder.conv_norm_out.weight":
return "encoder.norm_out.weight"
if key == "encoder.conv_norm_out.bias":
return "encoder.norm_out.bias"
if key == "decoder.conv_norm_out.weight":
return "decoder.norm_out.weight"
if key == "decoder.conv_norm_out.bias":
return "decoder.norm_out.bias"
m = re.match(r"^(encoder|decoder)\.mid_block\.resnets\.(\d+)\.(.+)$", key)
if m:
side, idx, rest = m.group(1), int(m.group(2)), m.group(3)
rest = rest.replace("conv_shortcut", "nin_shortcut")
return f"{side}.mid.block_{idx + 1}.{rest}"
m = re.match(r"^(encoder|decoder)\.mid_block\.attentions\.0\.(.+)$", key)
if m:
side, rest = m.group(1), m.group(2)
rest = (
rest.replace("group_norm.", "norm.")
.replace("to_q.", "q.")
.replace("to_k.", "k.")
.replace("to_v.", "v.")
.replace("to_out.0.", "proj_out.")
)
return f"{side}.mid.attn_1.{rest}"
m = re.match(r"^encoder\.down_blocks\.(\d+)\.resnets\.(\d+)\.(.+)$", key)
if m:
level, res_idx, rest = m.group(1), m.group(2), m.group(3)
rest = rest.replace("conv_shortcut", "nin_shortcut")
return f"encoder.down.{level}.block.{res_idx}.{rest}"
m = re.match(r"^encoder\.down_blocks\.(\d+)\.downsamplers\.0\.conv\.(.+)$", key)
if m:
return f"encoder.down.{m.group(1)}.downsample.conv.{m.group(2)}"
m = re.match(r"^decoder\.up_blocks\.(\d+)\.resnets\.(\d+)\.(.+)$", key)
if m:
diffusers_idx = int(m.group(1))
res_idx = m.group(2)
rest = m.group(3).replace("conv_shortcut", "nin_shortcut")
return (
f"decoder.up.{_NUM_RESOLUTIONS - 1 - diffusers_idx}.block.{res_idx}.{rest}"
)
m = re.match(r"^decoder\.up_blocks\.(\d+)\.upsamplers\.0\.conv\.(.+)$", key)
if m:
diffusers_idx = int(m.group(1))
return f"decoder.up.{_NUM_RESOLUTIONS - 1 - diffusers_idx}.upsample.conv.{m.group(2)}"
if key.startswith(
(
"encoder.conv_in.",
"encoder.conv_out.",
"decoder.conv_in.",
"decoder.conv_out.",
)
):
return key
return None

View File

@@ -0,0 +1,86 @@
import torch
from diffusers import AutoencoderKLQwenImage
from .._mixin import OstrisModelMixin
class QwenImageVAE(AutoencoderKLQwenImage, OstrisModelMixin):
"""The Qwen-Image (wan-style video) VAE, shared by the qwen_image family,
nucleus_image and krea2."""
aitk_subfolder = "vae"
class QwenImageVAEHolderMixin:
"""BaseModel-side encode_images/decode_latents for models whose self.vae
is the Qwen-Image VAE: it is a video VAE, so images ride in a single-frame
dim and latents are normalized with the config's latents_mean/std."""
# tile the decode when low_vram (decode only; encode stays untiled)
vae_decode_tiled_on_low_vram = False
def encode_images(self, image_list, device=None, dtype=None):
if device is None:
device = self.vae_device_torch
if dtype is None:
dtype = self.vae_torch_dtype
if self.vae.device == torch.device("cpu"):
self.vae.to(device)
self.vae.eval()
self.vae.requires_grad_(False)
image_list = [image.to(device, dtype=dtype) for image in image_list]
images = torch.stack(image_list).to(device, dtype=dtype)
images = images.unsqueeze(2) # add the frame dim
latents = self.vae.encode(images).latent_dist.sample()
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(
1, self.vae.config.z_dim, 1, 1, 1
).to(latents.device, latents.dtype)
latents = (latents - latents_mean) * latents_std
latents = latents.squeeze(2) # drop the frame dim
return latents.to(device, dtype=dtype)
def decode_latents(self, latents: torch.Tensor, device=None, dtype=None):
if device is None:
device = self.vae_device_torch
if dtype is None:
dtype = self.vae_torch_dtype
if self.vae.device == torch.device("cpu"):
self.vae.to(device)
latents = latents.to(device, dtype=dtype)
latents = latents.unsqueeze(2) # add the frame dim
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = (
torch.tensor(self.vae.config.latents_std)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents = latents * latents_std + latents_mean
# full-resolution decode spikes VRAM; models opt in to tiling it
tiled = self.vae_decode_tiled_on_low_vram and self.model_config.low_vram
if tiled:
self.vae.enable_tiling()
try:
images = self.vae.decode(latents).sample
finally:
if tiled:
self.vae.disable_tiling()
images = images.squeeze(2) # drop the frame dim
return images.to(device, dtype=dtype)

View File

@@ -10,7 +10,8 @@ from toolkit.memory_management.manager import MemoryManager
from toolkit.models.base_model import BaseModel
from toolkit.prompt_utils import PromptEmbeds
from transformers import AutoTokenizer, UMT5EncoderModel
from diffusers import WanPipeline, WanTransformer3DModel, AutoencoderKL
from diffusers import WanPipeline, AutoencoderKL
from toolkit.models.v2.diffusion_models.wan import WanTransformer3DModel
from .autoencoder_kl_wan import AutoencoderKLWan
import os
import sys
@@ -29,8 +30,6 @@ import os
import copy
from toolkit.config_modules import ModelConfig, GenerateImageConfig, ModelArch
import torch
from optimum.quanto import freeze, qfloat8, QTensor, qint4
from toolkit.util.quantize import quantize, get_qtype
from diffusers import FlowMatchEulerDiscreteScheduler, UniPCMultistepScheduler
from typing import TYPE_CHECKING, List
from toolkit.accelerator import unwrap_model
@@ -44,7 +43,7 @@ from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
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.loaders.umt5 import get_umt5_encoder
from toolkit.models.v2.text_encoders.umt5 import UMT5TextEncoder
# for generation only?
scheduler_configUniPC = {
@@ -344,11 +343,9 @@ class Wan21(BaseModel):
def load_wan_transformer(self, transformer_path, subfolder=None):
self.print_and_status_update("Loading transformer")
dtype = self.torch_dtype
transformer = WanTransformer3DModel.from_pretrained(
transformer_path,
subfolder=subfolder,
torch_dtype=dtype,
).to(dtype=dtype)
transformer = WanTransformer3DModel.load_model(
transformer_path, dtype=dtype, subfolder=subfolder
)
if self.model_config.split_model_over_gpus:
raise ValueError(
@@ -418,29 +415,9 @@ class Wan21(BaseModel):
self.print_and_status_update("Loading UMT5EncoderModel")
tokenizer, text_encoder = get_umt5_encoder(
model_path=te_path,
tokenizer_subfolder="tokenizer",
encoder_subfolder="text_encoder",
torch_dtype=dtype,
comfy_files=self._comfy_te_file
)
text_encoder.to(self.device_torch, dtype=dtype)
flush()
if self.model_config.quantize_te:
self.print_and_status_update("Quantizing UMT5EncoderModel")
quantize(text_encoder, weights=get_qtype(self.model_config.qtype))
freeze(text_encoder)
flush()
if self.model_config.layer_offloading and self.model_config.layer_offloading_text_encoder_percent > 0:
MemoryManager.attach(
text_encoder,
self.device_torch,
offload_percent=self.model_config.layer_offloading_text_encoder_percent
)
tokenizer = UMT5TextEncoder.load_tokenizer(te_path)
text_encoder = UMT5TextEncoder.load_model(te_path, dtype=dtype)
self.prepare_text_encoder(text_encoder, dtype=dtype)
if self.model_config.low_vram:
print("Moving transformer back to GPU")