Added legacy paths
This commit is contained in:
@@ -371,6 +371,43 @@ Decisions:
|
||||
unload/load the difference
|
||||
- [ ] Legacy `stable_diffusion_model.py` archs: grandfather or port last
|
||||
|
||||
## 100% mixin coverage (2026-08-28)
|
||||
|
||||
Every component of every non-legacy arch now loads through OstrisModelMixin —
|
||||
DiTs, text encoders, VAEs, vision encoders, connectors/vocoders, and the
|
||||
custom cases that previously bypassed it:
|
||||
|
||||
- New wrapper classes: llama, gemma3 + gemma4, mistral3 (×2), qwen3 base,
|
||||
qwen3-vl base + text-only, wan VAE, diffusers flux2 KL VAE, CLIP vision
|
||||
(first vision encoder), cosmos DiT, anima text conditioner, the full ltx2
|
||||
set (transformer, video/audio VAEs, connectors, vocoders).
|
||||
- Existing-class swaps: omnigen2 (mllm + VAE), boogu (TE + VAE), klein TE,
|
||||
hidream llama, ideogram4 TE, prx TE.
|
||||
- Custom restructures: minimax video/audio VAEs and TE, MageVAE
|
||||
(deferred-load ctor), the shared flux2_kl AutoEncoder (small-decoder sniff
|
||||
as a class hook; flux2 + ideogram4 route through it), ace_step's bundle
|
||||
(per-component class loaders), anima (modular-pipeline load replaced with
|
||||
component-wise v2 loads + update_components), hidream_o1's Qwen3VL DiT,
|
||||
z_image_l2p rebased onto the v2 class, ltx2's converter builds v2 classes.
|
||||
- Mixin: kwargs flow through the single-file chain (component ctor args like
|
||||
MageVAE's sample_posterior).
|
||||
|
||||
Verified with real weights: ltx2.3 + ltx2.5 (full v2 family), anima,
|
||||
ace_step_15 (new harness entry), minimax VAEs, plus the standing harness
|
||||
coverage. The legacy monolith archs joined the system per the inference-engine
|
||||
goal ("send a job with any base_model, reload/unload components on the fly"):
|
||||
six new wrapper classes complete the component vocabulary (UNet2DCondition,
|
||||
SD3/PixArt ×2/AuraFlow/Lumina2 transformers, Gemma2), and
|
||||
`adopt_component` (in-place class swap, OstrisLinear-style) rebinds
|
||||
pipeline-loaded components onto their wrappers at the monolith's single
|
||||
post-load funnel — covering every legacy arch without touching its fifteen
|
||||
load branches, with all pipeline references staying valid. Verified: sd1
|
||||
(UNet/KLVAE/CLIPTextEncoder all mixin instances + generation) and sdxl
|
||||
(dual-CLIP adoption + 1024² generation); harness gained sd1/sdxl entries,
|
||||
a legacy scheduler fallback, and the sampler-name pass-through. Full
|
||||
monolith decomposition (per-arch v2 loading with comfy candidates) remains
|
||||
future work, but every resident component is now poolable by the engine.
|
||||
|
||||
## Testing
|
||||
|
||||
- [x] `testing/test_model_loading.py`: per-arch load + one small sample through
|
||||
@@ -418,8 +455,12 @@ Decisions:
|
||||
flux2_klein_9b, prx_pixel, zeta_chroma, zimage_l2p, both qwen edit
|
||||
archs, f-lite. Fixes found by the run: ltx2.5's fp32 scale_shift
|
||||
tables promoted hidden states into bf16 linears under the diffusers
|
||||
class (ComfyUI casts per-op) — the DiT/connectors now cast to compute
|
||||
dtype after quantized attach (ConvRot storage immune); harness configs
|
||||
class — fixed ComfyUI-style (toolkit/util/mixed_precision.py):
|
||||
per-op input casting hooks on every weighted module + the stored-fp32
|
||||
tensors pinned against parent .to(dtype) casts (device moves pass
|
||||
through), so the tables stay genuinely fp32 at sample time. The
|
||||
mechanism is general — any mixed-precision comfy checkpoint on a
|
||||
diffusers-class arch can use it; harness configs
|
||||
for e1 resolution and zeta/l2p extras_name_or_path corrected to match
|
||||
the UI defaults.
|
||||
- [ ] Each newly migrated model adds its test in the same PR as its migration.
|
||||
|
||||
@@ -42,6 +42,24 @@ class BasicModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMix
|
||||
pass
|
||||
|
||||
|
||||
def adopt_component(module: torch.nn.Module, wrapper_cls: type) -> torch.nn.Module:
|
||||
"""Rebind an already-loaded component instance onto its v2 wrapper class
|
||||
(in-place class swap, like the OstrisLinear conversion). Used when a
|
||||
component arrives from machinery that builds the base class directly
|
||||
(e.g. a diffusers pipeline load), so every resident component is a mixin
|
||||
instance the inference engine can pool and manage."""
|
||||
if isinstance(module, wrapper_cls):
|
||||
return module
|
||||
base = wrapper_cls.__mro__[1]
|
||||
if not isinstance(module, base):
|
||||
raise TypeError(
|
||||
f"cannot adopt {type(module).__name__} into {wrapper_cls.__name__} "
|
||||
f"(expected a {base.__name__})"
|
||||
)
|
||||
module.__class__ = wrapper_cls
|
||||
return module
|
||||
|
||||
|
||||
class OstrisModelMixin:
|
||||
# ---- per-model configuration, override in subclasses ----
|
||||
# subfolder that holds this model inside a diffusers style checkpoint
|
||||
@@ -195,6 +213,7 @@ class OstrisModelMixin:
|
||||
config_path=config_path,
|
||||
config=config,
|
||||
subfolder=subfolder,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
if os.path.isdir(name_or_path):
|
||||
@@ -259,6 +278,7 @@ class OstrisModelMixin:
|
||||
config_path: Optional[str] = None,
|
||||
config=None,
|
||||
subfolder: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
state_dict = load_file(file_path)
|
||||
return cls.load_from_state_dict(
|
||||
@@ -267,6 +287,7 @@ class OstrisModelMixin:
|
||||
config_path=config_path,
|
||||
config=config,
|
||||
subfolder=subfolder,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -284,6 +305,7 @@ class OstrisModelMixin:
|
||||
config_path: Optional[str] = None,
|
||||
config=None,
|
||||
subfolder: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Build the model and load an already-read single-file state dict
|
||||
(the tail of the single-file path; also callable directly for
|
||||
|
||||
13
toolkit/models/v2/diffusion_models/auraflow.py
Normal file
13
toolkit/models/v2/diffusion_models/auraflow.py
Normal file
@@ -0,0 +1,13 @@
|
||||
from diffusers import AuraFlowTransformer2DModel as DiffusersAuraFlowTransformer2DModel
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class AuraFlowTransformer2DModel(
|
||||
DiffusersAuraFlowTransformer2DModel, OstrisModelMixin
|
||||
):
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["joint_transformer_blocks", "single_transformer_blocks"]
|
||||
15
toolkit/models/v2/diffusion_models/cosmos.py
Normal file
15
toolkit/models/v2/diffusion_models/cosmos.py
Normal file
@@ -0,0 +1,15 @@
|
||||
from diffusers.models import (
|
||||
CosmosTransformer3DModel as DiffusersCosmosTransformer3DModel,
|
||||
)
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class CosmosTransformer3DModel(DiffusersCosmosTransformer3DModel, OstrisModelMixin):
|
||||
"""Cosmos video/image DiT (anima)."""
|
||||
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["transformer_blocks"]
|
||||
34
toolkit/models/v2/diffusion_models/ltx2.py
Normal file
34
toolkit/models/v2/diffusion_models/ltx2.py
Normal file
@@ -0,0 +1,34 @@
|
||||
from diffusers.models.transformers import (
|
||||
LTX2VideoTransformer3DModel as DiffusersLTX2VideoTransformer3DModel,
|
||||
)
|
||||
from diffusers.pipelines.ltx2 import (
|
||||
LTX2TextConnectors as DiffusersLTX2TextConnectors,
|
||||
)
|
||||
from diffusers.pipelines.ltx2 import LTX2Vocoder as DiffusersLTX2Vocoder
|
||||
from diffusers.pipelines.ltx2 import (
|
||||
LTX2VocoderWithBWE as DiffusersLTX2VocoderWithBWE,
|
||||
)
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class LTX2VideoTransformer3DModel(
|
||||
DiffusersLTX2VideoTransformer3DModel, OstrisModelMixin
|
||||
):
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["transformer_blocks"]
|
||||
|
||||
|
||||
class LTX2TextConnectors(DiffusersLTX2TextConnectors, OstrisModelMixin):
|
||||
aitk_subfolder = "connectors"
|
||||
|
||||
|
||||
class LTX2Vocoder(DiffusersLTX2Vocoder, OstrisModelMixin):
|
||||
aitk_subfolder = "vocoder"
|
||||
|
||||
|
||||
class LTX2VocoderWithBWE(DiffusersLTX2VocoderWithBWE, OstrisModelMixin):
|
||||
aitk_subfolder = "vocoder"
|
||||
13
toolkit/models/v2/diffusion_models/lumina2.py
Normal file
13
toolkit/models/v2/diffusion_models/lumina2.py
Normal file
@@ -0,0 +1,13 @@
|
||||
from diffusers import Lumina2Transformer2DModel as DiffusersLumina2Transformer2DModel
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class Lumina2Transformer2DModel(
|
||||
DiffusersLumina2Transformer2DModel, OstrisModelMixin
|
||||
):
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["layers"]
|
||||
24
toolkit/models/v2/diffusion_models/pixart.py
Normal file
24
toolkit/models/v2/diffusion_models/pixart.py
Normal file
@@ -0,0 +1,24 @@
|
||||
from diffusers import PixArtTransformer2DModel as DiffusersPixArtTransformer2DModel
|
||||
from diffusers import Transformer2DModel as DiffusersTransformer2DModel
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class PixArtTransformer2DModel(
|
||||
DiffusersPixArtTransformer2DModel, OstrisModelMixin
|
||||
):
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["transformer_blocks"]
|
||||
|
||||
|
||||
class Transformer2DModel(DiffusersTransformer2DModel, OstrisModelMixin):
|
||||
"""The generic diffusers DiT the pixart sigma path loads."""
|
||||
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["transformer_blocks"]
|
||||
11
toolkit/models/v2/diffusion_models/sd3.py
Normal file
11
toolkit/models/v2/diffusion_models/sd3.py
Normal file
@@ -0,0 +1,11 @@
|
||||
from diffusers import SD3Transformer2DModel as DiffusersSD3Transformer2DModel
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class SD3Transformer2DModel(DiffusersSD3Transformer2DModel, OstrisModelMixin):
|
||||
aitk_subfolder = "transformer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["transformer_blocks"]
|
||||
13
toolkit/models/v2/diffusion_models/unet.py
Normal file
13
toolkit/models/v2/diffusion_models/unet.py
Normal file
@@ -0,0 +1,13 @@
|
||||
from diffusers import UNet2DConditionModel as DiffusersUNet2DConditionModel
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class UNet2DConditionModel(DiffusersUNet2DConditionModel, OstrisModelMixin):
|
||||
"""The SD1/SD2/SDXL-family UNet."""
|
||||
|
||||
aitk_subfolder = "unet"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["down_blocks", "up_blocks"]
|
||||
9
toolkit/models/v2/text_encoders/anima.py
Normal file
9
toolkit/models/v2/text_encoders/anima.py
Normal file
@@ -0,0 +1,9 @@
|
||||
from diffusers import AnimaTextConditioner as DiffusersAnimaTextConditioner
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class AnimaTextConditioner(DiffusersAnimaTextConditioner, OstrisModelMixin):
|
||||
"""Anima's learned text conditioner (rides next to the Qwen3 encoder)."""
|
||||
|
||||
aitk_subfolder = "text_conditioner"
|
||||
14
toolkit/models/v2/text_encoders/gemma2.py
Normal file
14
toolkit/models/v2/text_encoders/gemma2.py
Normal file
@@ -0,0 +1,14 @@
|
||||
from transformers import Gemma2Model
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
|
||||
class Gemma2ModelEncoder(Gemma2Model, OstrisTransformersMixin):
|
||||
"""Gemma2 base model (lumina2's text encoder — what AutoModel resolves)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["layers"]
|
||||
32
toolkit/models/v2/text_encoders/gemma3.py
Normal file
32
toolkit/models/v2/text_encoders/gemma3.py
Normal file
@@ -0,0 +1,32 @@
|
||||
from transformers import Gemma3ForConditionalGeneration
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
|
||||
class Gemma3TextEncoder(Gemma3ForConditionalGeneration, OstrisTransformersMixin):
|
||||
"""Gemma3 conditioning stack (ltx2 / ltx2.3)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
# both layouts seen across transformers versions; missing paths skip
|
||||
return ["model.language_model.layers", "language_model.model.layers"]
|
||||
|
||||
|
||||
try:
|
||||
from transformers.models.gemma4.modeling_gemma4 import Gemma4TextModel
|
||||
|
||||
class Gemma4TextEncoder(Gemma4TextModel, OstrisTransformersMixin):
|
||||
"""Gemma4 text decoder (ltx2.5's conditioning stack)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["layers"]
|
||||
|
||||
except ImportError:
|
||||
Gemma4TextEncoder = None
|
||||
14
toolkit/models/v2/text_encoders/llama.py
Normal file
14
toolkit/models/v2/text_encoders/llama.py
Normal file
@@ -0,0 +1,14 @@
|
||||
from transformers import LlamaForCausalLM
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
|
||||
class LlamaTextEncoder(LlamaForCausalLM, OstrisTransformersMixin):
|
||||
"""Llama causal-LM text encoder (hidream's text_encoder_4)."""
|
||||
|
||||
aitk_subfolder = "text_encoder_4"
|
||||
aitk_tokenizer_subfolder = "tokenizer_4"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["model.layers"]
|
||||
25
toolkit/models/v2/text_encoders/mistral3.py
Normal file
25
toolkit/models/v2/text_encoders/mistral3.py
Normal file
@@ -0,0 +1,25 @@
|
||||
from transformers import Mistral3ForConditionalGeneration, Mistral3Model
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
|
||||
class Mistral3TextEncoder(Mistral3ForConditionalGeneration, OstrisTransformersMixin):
|
||||
"""Mistral3 conditioning stack (flux2)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["model.language_model.layers", "language_model.model.layers"]
|
||||
|
||||
|
||||
class Mistral3ModelEncoder(Mistral3Model, OstrisTransformersMixin):
|
||||
"""The inner Mistral3 base model (ernie_image's text encoder)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["language_model.layers"]
|
||||
@@ -1,4 +1,4 @@
|
||||
from transformers import Qwen3ForCausalLM
|
||||
from transformers import Qwen3ForCausalLM, Qwen3Model
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
@@ -15,3 +15,14 @@ class Qwen3TextEncoder(Qwen3ForCausalLM, OstrisTransformersMixin):
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["model.layers"]
|
||||
|
||||
|
||||
class Qwen3ModelEncoder(Qwen3Model, OstrisTransformersMixin):
|
||||
"""The inner Qwen3 base model (anima's text encoder)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["layers"]
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers import Qwen3VLForConditionalGeneration
|
||||
from transformers import (
|
||||
Qwen3VLForConditionalGeneration,
|
||||
Qwen3VLModel,
|
||||
Qwen3VLTextModel,
|
||||
)
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
@@ -51,3 +55,35 @@ class Qwen3VLTextEncoder(Qwen3VLForConditionalGeneration, OstrisTransformersMixi
|
||||
"""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)
|
||||
|
||||
|
||||
class Qwen3VLModelEncoder(Qwen3VLModel, OstrisTransformersMixin):
|
||||
"""The inner Qwen3-VL base model (what AutoModel resolves): its
|
||||
last_hidden_state is the instruction feature boogu_image / ideogram4
|
||||
consume."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["language_model.layers"]
|
||||
|
||||
def drop_vision_tower(self):
|
||||
if getattr(self, "visual", None) is not None:
|
||||
self.visual = None
|
||||
return self
|
||||
|
||||
def patch_vision_patch_embed(self) -> int:
|
||||
return patch_qwen_vl_patch_embed(self)
|
||||
|
||||
|
||||
class Qwen3VLTextOnlyEncoder(Qwen3VLTextModel, OstrisTransformersMixin):
|
||||
"""The Qwen3-VL text tower alone (prx_pixel)."""
|
||||
|
||||
aitk_subfolder = "text_encoder"
|
||||
aitk_tokenizer_subfolder = "tokenizer"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["layers"]
|
||||
|
||||
10
toolkit/models/v2/vae/autoencoder_kl_flux2.py
Normal file
10
toolkit/models/v2/vae/autoencoder_kl_flux2.py
Normal file
@@ -0,0 +1,10 @@
|
||||
from diffusers import AutoencoderKLFlux2
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class Flux2KLVAE(AutoencoderKLFlux2, OstrisModelMixin):
|
||||
"""The diffusers Flux2 KL VAE (ernie_image; distinct from the hand-rolled
|
||||
BFL AutoEncoder in vae/flux2_kl.py that flux2/ideogram4 use)."""
|
||||
|
||||
aitk_subfolder = "vae"
|
||||
@@ -20,6 +20,8 @@ import torch.utils.checkpoint as ckpt
|
||||
from einops import rearrange
|
||||
from torch import Tensor, nn
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
@dataclass
|
||||
class AutoEncoderParams:
|
||||
@@ -354,7 +356,14 @@ class Decoder(nn.Module):
|
||||
return h
|
||||
|
||||
|
||||
class AutoEncoder(nn.Module):
|
||||
class AutoEncoder(nn.Module, OstrisModelMixin):
|
||||
@classmethod
|
||||
def aitk_config_from_state_dict(cls, state_dict):
|
||||
# small decoder builds report 96ch at the first decoder block
|
||||
if state_dict["decoder.up.0.block.0.conv1.bias"].shape[0] == 96:
|
||||
return AutoEncoderSmallDecoderParams()
|
||||
return AutoEncoderParams()
|
||||
|
||||
def __init__(self, params: AutoEncoderParams):
|
||||
super().__init__()
|
||||
self.params = params
|
||||
|
||||
16
toolkit/models/v2/vae/ltx2.py
Normal file
16
toolkit/models/v2/vae/ltx2.py
Normal file
@@ -0,0 +1,16 @@
|
||||
from diffusers.models.autoencoders import (
|
||||
AutoencoderKLLTX2Audio as DiffusersAutoencoderKLLTX2Audio,
|
||||
)
|
||||
from diffusers.models.autoencoders import (
|
||||
AutoencoderKLLTX2Video as DiffusersAutoencoderKLLTX2Video,
|
||||
)
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class LTX2VideoVAE(DiffusersAutoencoderKLLTX2Video, OstrisModelMixin):
|
||||
aitk_subfolder = "vae"
|
||||
|
||||
|
||||
class LTX2AudioVAE(DiffusersAutoencoderKLLTX2Audio, OstrisModelMixin):
|
||||
aitk_subfolder = "audio_vae"
|
||||
9
toolkit/models/v2/vae/wan.py
Normal file
9
toolkit/models/v2/vae/wan.py
Normal file
@@ -0,0 +1,9 @@
|
||||
from diffusers import AutoencoderKLWan
|
||||
|
||||
from .._mixin import OstrisModelMixin
|
||||
|
||||
|
||||
class WanVAE(AutoencoderKLWan, OstrisModelMixin):
|
||||
"""The wan 2.1/2.2 causal video VAE."""
|
||||
|
||||
aitk_subfolder = "vae"
|
||||
13
toolkit/models/v2/vision_encoders/clip_vision.py
Normal file
13
toolkit/models/v2/vision_encoders/clip_vision.py
Normal file
@@ -0,0 +1,13 @@
|
||||
from transformers import CLIPVisionModel
|
||||
|
||||
from .._mixin import OstrisTransformersMixin
|
||||
|
||||
|
||||
class CLIPVisionEncoder(CLIPVisionModel, OstrisTransformersMixin):
|
||||
"""CLIP vision tower (wan21 i2v image conditioning)."""
|
||||
|
||||
aitk_subfolder = "image_encoder"
|
||||
|
||||
@classmethod
|
||||
def get_transformer_block_names(cls):
|
||||
return ["vision_model.encoder.layers"]
|
||||
@@ -44,6 +44,7 @@ from typing import Any, Callable, Dict, List, Optional, Union
|
||||
from toolkit.models.wan21.wan_lora_convert import convert_to_diffusers, convert_to_original
|
||||
from toolkit.util.quantize import quantize_model
|
||||
from toolkit.models.v2.text_encoders.umt5 import UMT5TextEncoder
|
||||
from toolkit.models.v2.vae.wan import WanVAE
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
|
||||
# for generation only?
|
||||
@@ -431,11 +432,9 @@ class Wan21(BaseModel):
|
||||
|
||||
if self._wan_vae_path is not None:
|
||||
# load the vae from individual repo
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
self._wan_vae_path, torch_dtype=dtype).to(dtype=dtype)
|
||||
vae = WanVAE.load_model(self._wan_vae_path, dtype=dtype, subfolder="")
|
||||
else:
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
vae_path, subfolder="vae", torch_dtype=dtype).to(dtype=dtype)
|
||||
vae = WanVAE.load_model(vae_path, dtype=dtype)
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# WIP, coming soon ish
|
||||
from toolkit.models.v2.vision_encoders.clip_vision import CLIPVisionEncoder
|
||||
from functools import partial
|
||||
import torch
|
||||
import yaml
|
||||
@@ -314,10 +315,8 @@ class Wan21I2V(Wan21):
|
||||
self.model_config.extras_name_or_path ,
|
||||
subfolder="image_processor"
|
||||
)
|
||||
self.image_encoder = CLIPVisionModel.from_pretrained(
|
||||
self.model_config.extras_name_or_path,
|
||||
subfolder="image_encoder",
|
||||
torch_dtype=dtype,
|
||||
self.image_encoder = CLIPVisionEncoder.load_model(
|
||||
self.model_config.extras_name_or_path, dtype=dtype
|
||||
)
|
||||
except Exception as e:
|
||||
# load from name_or_path
|
||||
@@ -325,10 +324,8 @@ class Wan21I2V(Wan21):
|
||||
self.model_config.name_or_path_original,
|
||||
subfolder="image_processor"
|
||||
)
|
||||
self.image_encoder = CLIPVisionModel.from_pretrained(
|
||||
self.model_config.name_or_path_original,
|
||||
subfolder="image_encoder",
|
||||
torch_dtype=dtype,
|
||||
self.image_encoder = CLIPVisionEncoder.load_model(
|
||||
self.model_config.name_or_path_original, dtype=dtype
|
||||
)
|
||||
self.image_encoder.to(self.device_torch, dtype=dtype)
|
||||
self.image_encoder.eval()
|
||||
|
||||
@@ -101,6 +101,46 @@ DO_NOT_TRAIN_WEIGHTS = [
|
||||
DeviceStatePreset = Literal['cache_latents', 'generate']
|
||||
|
||||
|
||||
# diffusers-class name -> v2 wrapper for the legacy archs' components. The
|
||||
# adoption is an in-place class swap, so pipeline-held references stay valid.
|
||||
_V2_ADOPTION_MAP = {
|
||||
"UNet2DConditionModel": ("toolkit.models.v2.diffusion_models.unet", "UNet2DConditionModel"),
|
||||
"AutoencoderKL": ("toolkit.models.v2.vae.autoencoder_kl", "KLVAE"),
|
||||
"CLIPTextModel": ("toolkit.models.v2.text_encoders.clip", "CLIPTextEncoder"),
|
||||
"CLIPTextModelWithProjection": ("toolkit.models.v2.text_encoders.clip", "CLIPTextEncoderWithProjection"),
|
||||
"T5EncoderModel": ("toolkit.models.v2.text_encoders.t5", "T5TextEncoder"),
|
||||
"UMT5EncoderModel": ("toolkit.models.v2.text_encoders.umt5", "UMT5TextEncoder"),
|
||||
"SD3Transformer2DModel": ("toolkit.models.v2.diffusion_models.sd3", "SD3Transformer2DModel"),
|
||||
"PixArtTransformer2DModel": ("toolkit.models.v2.diffusion_models.pixart", "PixArtTransformer2DModel"),
|
||||
"Transformer2DModel": ("toolkit.models.v2.diffusion_models.pixart", "Transformer2DModel"),
|
||||
"AuraFlowTransformer2DModel": ("toolkit.models.v2.diffusion_models.auraflow", "AuraFlowTransformer2DModel"),
|
||||
"FluxTransformer2DModel": ("toolkit.models.v2.diffusion_models.flux", "FluxTransformer2DModel"),
|
||||
"Lumina2Transformer2DModel": ("toolkit.models.v2.diffusion_models.lumina2", "Lumina2Transformer2DModel"),
|
||||
"Gemma2Model": ("toolkit.models.v2.text_encoders.gemma2", "Gemma2ModelEncoder"),
|
||||
}
|
||||
|
||||
|
||||
def _adopt_v2(module):
|
||||
"""Rebind a loaded legacy component onto its v2 wrapper class so every
|
||||
resident component is an OstrisModelMixin instance the inference engine
|
||||
can pool and hot-swap. No-op for unknown or already-adopted classes."""
|
||||
import importlib
|
||||
|
||||
from toolkit.models.v2._mixin import OstrisModelMixin, adopt_component
|
||||
|
||||
if module is None or isinstance(module, OstrisModelMixin):
|
||||
return module
|
||||
entry = _V2_ADOPTION_MAP.get(type(module).__name__)
|
||||
if entry is None:
|
||||
return module
|
||||
try:
|
||||
wrapper = getattr(importlib.import_module(entry[0]), entry[1])
|
||||
return adopt_component(module, wrapper)
|
||||
except (ImportError, TypeError):
|
||||
return module
|
||||
|
||||
|
||||
|
||||
class BlankNetwork:
|
||||
|
||||
def __init__(self):
|
||||
@@ -1048,6 +1088,16 @@ class StableDiffusion:
|
||||
self.unet.requires_grad_(False)
|
||||
self.unet.eval()
|
||||
|
||||
# every resident component joins the v2 mixin system (in-place class
|
||||
# adoption for components the pipeline loaders built directly)
|
||||
_adopt_v2(self.unet)
|
||||
_adopt_v2(self.vae)
|
||||
if isinstance(text_encoder, list):
|
||||
for te in text_encoder:
|
||||
_adopt_v2(te)
|
||||
elif text_encoder is not None:
|
||||
_adopt_v2(text_encoder)
|
||||
|
||||
# load any loras we have
|
||||
if self.model_config.lora_path is not None and not self.is_flux and not self.is_lumina2:
|
||||
pipe.load_lora_weights(self.model_config.lora_path, adapter_name="lora1")
|
||||
|
||||
93
toolkit/util/mixed_precision.py
Normal file
93
toolkit/util/mixed_precision.py
Normal file
@@ -0,0 +1,93 @@
|
||||
"""ComfyUI-style handling for mixed-precision checkpoints (e.g. fp32
|
||||
scale_shift tables next to bf16 linears in one file).
|
||||
|
||||
Two pieces, used together on a loaded module tree:
|
||||
|
||||
- attach_per_op_casting(root): every weighted module casts its floating-point
|
||||
inputs to its own weight dtype at forward (what comfy's manual-cast ops do),
|
||||
so activations promoted to fp32 by a stored-fp32 tensor drop back to the
|
||||
layer's dtype at the next op instead of erroring in a bf16 matmul.
|
||||
- pin_stored_fp32(root): tensors that are fp32 after load stay fp32 through
|
||||
any parent-level dtype cast (`module.to(dtype=...)`, `.half()`, ...) —
|
||||
device moves still apply. Without this, a holder's blanket `.to(dtype)`
|
||||
would silently downcast the deliberately-fp32 pieces.
|
||||
"""
|
||||
|
||||
import types
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from toolkit.util.ostris_quant import OstrisLinear
|
||||
|
||||
|
||||
def _cast_args(args, kwargs, dtype):
|
||||
def cast(t):
|
||||
if torch.is_tensor(t) and t.is_floating_point() and t.dtype != dtype:
|
||||
return t.to(dtype)
|
||||
return t
|
||||
|
||||
return tuple(cast(a) for a in args), {k: cast(v) for k, v in kwargs.items()}
|
||||
|
||||
|
||||
def attach_per_op_casting(root: nn.Module) -> int:
|
||||
"""Register forward-pre-hooks casting inputs to each weighted module's own
|
||||
weight dtype. Covers Linear/Conv/Norm-style modules that own a floating
|
||||
``weight`` parameter, and OstrisLinear via its stored orig dtype (its
|
||||
``weight`` property would materialize the dequantized tensor). Returns the
|
||||
number of modules hooked."""
|
||||
hooked = 0
|
||||
for module in root.modules():
|
||||
if isinstance(module, OstrisLinear):
|
||||
def hook(mod, args, kwargs):
|
||||
return _cast_args(args, kwargs, mod.ostris_orig_dtype)
|
||||
|
||||
else:
|
||||
weight = module._parameters.get("weight", None)
|
||||
if weight is None or not weight.is_floating_point():
|
||||
continue
|
||||
|
||||
def hook(mod, args, kwargs):
|
||||
return _cast_args(args, kwargs, mod._parameters["weight"].dtype)
|
||||
|
||||
module.register_forward_pre_hook(hook, with_kwargs=True)
|
||||
hooked += 1
|
||||
return hooked
|
||||
|
||||
|
||||
def _tag_fp32(root: nn.Module):
|
||||
for t in list(root.parameters()) + list(root.buffers()):
|
||||
if t.is_floating_point() and t.dtype == torch.float32:
|
||||
t._pin_dtype = True
|
||||
if isinstance(t, nn.Parameter):
|
||||
t.data._pin_dtype = True
|
||||
|
||||
|
||||
def pin_stored_fp32(root: nn.Module):
|
||||
"""Mark every currently-fp32 param/buffer in ``root`` and wrap the root's
|
||||
``_apply`` so parent dtype casts skip them (device moves still apply).
|
||||
The wrapper probes the cast function with a scalar to detect whether it
|
||||
changes float dtypes, so ``.to(device)`` passes through untouched."""
|
||||
_tag_fp32(root)
|
||||
orig_apply = root._apply
|
||||
|
||||
def _apply(self, fn, *args, **kwargs):
|
||||
probe = fn(torch.zeros((), dtype=torch.float32))
|
||||
if probe.dtype == torch.float32:
|
||||
return orig_apply(fn, *args, **kwargs)
|
||||
|
||||
def fn_pinned(t):
|
||||
if getattr(t, "_pin_dtype", False):
|
||||
out = t if t.device == probe.device else t.to(probe.device)
|
||||
out._pin_dtype = True
|
||||
return out
|
||||
return fn(t)
|
||||
|
||||
result = orig_apply(fn_pinned, *args, **kwargs)
|
||||
# _apply may rewrap tensors (dropping attributes); the invariant is
|
||||
# simple — the pinned set IS the fp32 set — so re-tag
|
||||
_tag_fp32(root)
|
||||
return result
|
||||
|
||||
root._apply = types.MethodType(_apply, root)
|
||||
return root
|
||||
Reference in New Issue
Block a user