Phase 1
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()`
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
14
toolkit/models/v2/diffusion_models/flux.py
Normal file
14
toolkit/models/v2/diffusion_models/flux.py
Normal 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"]
|
||||
18
toolkit/models/v2/diffusion_models/hidream.py
Normal file
18
toolkit/models/v2/diffusion_models/hidream.py
Normal 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"]
|
||||
15
toolkit/models/v2/diffusion_models/nucleus_image.py
Normal file
15
toolkit/models/v2/diffusion_models/nucleus_image.py
Normal 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"]
|
||||
29
toolkit/models/v2/diffusion_models/qwen_image.py
Normal file
29
toolkit/models/v2/diffusion_models/qwen_image.py
Normal 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
|
||||
13
toolkit/models/v2/diffusion_models/wan.py
Normal file
13
toolkit/models/v2/diffusion_models/wan.py
Normal 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"]
|
||||
27
toolkit/models/v2/text_encoders/clip.py
Normal file
27
toolkit/models/v2/text_encoders/clip.py
Normal 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"]
|
||||
22
toolkit/models/v2/text_encoders/qwen25_vl.py
Normal file
22
toolkit/models/v2/text_encoders/qwen25_vl.py
Normal 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
|
||||
17
toolkit/models/v2/text_encoders/qwen3.py
Normal file
17
toolkit/models/v2/text_encoders/qwen3.py
Normal 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"]
|
||||
53
toolkit/models/v2/text_encoders/qwen3_vl.py
Normal file
53
toolkit/models/v2/text_encoders/qwen3_vl.py
Normal 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)
|
||||
16
toolkit/models/v2/text_encoders/t5.py
Normal file
16
toolkit/models/v2/text_encoders/t5.py
Normal 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"]
|
||||
49
toolkit/models/v2/text_encoders/umt5.py
Normal file
49
toolkit/models/v2/text_encoders/umt5.py
Normal 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
|
||||
)
|
||||
10
toolkit/models/v2/vae/autoencoder_kl.py
Normal file
10
toolkit/models/v2/vae/autoencoder_kl.py
Normal 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"
|
||||
539
toolkit/models/v2/vae/flux2_kl.py
Normal file
539
toolkit/models/v2/vae/flux2_kl.py
Normal 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
|
||||
86
toolkit/models/v2/vae/qwen_image.py
Normal file
86
toolkit/models/v2/vae/qwen_image.py
Normal 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)
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user