diff --git a/extensions_built_in/captioner/Qwen3VLCaptioner.py b/extensions_built_in/captioner/Qwen3VLCaptioner.py index 3c8a5c4..b4ec37e 100644 --- a/extensions_built_in/captioner/Qwen3VLCaptioner.py +++ b/extensions_built_in/captioner/Qwen3VLCaptioner.py @@ -12,6 +12,8 @@ from optimum.quanto import freeze from toolkit.basic import flush from toolkit.util.quantize import quantize, get_qtype +from toolkit.models.v2.text_encoders.qwen3_vl import patch_qwen_vl_patch_embed + from .BaseCaptioner import BaseCaptioner import transformers import logging @@ -19,29 +21,6 @@ import traceback import warnings -def patch_qwen_vl_patch_embed(model): - """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 - - # transformers.logging.set_verbosity_error() warnings.filterwarnings("ignore") logging.disable(logging.WARNING) diff --git a/extensions_built_in/diffusion_models/boogu_image/boogu_image.py b/extensions_built_in/diffusion_models/boogu_image/boogu_image.py index 8241561..5f70298 100644 --- a/extensions_built_in/diffusion_models/boogu_image/boogu_image.py +++ b/extensions_built_in/diffusion_models/boogu_image/boogu_image.py @@ -30,6 +30,7 @@ from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds from toolkit.basic import flush from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.models.base_model import BaseModel +from toolkit.models.v2.text_encoders.qwen3_vl import patch_qwen_vl_patch_embed from toolkit.samplers.custom_flowmatch_sampler import ( CustomFlowMatchEulerDiscreteScheduler, ) @@ -71,33 +72,6 @@ SYSTEM_PROMPT_T2I = ( HF_TOKEN = os.getenv("HF_TOKEN", None) -def patch_qwen_vl_patch_embed(model) -> int: - """Swap Qwen-VL's vision ``patch_embed`` Conv3d for the equivalent ``F.linear``. - - Qwen-VL's patch_embed is a Conv3d whose kernel == stride, i.e. just a linear - projection of each flattened patch. bf16 Conv3d has no fast cuDNN kernel and - falls back to a slow path that effectively locks up image caching for the edit - (TI2I) model. The weight is read lazily so this survives later ``.to()`` moves. - Returns the number of patch_embed modules patched. (Vendored from - extensions_built_in/captioner/Qwen3VLCaptioner.py.) - """ - 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 BooguImageModel(BaseModel): arch = "boogu_image" # Default HF repo when model.name_or_path is unset (overridden by the edit model). @@ -167,8 +141,8 @@ class BooguImageModel(BaseModel): # deserialize -- use the bf16 repo and let ai-toolkit quantize if wanted. self.print_and_status_update("Loading transformer") try: - transformer = BooguImageTransformer2DModel.from_pretrained( - base, subfolder="transformer", torch_dtype=dtype, token=HF_TOKEN + transformer = BooguImageTransformer2DModel.load_model( + base, dtype=dtype, token=HF_TOKEN ) except OSError as e: raise OSError( diff --git a/extensions_built_in/diffusion_models/boogu_image/src/transformer.py b/extensions_built_in/diffusion_models/boogu_image/src/transformer.py index ce93406..090a089 100644 --- a/extensions_built_in/diffusion_models/boogu_image/src/transformer.py +++ b/extensions_built_in/diffusion_models/boogu_image/src/transformer.py @@ -21,6 +21,8 @@ from typing import Any, Dict, List, Optional, Tuple, Union import torch import torch.nn as nn from diffusers.configuration_utils import ConfigMixin, register_to_config + +from toolkit.models.v2._mixin import OstrisModelMixin from diffusers.loaders import PeftAdapterMixin from diffusers.loaders.single_file_model import FromOriginalModelMixin from diffusers.models.attention_processor import Attention @@ -490,10 +492,17 @@ class BooguImageDoubleStreamTransformerBlock(nn.Module): class BooguImageTransformer2DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin + ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, OstrisModelMixin ): """Boogu-Image transformer with mixed double-stream -> single-stream topology.""" + aitk_subfolder = "transformer" + + @classmethod + def get_transformer_block_names(cls): + return ["double_stream_layers", "single_stream_layers"] + + _supports_gradient_checkpointing = True _no_split_modules = [ "BooguImageTransformerBlock", diff --git a/extensions_built_in/diffusion_models/chroma/chroma_model.py b/extensions_built_in/diffusion_models/chroma/chroma_model.py index 15da572..0cd0166 100644 --- a/extensions_built_in/diffusion_models/chroma/chroma_model.py +++ b/extensions_built_in/diffusion_models/chroma/chroma_model.py @@ -5,8 +5,9 @@ import torch from toolkit.config_modules import GenerateImageConfig, ModelConfig from PIL import Image from toolkit.models.base_model import BaseModel +from toolkit.models.v2.text_encoders.t5 import T5TextEncoder +from toolkit.models.v2.vae.autoencoder_kl import KLVAE from toolkit.basic import flush -from diffusers import AutoencoderKL # from toolkit.pixel_shuffle_encoder import AutoencoderPixelMixer from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler @@ -14,7 +15,6 @@ from toolkit.dequantize import patch_dequantization_on_save from toolkit.accelerator import unwrap_model from optimum.quanto import freeze, QTensor from toolkit.util.quantize import quantize, get_qtype -from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer from .pipeline import ChromaPipeline, prepare_latent_image_ids from einops import rearrange, repeat import random @@ -148,37 +148,17 @@ class ChromaModel(BaseModel): self.print_and_status_update("Loading transformer") - chroma_state_dict = load_file(model_path, 'cpu') - - # determine number of double and single blocks - double_blocks = 0 - single_blocks = 0 - for key in chroma_state_dict.keys(): - if "double_blocks" in key: - block_num = int(key.split(".")[1]) + 1 - if block_num > double_blocks: - double_blocks = block_num - elif "single_blocks" in key: - block_num = int(key.split(".")[1]) + 1 - if block_num > single_blocks: - single_blocks = block_num - print(f"Double Blocks: {double_blocks}") - print(f"Single Blocks: {single_blocks}") - - chroma_params.depth = double_blocks - chroma_params.depth_single_blocks = single_blocks - transformer = Chroma(chroma_params) - + if model_path.endswith(".safetensors"): + transformer = Chroma.load_model(model_path, dtype=dtype) + else: + transformer = Chroma.load_from_state_dict(load_file(model_path, "cpu"), dtype) # add dtype, not sure why it doesnt have it transformer.dtype = dtype - # load the state dict into the model - transformer.load_state_dict(chroma_state_dict) - transformer.to(self.quantize_device, dtype=dtype) - + transformer.config = FakeConfig() - transformer.config.num_layers = double_blocks - transformer.config.num_single_layers = single_blocks + transformer.config.num_layers = transformer.params.depth + transformer.config.num_single_layers = transformer.params.depth_single_blocks if self.model_config.quantize: # patch the state dict method @@ -195,21 +175,9 @@ class ChromaModel(BaseModel): flush() self.print_and_status_update("Loading T5") - tokenizer_2 = T5TokenizerFast.from_pretrained( - extras_path, subfolder="tokenizer_2", torch_dtype=dtype - ) - text_encoder_2 = T5EncoderModel.from_pretrained( - extras_path, subfolder="text_encoder_2", torch_dtype=dtype - ) - text_encoder_2.to(self.device_torch, dtype=dtype) - flush() - - if self.model_config.quantize_te: - self.print_and_status_update("Quantizing T5") - quantize(text_encoder_2, weights=get_qtype( - self.model_config.qtype)) - freeze(text_encoder_2) - flush() + tokenizer_2 = T5TextEncoder.load_tokenizer(extras_path) + text_encoder_2 = T5TextEncoder.load_model(extras_path, dtype=dtype) + self.prepare_text_encoder(text_encoder_2, dtype=dtype) # self.print_and_status_update("Loading CLIP") text_encoder = FakeCLIP() @@ -219,12 +187,7 @@ class ChromaModel(BaseModel): self.noise_scheduler = ChromaModel.get_train_scheduler() self.print_and_status_update("Loading VAE") - vae = AutoencoderKL.from_pretrained( - extras_path, - subfolder="vae", - torch_dtype=dtype - ) - vae = vae.to(self.device_torch, dtype=dtype) + vae = KLVAE.load_model(extras_path, dtype=dtype, device=self.device_torch) self.print_and_status_update("Making pipe") diff --git a/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py b/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py index e5a79fe..0a6133b 100644 --- a/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py +++ b/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py @@ -5,8 +5,8 @@ import torch from toolkit.config_modules import GenerateImageConfig, ModelConfig from PIL import Image from toolkit.models.base_model import BaseModel +from toolkit.models.v2.text_encoders.t5 import T5TextEncoder from toolkit.basic import flush -from diffusers import AutoencoderKL # from toolkit.pixel_shuffle_encoder import AutoencoderPixelMixer from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler @@ -14,7 +14,6 @@ from toolkit.dequantize import patch_dequantization_on_save from toolkit.accelerator import unwrap_model from optimum.quanto import freeze, QTensor from toolkit.util.quantize import quantize, get_qtype -from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer from .pipeline import ChromaPipeline, prepare_latent_image_ids from einops import rearrange, repeat import random @@ -152,38 +151,16 @@ class ChromaRadianceModel(BaseModel): if model_path.endswith('.pth') or model_path.endswith('.pt'): chroma_state_dict = torch.load(model_path, map_location='cpu', weights_only=True) + transformer = Chroma.load_from_state_dict(chroma_state_dict, dtype) else: - chroma_state_dict = load_file(model_path, 'cpu') - - # determine number of double and single blocks - double_blocks = 0 - single_blocks = 0 - for key in chroma_state_dict.keys(): - if "double_blocks" in key: - block_num = int(key.split(".")[1]) + 1 - if block_num > double_blocks: - double_blocks = block_num - elif "single_blocks" in key: - block_num = int(key.split(".")[1]) + 1 - if block_num > single_blocks: - single_blocks = block_num - print(f"Double Blocks: {double_blocks}") - print(f"Single Blocks: {single_blocks}") - - chroma_params.depth = double_blocks - chroma_params.depth_single_blocks = single_blocks - transformer = Chroma(chroma_params) - + transformer = Chroma.load_model(model_path, dtype=dtype) # add dtype, not sure why it doesnt have it transformer.dtype = dtype - # load the state dict into the model - transformer.load_state_dict(chroma_state_dict) - transformer.to(self.quantize_device, dtype=dtype) - + transformer.config = FakeConfig() - transformer.config.num_layers = double_blocks - transformer.config.num_single_layers = single_blocks + transformer.config.num_layers = transformer.params.depth + transformer.config.num_single_layers = transformer.params.depth_single_blocks if self.model_config.quantize: # patch the state dict method @@ -200,21 +177,9 @@ class ChromaRadianceModel(BaseModel): flush() self.print_and_status_update("Loading T5") - tokenizer_2 = T5TokenizerFast.from_pretrained( - extras_path, subfolder="tokenizer_2", torch_dtype=dtype - ) - text_encoder_2 = T5EncoderModel.from_pretrained( - extras_path, subfolder="text_encoder_2", torch_dtype=dtype - ) - text_encoder_2.to(self.device_torch, dtype=dtype) - flush() - - if self.model_config.quantize_te: - self.print_and_status_update("Quantizing T5") - quantize(text_encoder_2, weights=get_qtype( - self.model_config.qtype)) - freeze(text_encoder_2) - flush() + tokenizer_2 = T5TextEncoder.load_tokenizer(extras_path) + text_encoder_2 = T5TextEncoder.load_model(extras_path, dtype=dtype) + self.prepare_text_encoder(text_encoder_2, dtype=dtype) # self.print_and_status_update("Loading CLIP") text_encoder = FakeCLIP() diff --git a/extensions_built_in/diffusion_models/chroma/src/model.py b/extensions_built_in/diffusion_models/chroma/src/model.py index ebdf69d..5ef4dcf 100644 --- a/extensions_built_in/diffusion_models/chroma/src/model.py +++ b/extensions_built_in/diffusion_models/chroma/src/model.py @@ -1,4 +1,6 @@ -from dataclasses import dataclass +from dataclasses import dataclass, replace + +from toolkit.models.v2._mixin import OstrisModelMixin import torch from torch import Tensor, nn @@ -86,11 +88,40 @@ def modify_mask_to_attend_padding(mask, max_seq_length, num_extra_padding=8): return modified_mask -class Chroma(nn.Module): +class Chroma(nn.Module, OstrisModelMixin): """ Transformer model for flow matching on sequences. """ + @classmethod + def aitk_config_from_state_dict(cls, state_dict): + # block counts come from the checkpoint's key indices + double_blocks = 0 + single_blocks = 0 + for key in state_dict.keys(): + if "double_blocks" in key: + block_num = int(key.split(".")[1]) + 1 + if block_num > double_blocks: + double_blocks = block_num + elif "single_blocks" in key: + block_num = int(key.split(".")[1]) + 1 + if block_num > single_blocks: + single_blocks = block_num + print(f"Double Blocks: {double_blocks}") + print(f"Single Blocks: {single_blocks}") + return replace( + chroma_params, depth=double_blocks, depth_single_blocks=single_blocks + ) + + @classmethod + def aitk_from_config(cls, config): + with torch.device("meta"): + return cls(config) + + @classmethod + def get_transformer_block_names(cls): + return ["double_blocks", "single_blocks"] + def __init__(self, params: ChromaParams): super().__init__() self.params = params diff --git a/extensions_built_in/diffusion_models/chroma/src/radiance.py b/extensions_built_in/diffusion_models/chroma/src/radiance.py index d328f26..e6f7be7 100644 --- a/extensions_built_in/diffusion_models/chroma/src/radiance.py +++ b/extensions_built_in/diffusion_models/chroma/src/radiance.py @@ -1,4 +1,6 @@ -from dataclasses import dataclass +from dataclasses import dataclass, replace + +from toolkit.models.v2._mixin import OstrisModelMixin import torch from torch import Tensor, nn @@ -100,11 +102,40 @@ def modify_mask_to_attend_padding(mask, max_seq_length, num_extra_padding=8): return modified_mask -class Chroma(nn.Module): +class Chroma(nn.Module, OstrisModelMixin): """ Transformer model for flow matching on sequences. """ + @classmethod + def aitk_config_from_state_dict(cls, state_dict): + # block counts come from the checkpoint's key indices + double_blocks = 0 + single_blocks = 0 + for key in state_dict.keys(): + if "double_blocks" in key: + block_num = int(key.split(".")[1]) + 1 + if block_num > double_blocks: + double_blocks = block_num + elif "single_blocks" in key: + block_num = int(key.split(".")[1]) + 1 + if block_num > single_blocks: + single_blocks = block_num + print(f"Double Blocks: {double_blocks}") + print(f"Single Blocks: {single_blocks}") + return replace( + chroma_params, depth=double_blocks, depth_single_blocks=single_blocks + ) + + @classmethod + def aitk_from_config(cls, config): + with torch.device("meta"): + return cls(config) + + @classmethod + def get_transformer_block_names(cls): + return ["double_blocks", "single_blocks"] + def __init__(self, params: ChromaParams): super().__init__() self.params = params diff --git a/extensions_built_in/diffusion_models/ernie_image/ernie_image.py b/extensions_built_in/diffusion_models/ernie_image/ernie_image.py index bb296f6..c9d76a6 100644 --- a/extensions_built_in/diffusion_models/ernie_image/ernie_image.py +++ b/extensions_built_in/diffusion_models/ernie_image/ernie_image.py @@ -79,20 +79,14 @@ class ErnieImageModel(BaseModel): self.print_and_status_update("Loading transformer") - transformer_path = model_path - transformer_subfolder = "transformer" - if os.path.exists(transformer_path): - transformer_subfolder = None - transformer_path = os.path.join(transformer_path, "transformer") + if os.path.exists(model_path): # check if the path is a full checkpoint. te_folder_path = os.path.join(model_path, "text_encoder") # if we have the te, this folder is a full checkpoint, use it as the base if os.path.exists(te_folder_path): base_model_path = model_path - transformer = ErnieImageTransformer2DModel.from_pretrained( - transformer_path, subfolder=transformer_subfolder, torch_dtype=dtype - ) + transformer = ErnieImageTransformer2DModel.load_model(model_path, dtype=dtype) if self.model_config.quantize: self.print_and_status_update("Quantizing Transformer") diff --git a/extensions_built_in/diffusion_models/ernie_image/transformer.py b/extensions_built_in/diffusion_models/ernie_image/transformer.py index 3d27efe..ca3a2ae 100644 --- a/extensions_built_in/diffusion_models/ernie_image/transformer.py +++ b/extensions_built_in/diffusion_models/ernie_image/transformer.py @@ -33,6 +33,8 @@ from diffusers.models.attention_dispatch import dispatch_attention_fn from diffusers.models.attention_processor import Attention from diffusers.models.embeddings import TimestepEmbedding, Timesteps from diffusers.models.modeling_utils import ModelMixin + +from toolkit.models.v2._mixin import OstrisModelMixin from diffusers.models.normalization import RMSNorm @@ -288,8 +290,14 @@ class ErnieImageAdaLNContinuous(nn.Module): return x -class ErnieImageTransformer2DModel(ModelMixin, ConfigMixin): +class ErnieImageTransformer2DModel(ModelMixin, ConfigMixin, OstrisModelMixin): _supports_gradient_checkpointing = True + aitk_subfolder = "transformer" + + @classmethod + def get_transformer_block_names(cls): + return ["layers"] + @register_to_config def __init__( diff --git a/extensions_built_in/diffusion_models/f_light/f_light.py b/extensions_built_in/diffusion_models/f_light/f_light.py index 813e6c3..2f3795c 100644 --- a/extensions_built_in/diffusion_models/f_light/f_light.py +++ b/extensions_built_in/diffusion_models/f_light/f_light.py @@ -6,15 +6,15 @@ import yaml from toolkit.config_modules import GenerateImageConfig, ModelConfig from PIL import Image from toolkit.models.base_model import BaseModel +from toolkit.models.v2.text_encoders.t5 import T5TextEncoder +from toolkit.models.v2.vae.autoencoder_kl import KLVAE from toolkit.basic import flush -from diffusers import AutoencoderKL from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler from toolkit.dequantize import patch_dequantization_on_save from toolkit.accelerator import unwrap_model from optimum.quanto import freeze, QTensor from toolkit.util.quantize import quantize, get_qtype -from transformers import T5TokenizerFast, T5EncoderModel from .src import FLitePipeline, DiT if TYPE_CHECKING: @@ -74,11 +74,7 @@ class FLiteModel(BaseModel): self.print_and_status_update("Loading transformer") - transformer = DiT.from_pretrained( - model_path, - subfolder="dit_model", - torch_dtype=dtype, - ) + transformer = DiT.load_model(model_path, dtype=dtype) transformer.to(self.quantize_device, dtype=dtype) @@ -97,31 +93,16 @@ class FLiteModel(BaseModel): flush() self.print_and_status_update("Loading T5") - tokenizer = T5TokenizerFast.from_pretrained( - extras_path, subfolder="tokenizer", torch_dtype=dtype + tokenizer = T5TextEncoder.load_tokenizer(extras_path, subfolder="tokenizer") + text_encoder = T5TextEncoder.load_model( + extras_path, dtype=dtype, subfolder="text_encoder" ) - text_encoder = T5EncoderModel.from_pretrained( - extras_path, subfolder="text_encoder", torch_dtype=dtype - ) - text_encoder.to(self.device_torch, dtype=dtype) - flush() - - if self.model_config.quantize_te: - self.print_and_status_update("Quantizing T5") - quantize(text_encoder, weights=get_qtype( - self.model_config.qtype)) - freeze(text_encoder) - flush() + self.prepare_text_encoder(text_encoder, dtype=dtype) self.noise_scheduler = FLiteModel.get_train_scheduler() self.print_and_status_update("Loading VAE") - vae = AutoencoderKL.from_pretrained( - extras_path, - subfolder="vae", - torch_dtype=dtype - ) - vae = vae.to(self.device_torch, dtype=dtype) + vae = KLVAE.load_model(extras_path, dtype=dtype, device=self.device_torch) self.print_and_status_update("Making pipe") diff --git a/extensions_built_in/diffusion_models/f_light/src/model.py b/extensions_built_in/diffusion_models/f_light/src/model.py index 903d492..26b3a9c 100644 --- a/extensions_built_in/diffusion_models/f_light/src/model.py +++ b/extensions_built_in/diffusion_models/f_light/src/model.py @@ -7,6 +7,8 @@ import torch.nn.functional as F from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin from diffusers.models.modeling_utils import ModelMixin + +from toolkit.models.v2._mixin import OstrisModelMixin from diffusers.utils.accelerate_utils import apply_forward_hook from einops import rearrange from peft import get_peft_model_state_dict, set_peft_model_state_dict @@ -302,7 +304,13 @@ def apply_rotary_emb(x, cos, sin): return torch.cat([y1, y2], 3).to(dtype=orig_dtype) -class DiT(ModelMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin): # type: ignore[misc] +class DiT(ModelMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin, OstrisModelMixin): # type: ignore[misc] + aitk_subfolder = "dit_model" + + @classmethod + def get_transformer_block_names(cls): + return ["blocks"] + _supports_gradient_checkpointing = True @register_to_config diff --git a/extensions_built_in/diffusion_models/flux2/flux2_model.py b/extensions_built_in/diffusion_models/flux2/flux2_model.py index e39e468..20aedf0 100644 --- a/extensions_built_in/diffusion_models/flux2/flux2_model.py +++ b/extensions_built_in/diffusion_models/flux2/flux2_model.py @@ -21,7 +21,11 @@ from toolkit.util.quantize import quantize, get_qtype, quantize_model from transformers import AutoProcessor, Mistral3ForConditionalGeneration from .src.model import Flux2, Flux2Params from .src.pipeline import Flux2Pipeline -from .src.autoencoder import AutoEncoder, AutoEncoderParams, AutoEncoderSmallDecoderParams +from toolkit.models.v2.vae.flux2_kl import ( + AutoEncoder, + AutoEncoderParams, + AutoEncoderSmallDecoderParams, +) from safetensors.torch import load_file, save_file from PIL import Image import torch.nn.functional as F diff --git a/extensions_built_in/diffusion_models/flux2/src/autoencoder.py b/extensions_built_in/diffusion_models/flux2/src/autoencoder.py deleted file mode 100644 index 3d8230e..0000000 --- a/extensions_built_in/diffusion_models/flux2/src/autoencoder.py +++ /dev/null @@ -1,435 +0,0 @@ -from dataclasses import dataclass, field - -import torch -from einops import rearrange -from torch import Tensor, nn -import math -import torch.utils.checkpoint as ckpt - - -@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 diff --git a/extensions_built_in/diffusion_models/flux2/src/pipeline.py b/extensions_built_in/diffusion_models/flux2/src/pipeline.py index 8e64639..aec4dd0 100644 --- a/extensions_built_in/diffusion_models/flux2/src/pipeline.py +++ b/extensions_built_in/diffusion_models/flux2/src/pipeline.py @@ -12,7 +12,7 @@ from diffusers.utils import ( from diffusers.utils.torch_utils import randn_tensor from diffusers.pipelines.pipeline_utils import DiffusionPipeline from diffusers.utils import BaseOutput -from .autoencoder import AutoEncoder +from toolkit.models.v2.vae.flux2_kl import AutoEncoder from .model import Flux2 from einops import rearrange from transformers import AutoProcessor, Mistral3ForConditionalGeneration diff --git a/extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py b/extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py index e9eeee5..927477f 100644 --- a/extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py +++ b/extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py @@ -8,7 +8,11 @@ from toolkit import train_tools from toolkit.config_modules import GenerateImageConfig, ModelConfig from PIL import Image from toolkit.models.base_model import BaseModel -from diffusers import FluxTransformer2DModel, AutoencoderKL, FluxKontextPipeline +from toolkit.models.v2.text_encoders.t5 import T5TextEncoder +from toolkit.models.v2.text_encoders.clip import CLIPTextEncoder +from toolkit.models.v2.vae.autoencoder_kl import KLVAE +from diffusers import FluxKontextPipeline +from toolkit.models.v2.diffusion_models.flux import FluxTransformer2DModel from toolkit.basic import flush from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler @@ -18,7 +22,6 @@ from toolkit.accelerator import get_accelerator, unwrap_model from optimum.quanto import freeze, QTensor from toolkit.util.mask import generate_random_mask, random_dialate_mask from toolkit.util.quantize import quantize, get_qtype -from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer from einops import rearrange, repeat import random import torch.nn.functional as F @@ -80,11 +83,7 @@ class FluxKontextModel(BaseModel): # so we need this for the VAE, te, etc base_model_path = self.model_config.extras_name_or_path - transformer_path = model_path - transformer_subfolder = 'transformer' - if os.path.exists(transformer_path): - transformer_subfolder = None - transformer_path = os.path.join(transformer_path, 'transformer') + if os.path.exists(model_path): # check if the path is a full checkpoint. te_folder_path = os.path.join(model_path, 'text_encoder') # if we have the te, this folder is a full checkpoint, use it as the base @@ -92,11 +91,7 @@ class FluxKontextModel(BaseModel): base_model_path = model_path self.print_and_status_update("Loading transformer") - transformer = FluxTransformer2DModel.from_pretrained( - transformer_path, - subfolder=transformer_subfolder, - torch_dtype=dtype - ) + transformer = FluxTransformer2DModel.load_model(model_path, dtype=dtype) transformer.to(self.quantize_device, dtype=dtype) if self.model_config.quantize: @@ -114,32 +109,18 @@ class FluxKontextModel(BaseModel): flush() self.print_and_status_update("Loading T5") - tokenizer_2 = T5TokenizerFast.from_pretrained( - base_model_path, subfolder="tokenizer_2", torch_dtype=dtype - ) - text_encoder_2 = T5EncoderModel.from_pretrained( - base_model_path, subfolder="text_encoder_2", torch_dtype=dtype - ) - text_encoder_2.to(self.device_torch, dtype=dtype) - flush() - - if self.model_config.quantize_te: - self.print_and_status_update("Quantizing T5") - quantize(text_encoder_2, weights=get_qtype( - self.model_config.qtype)) - freeze(text_encoder_2) - flush() + tokenizer_2 = T5TextEncoder.load_tokenizer(base_model_path) + text_encoder_2 = T5TextEncoder.load_model(base_model_path, dtype=dtype) + self.prepare_text_encoder(text_encoder_2, dtype=dtype) self.print_and_status_update("Loading CLIP") - text_encoder = CLIPTextModel.from_pretrained( - base_model_path, subfolder="text_encoder", torch_dtype=dtype) - tokenizer = CLIPTokenizer.from_pretrained( - base_model_path, subfolder="tokenizer", torch_dtype=dtype) - text_encoder.to(self.device_torch, dtype=dtype) + text_encoder = CLIPTextEncoder.load_model( + base_model_path, dtype=dtype, device=self.device_torch + ) + tokenizer = CLIPTextEncoder.load_tokenizer(base_model_path, use_fast=False) self.print_and_status_update("Loading VAE") - vae = AutoencoderKL.from_pretrained( - base_model_path, subfolder="vae", torch_dtype=dtype) + vae = KLVAE.load_model(base_model_path, dtype=dtype) self.noise_scheduler = FluxKontextModel.get_train_scheduler() diff --git a/extensions_built_in/diffusion_models/hidream/hidream_e1_model.py b/extensions_built_in/diffusion_models/hidream/hidream_e1_model.py index 5306ad5..2981807 100644 --- a/extensions_built_in/diffusion_models/hidream/hidream_e1_model.py +++ b/extensions_built_in/diffusion_models/hidream/hidream_e1_model.py @@ -7,7 +7,7 @@ from toolkit.accelerator import unwrap_model import torch from toolkit.prompt_utils import PromptEmbeds from toolkit.config_modules import GenerateImageConfig -from diffusers.models import HiDreamImageTransformer2DModel +from toolkit.models.v2.diffusion_models.hidream import HiDreamImageTransformer2DModel import torch.nn.functional as F from PIL import Image diff --git a/extensions_built_in/diffusion_models/hidream/hidream_model.py b/extensions_built_in/diffusion_models/hidream/hidream_model.py index 71921aa..3d72f04 100644 --- a/extensions_built_in/diffusion_models/hidream/hidream_model.py +++ b/extensions_built_in/diffusion_models/hidream/hidream_model.py @@ -9,7 +9,10 @@ from toolkit import train_tools from toolkit.config_modules import GenerateImageConfig, ModelConfig from PIL import Image from toolkit.models.base_model import BaseModel -from diffusers import AutoencoderKL, TorchAoConfig +from toolkit.models.v2.text_encoders.t5 import T5TextEncoder +from toolkit.models.v2.text_encoders.clip import CLIPTextEncoderWithProjection +from toolkit.models.v2.vae.autoencoder_kl import KLVAE +from diffusers import TorchAoConfig from toolkit.basic import flush from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler @@ -19,7 +22,6 @@ from toolkit.accelerator import get_accelerator, unwrap_model from optimum.quanto import freeze, QTensor from toolkit.util.mask import generate_random_mask, random_dialate_mask from toolkit.util.quantize import quantize, get_qtype -from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer, TorchAoConfig as TorchAoConfigTransformers from .src.pipelines.hidream_image.pipeline_hidream_image import HiDreamImagePipeline from .src.models.transformers.transformer_hidream_image import HiDreamImageTransformer2DModel from .src.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler @@ -28,15 +30,6 @@ from einops import rearrange, repeat import random import torch.nn.functional as F from tqdm import tqdm -from transformers import ( - CLIPTextModelWithProjection, - CLIPTokenizer, - T5EncoderModel, - T5Tokenizer, - LlamaForCausalLM, - PreTrainedTokenizerFast -) - if TYPE_CHECKING: from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO @@ -125,10 +118,8 @@ class HidreamModel(BaseModel): self.print_and_status_update("Loading transformer") - transformer = self.hidream_transformer_class.from_pretrained( - model_path, - subfolder="transformer", - torch_dtype=torch.bfloat16 + transformer = self.hidream_transformer_class.load_model( + model_path, dtype=torch.bfloat16 ) if not self.low_vram: @@ -164,44 +155,34 @@ class HidreamModel(BaseModel): self.print_and_status_update("Loading vae") - vae = AutoencoderKL.from_pretrained( - extras_path, - subfolder="vae", - torch_dtype=torch.bfloat16 - ).to(self.device_torch, dtype=dtype) + vae = KLVAE.load_model(extras_path, dtype=torch.bfloat16).to( + self.device_torch, dtype=dtype + ) self.print_and_status_update("Loading clip encoders") - text_encoder = CLIPTextModelWithProjection.from_pretrained( - extras_path, - subfolder="text_encoder", - torch_dtype=torch.bfloat16 + text_encoder = CLIPTextEncoderWithProjection.load_model( + extras_path, dtype=torch.bfloat16 ).to(self.device_torch, dtype=dtype) - - tokenizer = CLIPTokenizer.from_pretrained( - extras_path, - subfolder="tokenizer" + + tokenizer = CLIPTextEncoderWithProjection.load_tokenizer( + extras_path, use_fast=False ) - - text_encoder_2 = CLIPTextModelWithProjection.from_pretrained( - extras_path, - subfolder="text_encoder_2", - torch_dtype=torch.bfloat16 + + text_encoder_2 = CLIPTextEncoderWithProjection.load_model( + extras_path, dtype=torch.bfloat16, subfolder="text_encoder_2" ).to(self.device_torch, dtype=dtype) - - tokenizer_2 = CLIPTokenizer.from_pretrained( - extras_path, - subfolder="tokenizer_2" + + tokenizer_2 = CLIPTextEncoderWithProjection.load_tokenizer( + extras_path, subfolder="tokenizer_2", use_fast=False ) flush() self.print_and_status_update("Loading T5 encoders") - text_encoder_3 = T5EncoderModel.from_pretrained( - extras_path, - subfolder="text_encoder_3", - torch_dtype=torch.bfloat16 + text_encoder_3 = T5TextEncoder.load_model( + extras_path, dtype=torch.bfloat16, subfolder="text_encoder_3" ).to(self.device_torch, dtype=dtype) if self.model_config.quantize_te: @@ -211,9 +192,8 @@ class HidreamModel(BaseModel): freeze(text_encoder_3) flush() - tokenizer_3 = T5Tokenizer.from_pretrained( - extras_path, - subfolder="tokenizer_3" + tokenizer_3 = T5TextEncoder.load_tokenizer( + extras_path, subfolder="tokenizer_3", use_fast=False ) flush() diff --git a/extensions_built_in/diffusion_models/hidream/src/models/transformers/transformer_hidream_image.py b/extensions_built_in/diffusion_models/hidream/src/models/transformers/transformer_hidream_image.py index f7eb104..1c23bf7 100644 --- a/extensions_built_in/diffusion_models/hidream/src/models/transformers/transformer_hidream_image.py +++ b/extensions_built_in/diffusion_models/hidream/src/models/transformers/transformer_hidream_image.py @@ -8,6 +8,8 @@ from einops import repeat from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin from diffusers.models.modeling_utils import ModelMixin + +from toolkit.models.v2._mixin import OstrisModelMixin from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers from diffusers.utils.torch_utils import maybe_allow_in_graph from diffusers.models.modeling_outputs import Transformer2DModelOutput @@ -228,9 +230,15 @@ class HiDreamImageBlock(nn.Module): ) class HiDreamImageTransformer2DModel( - ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin + ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, OstrisModelMixin ): _supports_gradient_checkpointing = True + aitk_subfolder = "transformer" + + @classmethod + def get_transformer_block_names(cls): + return ["double_stream_blocks", "single_stream_blocks"] + _no_split_modules = ["HiDreamImageBlock"] @register_to_config diff --git a/extensions_built_in/diffusion_models/ideogram4/ideogram4.py b/extensions_built_in/diffusion_models/ideogram4/ideogram4.py index 390f7a4..2048bdd 100644 --- a/extensions_built_in/diffusion_models/ideogram4/ideogram4.py +++ b/extensions_built_in/diffusion_models/ideogram4/ideogram4.py @@ -26,7 +26,11 @@ from huggingface_hub.errors import EntryNotFoundError from transformers import AutoModel, AutoTokenizer from .src.transformer import Ideogram4Config, Ideogram4Transformer2DModel -from .src.vae import AutoEncoder, AutoEncoderParams, convert_diffusers_state_dict +from toolkit.models.v2.vae.flux2_kl import ( + AutoEncoder, + AutoEncoderParams, + convert_diffusers_state_dict, +) from .src.latent_norm import get_latent_norm from .src.pipeline import ( Ideogram4Pipeline, diff --git a/extensions_built_in/diffusion_models/krea2/krea2.py b/extensions_built_in/diffusion_models/krea2/krea2.py index 22d75a8..b0fa0a6 100644 --- a/extensions_built_in/diffusion_models/krea2/krea2.py +++ b/extensions_built_in/diffusion_models/krea2/krea2.py @@ -25,18 +25,21 @@ from safetensors.torch import load_file, save_file import huggingface_hub from huggingface_hub.errors import EntryNotFoundError -from diffusers import AutoencoderKLQwenImage from transformers import ( AutoProcessor, AutoTokenizer, Qwen2TokenizerFast, - Qwen3VLForConditionalGeneration, ) from optimum.quanto import freeze from toolkit.config_modules import GenerateImageConfig, ModelConfig, NetworkConfig from toolkit.lora_special import LoRASpecialNetwork from toolkit.models.base_model import BaseModel +from toolkit.models.v2.vae.qwen_image import QwenImageVAE, QwenImageVAEHolderMixin +from toolkit.models.v2.text_encoders.qwen3_vl import ( + Qwen3VLTextEncoder, + patch_qwen_vl_patch_embed, +) from toolkit.basic import flush from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import ( @@ -100,30 +103,6 @@ QWEN_IMAGE_VAE_PATH = "Qwen/Qwen-Image" HF_TOKEN = os.getenv("HF_TOKEN", None) -def patch_qwen_vl_patch_embed(model): - """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. (Same patch as the - Qwen3VLCaptioner extension.)""" - 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 - - def _load_mmdit_state_dict(name_or_path: str, filename: Optional[str]) -> dict: """Load the MMDiT weights from a local safetensors file/dir or the HF hub. @@ -163,7 +142,7 @@ def _load_mmdit_state_dict(name_or_path: str, filename: Optional[str]) -> dict: return load_file(path) -class Krea2Model(BaseModel): +class Krea2Model(QwenImageVAEHolderMixin, BaseModel): arch = "krea2" def __init__( @@ -268,8 +247,8 @@ class Krea2Model(BaseModel): processor = Qwen2TokenizerFast.from_pretrained( te_path, max_length=self.max_text_length, token=HF_TOKEN ) - text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( - te_path, torch_dtype=dtype, token=HF_TOKEN + text_encoder = Qwen3VLTextEncoder.load_model( + te_path, dtype=dtype, subfolder="", token=HF_TOKEN ) vl_processor = None if self.is_edit: @@ -281,8 +260,7 @@ class Krea2Model(BaseModel): else: # We only ever encode text, so the vision tower is dead weight -- drop it to # free VRAM and skip loading its (bf16-slow) Conv3d patch_embed onto the GPU. - if getattr(text_encoder.model, "visual", None) is not None: - text_encoder.model.visual = None + text_encoder.drop_vision_tower() text_encoder.eval() text_encoder.requires_grad_(False) flush() @@ -291,8 +269,8 @@ class Krea2Model(BaseModel): def _load_vae(self): vae_path = self.model_config.model_kwargs.get("vae_path", QWEN_IMAGE_VAE_PATH) self.print_and_status_update(f"Loading Qwen-Image VAE from {vae_path}") - vae = AutoencoderKLQwenImage.from_pretrained( - vae_path, subfolder="vae", torch_dtype=self.vae_torch_dtype, token=HF_TOKEN + vae = QwenImageVAE.load_model( + vae_path, dtype=self.vae_torch_dtype, token=HF_TOKEN ) vae.eval() vae.requires_grad_(False) @@ -772,76 +750,9 @@ class Krea2Model(BaseModel): return False # ------------------------------------------------------------------ - # VAE (Qwen-Image AutoencoderKLQwenImage -- same handling as qwen_image arch) + # VAE (Qwen-Image AutoencoderKLQwenImage -- shared QwenImageVAEHolderMixin) # ------------------------------------------------------------------ - def encode_images(self, image_list: List[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) - 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) - - # AutoencoderKLQwenImage is a video VAE: add a frame dim. - images = images.unsqueeze(2) - 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 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 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; tile it when low on VRAM (decode - # only -- encode stays untiled). - tiled = 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 frame dim - return images.to(device, dtype=dtype) - + vae_decode_tiled_on_low_vram = True # ------------------------------------------------------------------ # Saving / bookkeeping # ------------------------------------------------------------------ diff --git a/extensions_built_in/diffusion_models/mageflow/mageflow.py b/extensions_built_in/diffusion_models/mageflow/mageflow.py index 589f9f2..49b3530 100644 --- a/extensions_built_in/diffusion_models/mageflow/mageflow.py +++ b/extensions_built_in/diffusion_models/mageflow/mageflow.py @@ -34,7 +34,8 @@ from torchvision.transforms.functional import to_tensor from safetensors.torch import load_file, save_file import huggingface_hub -from transformers import AutoProcessor, AutoTokenizer, Qwen3VLForConditionalGeneration +from transformers import AutoProcessor, AutoTokenizer +from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder from optimum.quanto import freeze from toolkit.config_modules import GenerateImageConfig, ModelConfig @@ -209,8 +210,11 @@ class MageFlowModel(BaseModel): self.print_and_status_update(f"Loading Qwen3-VL text encoder from {te_path}") tokenizer = AutoTokenizer.from_pretrained(te_path, token=HF_TOKEN, **te_kwargs) - text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( - te_path, torch_dtype=dtype, token=HF_TOKEN, **te_kwargs + text_encoder = Qwen3VLTextEncoder.load_model( + te_path, + dtype=dtype, + subfolder=te_kwargs.get("subfolder", ""), + token=HF_TOKEN, ) vl_processor = None if self.is_edit: @@ -224,8 +228,7 @@ class MageFlowModel(BaseModel): else: # We only ever encode text, so the vision tower is dead weight -- # drop it to free VRAM. - if getattr(text_encoder.model, "visual", None) is not None: - text_encoder.model.visual = None + text_encoder.drop_vision_tower() text_encoder.eval() text_encoder.requires_grad_(False) flush() diff --git a/extensions_built_in/diffusion_models/mageflow/src/text_encoder.py b/extensions_built_in/diffusion_models/mageflow/src/text_encoder.py index 4f3508a..0f05ed6 100644 --- a/extensions_built_in/diffusion_models/mageflow/src/text_encoder.py +++ b/extensions_built_in/diffusion_models/mageflow/src/text_encoder.py @@ -18,7 +18,8 @@ import math from typing import List, Optional import torch -import torch.nn.functional as F + +from toolkit.models.v2.text_encoders.qwen3_vl import patch_qwen_vl_patch_embed # Prompt templates from the reference (mage_flow/models/utils.py). ``start_idx`` @@ -57,30 +58,6 @@ def edit_prompt_body(instruction: str, num_refs: int) -> str: return prefix + instruction -def patch_qwen_vl_patch_embed(model): - """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. (Same patch as the krea2 - extension / Qwen3VLCaptioner.)""" - 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 - - def resize_vl_images( images: List[torch.Tensor], max_long_edge: int = 384 ) -> List["PIL.Image.Image"]: diff --git a/extensions_built_in/diffusion_models/nucleus_image/nucleus_image_model.py b/extensions_built_in/diffusion_models/nucleus_image/nucleus_image_model.py index 54ec7ff..ee80c39 100644 --- a/extensions_built_in/diffusion_models/nucleus_image/nucleus_image_model.py +++ b/extensions_built_in/diffusion_models/nucleus_image/nucleus_image_model.py @@ -6,21 +6,25 @@ import torch import yaml from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.models.base_model import BaseModel +from toolkit.models.v2.vae.qwen_image import QwenImageVAE, QwenImageVAEHolderMixin +from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder from toolkit.basic import flush from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import ( CustomFlowMatchEulerDiscreteScheduler, ) from toolkit.accelerator import unwrap_model -from optimum.quanto import freeze -from toolkit.util.quantize import quantize, get_qtype, quantize_model +from toolkit.util.quantize import quantize_model from toolkit.memory_management import MemoryManager -from transformers import Qwen3VLForConditionalGeneration, Qwen3VLProcessor +from transformers import Qwen3VLProcessor import torch.nn.functional as F try: - from diffusers import NucleusMoEImagePipeline, NucleusMoEImageTransformer2DModel, AutoencoderKLQwenImage + from diffusers import NucleusMoEImagePipeline, AutoencoderKLQwenImage + from toolkit.models.v2.diffusion_models.nucleus_image import ( + NucleusMoEImageTransformer2DModel, + ) from diffusers.models.transformers.transformer_nucleusmoe_image import SwiGLUExperts except ImportError: raise ImportError( @@ -46,7 +50,7 @@ scheduler_config = { } -class NucleusImageModel(BaseModel): +class NucleusImageModel(QwenImageVAEHolderMixin, BaseModel): arch = "nucleus_image" def __init__( @@ -81,20 +85,14 @@ class NucleusImageModel(BaseModel): self.print_and_status_update("Loading transformer") - transformer_path = model_path - transformer_subfolder = "transformer" - if os.path.exists(transformer_path): - transformer_subfolder = None - transformer_path = os.path.join(transformer_path, "transformer") + if os.path.exists(model_path): # check if the path is a full checkpoint. te_folder_path = os.path.join(model_path, "text_encoder") # if we have the te, this folder is a full checkpoint, use it as the base if os.path.exists(te_folder_path): base_model_path = model_path - transformer = NucleusMoEImageTransformer2DModel.from_pretrained( - transformer_path, subfolder=transformer_subfolder, torch_dtype=dtype - ) + transformer = NucleusMoEImageTransformer2DModel.load_model(model_path, dtype=dtype) # handle versions of pytorch that don't have grouped mm, by disabling it in the SwiGLUExperts if not hasattr(torch.nn.functional, "grouped_mm"): @@ -129,33 +127,13 @@ class NucleusImageModel(BaseModel): tokenizer = Qwen3VLProcessor.from_pretrained( base_model_path, subfolder="processor", torch_dtype=dtype ) - text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( - base_model_path, subfolder="text_encoder", torch_dtype=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: - self.print_and_status_update("Quantizing Text Encoder") - quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) - freeze(text_encoder) - flush() + text_encoder = Qwen3VLTextEncoder.load_model(base_model_path, dtype=dtype) + self.prepare_text_encoder(text_encoder, dtype=dtype) self.print_and_status_update("Loading VAE") - vae = AutoencoderKLQwenImage.from_pretrained( - base_model_path, subfolder="vae", torch_dtype=dtype - ).to(self.device_torch, dtype=dtype) + vae = QwenImageVAE.load_model( + base_model_path, dtype=dtype, device=self.device_torch + ) self.noise_scheduler = NucleusImageModel.get_train_scheduler() @@ -214,41 +192,6 @@ class NucleusImageModel(BaseModel): return pipeline - def encode_images(self, image_list: List[torch.Tensor], device=None, dtype=None): - if device is None: - device = self.vae_device_torch - if dtype is None: - dtype = self.vae_torch_dtype - - # Move to vae to device if on cpu - if self.vae.device == torch.device("cpu"): - self.vae.to(device) - self.vae.eval() - self.vae.requires_grad_(False) - # move to device and dtype - image_list = [image.to(device, dtype=dtype) for image in image_list] - images = torch.stack(image_list).to(device, dtype=dtype) - # it uses wan vae, so add dim for frame count - - images = images.unsqueeze(2) - 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.to(device, dtype=dtype) - - latents = latents.squeeze(2) # remove the frame count dimension - - return latents - def generate_single_image( self, pipeline: NucleusMoEImagePipeline, diff --git a/extensions_built_in/diffusion_models/omnigen2/__init__.py b/extensions_built_in/diffusion_models/omnigen2/__init__.py index edb10bd..77741be 100644 --- a/extensions_built_in/diffusion_models/omnigen2/__init__.py +++ b/extensions_built_in/diffusion_models/omnigen2/__init__.py @@ -96,8 +96,8 @@ class OmniGen2Model(BaseModel): self.print_and_status_update("Loading transformer") - transformer = OmniGen2Transformer2DModel.from_pretrained( - model_path, subfolder="transformer", torch_dtype=torch.bfloat16 + transformer = OmniGen2Transformer2DModel.load_model( + model_path, dtype=torch.bfloat16 ) if not self.low_vram: diff --git a/extensions_built_in/diffusion_models/omnigen2/src/models/transformers/transformer_omnigen2.py b/extensions_built_in/diffusion_models/omnigen2/src/models/transformers/transformer_omnigen2.py index 8e7fef6..abb3e0c 100644 --- a/extensions_built_in/diffusion_models/omnigen2/src/models/transformers/transformer_omnigen2.py +++ b/extensions_built_in/diffusion_models/omnigen2/src/models/transformers/transformer_omnigen2.py @@ -15,6 +15,8 @@ from diffusers.models.attention_processor import Attention from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_utils import ModelMixin +from toolkit.models.v2._mixin import OstrisModelMixin + from ..attention_processor import OmniGen2AttnProcessorFlash2Varlen, OmniGen2AttnProcessor from .repo import OmniGen2RotaryPosEmbed from .block_lumina2 import LuminaLayerNormContinuous, LuminaRMSNormZero, LuminaFeedForward, Lumina2CombinedTimestepCaptionEmbedding @@ -177,7 +179,15 @@ class OmniGen2TransformerBlock(nn.Module): return hidden_states -class OmniGen2Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): +class OmniGen2Transformer2DModel( + ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, OstrisModelMixin +): + aitk_subfolder = "transformer" + + @classmethod + def get_transformer_block_names(cls): + return ["noise_refiner", "ref_image_refiner", "context_refiner", "layers"] + """ OmniGen2 Transformer 2D Model. diff --git a/extensions_built_in/diffusion_models/prx_pixel_t2i/prx_pixel_t2i.py b/extensions_built_in/diffusion_models/prx_pixel_t2i/prx_pixel_t2i.py index b6f2137..ffea929 100644 --- a/extensions_built_in/diffusion_models/prx_pixel_t2i/prx_pixel_t2i.py +++ b/extensions_built_in/diffusion_models/prx_pixel_t2i/prx_pixel_t2i.py @@ -126,9 +126,7 @@ class PRXPixelT2IModel(BaseModel): self.print_and_status_update("Loading transformer") # from_pretrained reads config.json (bottleneck_size, resolution_embeds, # in_channels=3, ...) and the safetensors in one shot. - transformer = PRXTransformer2DModel.from_pretrained( - model_path, subfolder="transformer", torch_dtype=dtype - ) + transformer = PRXTransformer2DModel.load_model(model_path, dtype=dtype) transformer.to(dtype=dtype) flush() diff --git a/extensions_built_in/diffusion_models/prx_pixel_t2i/src/transformer_prx.py b/extensions_built_in/diffusion_models/prx_pixel_t2i/src/transformer_prx.py index e83d8a1..1dd6dee 100644 --- a/extensions_built_in/diffusion_models/prx_pixel_t2i/src/transformer_prx.py +++ b/extensions_built_in/diffusion_models/prx_pixel_t2i/src/transformer_prx.py @@ -40,6 +40,8 @@ from diffusers.models.attention_dispatch import dispatch_attention_fn from diffusers.models.embeddings import get_timestep_embedding from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_utils import ModelMixin + +from toolkit.models.v2._mixin import OstrisModelMixin from diffusers.models.normalization import RMSNorm @@ -666,7 +668,7 @@ def seq2img(seq: torch.Tensor, patch_size: int, shape: torch.Tensor) -> torch.Te return seq -class PRXTransformer2DModel(ModelMixin, ConfigMixin, AttentionMixin): +class PRXTransformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, OstrisModelMixin): r""" Transformer-based 2D model for text to image generation. @@ -702,6 +704,12 @@ class PRXTransformer2DModel(ModelMixin, ConfigMixin, AttentionMixin): `PRXResolutionEmbedder`. Used by the PRX-7B variant. """ + aitk_subfolder = "transformer" + + @classmethod + def get_transformer_block_names(cls): + return ["blocks"] + config_name = "config.json" _supports_gradient_checkpointing = True diff --git a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py index 5199321..54116a5 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py @@ -7,28 +7,25 @@ from toolkit import train_tools from toolkit.config_modules import GenerateImageConfig, ModelConfig from PIL import Image from toolkit.models.base_model import BaseModel +from toolkit.models.v2.vae.qwen_image import QwenImageVAE, QwenImageVAEHolderMixin from toolkit.basic import flush from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import ( CustomFlowMatchEulerDiscreteScheduler, ) from toolkit.accelerator import get_accelerator, unwrap_model -from optimum.quanto import freeze, QTensor -from toolkit.util.quantize import quantize, get_qtype, quantize_model +from toolkit.util.quantize import quantize_model import torch.nn.functional as F from toolkit.memory_management import MemoryManager from safetensors.torch import load_file from diffusers import ( QwenImagePipeline, - QwenImageTransformer2DModel, AutoencoderKLQwenImage, ) -from transformers import ( - Qwen2_5_VLForConditionalGeneration, - Qwen2Tokenizer, - Qwen2VLProcessor, -) +from transformers import Qwen2VLProcessor +from toolkit.models.v2.diffusion_models.qwen_image import QwenImageTransformer2DModel +from toolkit.models.v2.text_encoders.qwen25_vl import Qwen25VLTextEncoder from tqdm import tqdm from toolkit.util.qwen_vae_gradient_checkpointing import patch_qwen_vae_gradient_checkpointing @@ -55,7 +52,7 @@ scheduler_config = { } -class QwenImageModel(BaseModel): +class QwenImageModel(QwenImageVAEHolderMixin, BaseModel): arch = "qwen_image" _qwen_image_keep_visual = False _qwen_pipeline = QwenImagePipeline @@ -97,31 +94,14 @@ class QwenImageModel(BaseModel): self.print_and_status_update("Loading transformer") - if model_path.endswith(".safetensors"): - # load the safetensors file - transformer = QwenImageTransformer2DModel.from_single_file( - model_path, - config="Qwen/Qwen-Image", - subfolder="transformer", - torch_dtype=model_dtype, - ) - transformer.to(model_dtype) + if not model_path.endswith(".safetensors") and os.path.exists(model_path): + # check if the path is a full checkpoint. + te_folder_path = os.path.join(model_path, "text_encoder") + # if we have the te, this folder is a full checkpoint, use it as the base + if os.path.exists(te_folder_path): + base_model_path = model_path - else: - transformer_path = model_path - transformer_subfolder = "transformer" - if os.path.exists(transformer_path): - transformer_subfolder = None - transformer_path = os.path.join(transformer_path, "transformer") - # check if the path is a full checkpoint. - te_folder_path = os.path.join(model_path, "text_encoder") - # if we have the te, this folder is a full checkpoint, use it as the base - if os.path.exists(te_folder_path): - base_model_path = model_path - - transformer = QwenImageTransformer2DModel.from_pretrained( - transformer_path, subfolder=transformer_subfolder, torch_dtype=dtype - ) + transformer = QwenImageTransformer2DModel.load_model(model_path, dtype=model_dtype) if self.model_config.quantize: self.print_and_status_update("Quantizing Transformer") @@ -145,41 +125,18 @@ class QwenImageModel(BaseModel): flush() self.print_and_status_update("Text Encoder") - tokenizer = Qwen2Tokenizer.from_pretrained( - base_model_path, subfolder="tokenizer", torch_dtype=dtype - ) - text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained( - base_model_path, subfolder="text_encoder", torch_dtype=dtype - ) + tokenizer = Qwen25VLTextEncoder.load_tokenizer(base_model_path, use_fast=False) + text_encoder = Qwen25VLTextEncoder.load_model(base_model_path, dtype=dtype) # remove the visual model as it is not needed for image generation self.processor = None if not self._qwen_image_keep_visual: - text_encoder.model.visual = None + text_encoder.drop_vision_tower() - 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: - self.print_and_status_update("Quantizing Text Encoder") - quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) - freeze(text_encoder) - flush() + self.prepare_text_encoder(text_encoder, dtype=dtype) self.print_and_status_update("Loading VAE") - vae = AutoencoderKLQwenImage.from_pretrained( - base_model_path, subfolder="vae", torch_dtype=dtype - ) + vae = QwenImageVAE.load_model(base_model_path, dtype=dtype) self.noise_scheduler = QwenImageModel.get_train_scheduler() @@ -417,69 +374,3 @@ class QwenImageModel(BaseModel): lora_keys_use_comfy_prefix = True - def encode_images(self, image_list: List[torch.Tensor], device=None, dtype=None): - if device is None: - device = self.vae_device_torch - if dtype is None: - dtype = self.vae_torch_dtype - - # Move to vae to device if on cpu - if self.vae.device == torch.device("cpu"): - self.vae.to(device) - self.vae.eval() - self.vae.requires_grad_(False) - # move to device and dtype - image_list = [image.to(device, dtype=dtype) for image in image_list] - images = torch.stack(image_list).to(device, dtype=dtype) - # it uses wan vae, so add dim for frame count - - images = images.unsqueeze(2) - 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.to(device, dtype=dtype) - - latents = latents.squeeze(2) # remove the frame count dimension - - return latents - - 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) - - # add frame count dim for wan vae - latents = latents.unsqueeze(2) - - 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 - - images = self.vae.decode(latents).sample - - images = images.squeeze(2) # remove the frame count dimension - - return images.to(device, dtype=dtype) diff --git a/extensions_built_in/diffusion_models/wan22/wan22_14b_model.py b/extensions_built_in/diffusion_models/wan22/wan22_14b_model.py index 37bcb7a..3a26897 100644 --- a/extensions_built_in/diffusion_models/wan22/wan22_14b_model.py +++ b/extensions_built_in/diffusion_models/wan22/wan22_14b_model.py @@ -17,7 +17,7 @@ from toolkit.samplers.custom_flowmatch_sampler import ( ) from toolkit.util.quantize import quantize_model from .wan22_pipeline import Wan22Pipeline -from diffusers import WanTransformer3DModel +from toolkit.models.v2.diffusion_models.wan import WanTransformer3DModel from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO from torchvision.transforms import functional as TF @@ -287,11 +287,9 @@ class Wan2214bModel(Wan21): self.print_and_status_update("Loading transformer 1") dtype = self.torch_dtype - transformer_1 = WanTransformer3DModel.from_pretrained( - transformer_path_1, - subfolder=subfolder_1, - torch_dtype=dtype, - ).to(dtype=dtype) + transformer_1 = WanTransformer3DModel.load_model( + transformer_path_1, dtype=dtype, subfolder=subfolder_1 + ) flush() @@ -317,11 +315,9 @@ class Wan2214bModel(Wan21): self.print_and_status_update("Loading transformer 2") dtype = self.torch_dtype - transformer_2 = WanTransformer3DModel.from_pretrained( - transformer_path_2, - subfolder=subfolder_2, - torch_dtype=dtype, - ).to(dtype=dtype) + transformer_2 = WanTransformer3DModel.load_model( + transformer_path_2, dtype=dtype, subfolder=subfolder_2 + ) flush() diff --git a/extensions_built_in/diffusion_models/z_image/z_image.py b/extensions_built_in/diffusion_models/z_image/z_image.py index 324b76a..67ffd90 100644 --- a/extensions_built_in/diffusion_models/z_image/z_image.py +++ b/extensions_built_in/diffusion_models/z_image/z_image.py @@ -13,19 +13,16 @@ from toolkit.samplers.custom_flowmatch_sampler import ( CustomFlowMatchEulerDiscreteScheduler, ) from toolkit.accelerator import unwrap_model -from optimum.quanto import freeze from toolkit.util.quantize import ( - quantize, - get_qtype, quantize_model, dequantize_if_quantized, ) from toolkit.memory_management import MemoryManager from toolkit.metadata import get_meta_for_safetensors +from toolkit.models.v2.text_encoders.qwen3 import Qwen3TextEncoder +from toolkit.models.v2.vae.autoencoder_kl import KLVAE from safetensors.torch import load_file, save_file -from transformers import AutoTokenizer, Qwen3ForCausalLM -from diffusers import AutoencoderKL try: from diffusers import ZImagePipeline @@ -255,36 +252,12 @@ class ZImageModel(BaseModel): flush() self.print_and_status_update("Text Encoder") - tokenizer = AutoTokenizer.from_pretrained( - base_model_path, subfolder="tokenizer", torch_dtype=dtype - ) - text_encoder = Qwen3ForCausalLM.from_pretrained( - base_model_path, subfolder="text_encoder", torch_dtype=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: - self.print_and_status_update("Quantizing Text Encoder") - quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) - freeze(text_encoder) - flush() + tokenizer = Qwen3TextEncoder.load_tokenizer(base_model_path) + text_encoder = Qwen3TextEncoder.load_model(base_model_path, dtype=dtype) + self.prepare_text_encoder(text_encoder, dtype=dtype) self.print_and_status_update("Loading VAE") - vae = AutoencoderKL.from_pretrained( - base_model_path, subfolder="vae", torch_dtype=dtype - ) + vae = KLVAE.load_model(base_model_path, dtype=dtype) self.noise_scheduler = ZImageModel.get_train_scheduler() diff --git a/extensions_built_in/diffusion_models/z_image/z_image_l2p_model.py b/extensions_built_in/diffusion_models/z_image/z_image_l2p_model.py index 34fd183..46ca0d0 100644 --- a/extensions_built_in/diffusion_models/z_image/z_image_l2p_model.py +++ b/extensions_built_in/diffusion_models/z_image/z_image_l2p_model.py @@ -9,11 +9,10 @@ import torch.nn.functional as F import yaml from toolkit.basic import flush from toolkit.accelerator import unwrap_model -from optimum.quanto import freeze -from toolkit.util.quantize import quantize, get_qtype, quantize_model +from toolkit.util.quantize import quantize_model from toolkit.memory_management import MemoryManager -from transformers import AutoTokenizer, Qwen3ForCausalLM +from toolkit.models.v2.text_encoders.qwen3 import Qwen3TextEncoder from toolkit.models.FakeVAE import FakeVAE from toolkit.paths import MODELS_PATH from safetensors.torch import load_file, save_file @@ -468,31 +467,9 @@ class ZImageL2PModel(ZImageModel): flush() self.print_and_status_update("Text Encoder") - tokenizer = AutoTokenizer.from_pretrained( - base_model_path, subfolder="tokenizer", torch_dtype=dtype - ) - text_encoder = Qwen3ForCausalLM.from_pretrained( - base_model_path, subfolder="text_encoder", torch_dtype=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: - self.print_and_status_update("Quantizing Text Encoder") - quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) - freeze(text_encoder) - flush() + tokenizer = Qwen3TextEncoder.load_tokenizer(base_model_path) + text_encoder = Qwen3TextEncoder.load_model(base_model_path, dtype=dtype) + self.prepare_text_encoder(text_encoder, dtype=dtype) self.print_and_status_update("Loading VAE") # vae = AutoencoderKL.from_pretrained( diff --git a/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_model.py b/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_model.py index a4f3c0c..b4a3bb3 100644 --- a/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_model.py +++ b/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_model.py @@ -13,14 +13,13 @@ from toolkit.samplers.custom_flowmatch_sampler import ( CustomFlowMatchEulerDiscreteScheduler, ) from toolkit.accelerator import unwrap_model -from optimum.quanto import freeze -from toolkit.util.quantize import quantize, get_qtype, quantize_model +from toolkit.util.quantize import quantize_model from toolkit.memory_management import MemoryManager from safetensors.torch import load_file from optimum.quanto import QTensor from toolkit.metadata import get_meta_for_safetensors from safetensors.torch import load_file, save_file -from transformers import AutoTokenizer, Qwen3ForCausalLM +from toolkit.models.v2.text_encoders.qwen3 import Qwen3TextEncoder from diffusers import AutoencoderKL from toolkit.models.FakeVAE import FakeVAE from .zeta_chroma_transformer import ZImageDCT, ZImageDCTParams, vae_flatten, vae_unflatten, prepare_latent_image_ids, make_text_position_ids @@ -142,31 +141,9 @@ class ZetaChromaModel(BaseModel): flush() self.print_and_status_update("Text Encoder") - tokenizer = AutoTokenizer.from_pretrained( - base_model_path, subfolder="tokenizer", torch_dtype=dtype - ) - text_encoder = Qwen3ForCausalLM.from_pretrained( - base_model_path, subfolder="text_encoder", torch_dtype=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: - self.print_and_status_update("Quantizing Text Encoder") - quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) - freeze(text_encoder) - flush() + tokenizer = Qwen3TextEncoder.load_tokenizer(base_model_path) + text_encoder = Qwen3TextEncoder.load_model(base_model_path, dtype=dtype) + self.prepare_text_encoder(text_encoder, dtype=dtype) self.print_and_status_update("Loading VAE") vae = FakeVAE(scaling_factor=1.0) diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index ba4a1d7..8f06dfb 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -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: diff --git a/toolkit/models/loaders/umt5.py b/toolkit/models/loaders/umt5.py index f53db3c..3c7fd30 100644 --- a/toolkit/models/loaders/umt5.py +++ b/toolkit/models/loaders/umt5.py @@ -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="", - unk_token="", - pad_token="", - _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 diff --git a/toolkit/models/v2/PLANNING.md b/toolkit/models/v2/PLANNING.md index 796723f..87e0c9b 100644 --- a/toolkit/models/v2/PLANNING.md +++ b/toolkit/models/v2/PLANNING.md @@ -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()` diff --git a/toolkit/models/v2/_mixin.py b/toolkit/models/v2/_mixin.py index 4e82640..cacc49f 100644 --- a/toolkit/models/v2/_mixin.py +++ b/toolkit/models/v2/_mixin.py @@ -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) diff --git a/toolkit/models/v2/diffusion_models/flux.py b/toolkit/models/v2/diffusion_models/flux.py new file mode 100644 index 0000000..f2fb094 --- /dev/null +++ b/toolkit/models/v2/diffusion_models/flux.py @@ -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"] diff --git a/toolkit/models/v2/diffusion_models/hidream.py b/toolkit/models/v2/diffusion_models/hidream.py new file mode 100644 index 0000000..6ec8f3e --- /dev/null +++ b/toolkit/models/v2/diffusion_models/hidream.py @@ -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"] diff --git a/toolkit/models/v2/diffusion_models/nucleus_image.py b/toolkit/models/v2/diffusion_models/nucleus_image.py new file mode 100644 index 0000000..62092ce --- /dev/null +++ b/toolkit/models/v2/diffusion_models/nucleus_image.py @@ -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"] diff --git a/toolkit/models/v2/diffusion_models/qwen_image.py b/toolkit/models/v2/diffusion_models/qwen_image.py new file mode 100644 index 0000000..6d5b7d7 --- /dev/null +++ b/toolkit/models/v2/diffusion_models/qwen_image.py @@ -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 diff --git a/toolkit/models/v2/diffusion_models/wan.py b/toolkit/models/v2/diffusion_models/wan.py new file mode 100644 index 0000000..2ba9e88 --- /dev/null +++ b/toolkit/models/v2/diffusion_models/wan.py @@ -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"] diff --git a/toolkit/models/v2/text_encoders/clip.py b/toolkit/models/v2/text_encoders/clip.py new file mode 100644 index 0000000..6aac3bb --- /dev/null +++ b/toolkit/models/v2/text_encoders/clip.py @@ -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"] diff --git a/toolkit/models/v2/text_encoders/qwen25_vl.py b/toolkit/models/v2/text_encoders/qwen25_vl.py new file mode 100644 index 0000000..ba10fe3 --- /dev/null +++ b/toolkit/models/v2/text_encoders/qwen25_vl.py @@ -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 diff --git a/toolkit/models/v2/text_encoders/qwen3.py b/toolkit/models/v2/text_encoders/qwen3.py new file mode 100644 index 0000000..328712b --- /dev/null +++ b/toolkit/models/v2/text_encoders/qwen3.py @@ -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"] diff --git a/toolkit/models/v2/text_encoders/qwen3_vl.py b/toolkit/models/v2/text_encoders/qwen3_vl.py new file mode 100644 index 0000000..dd42e22 --- /dev/null +++ b/toolkit/models/v2/text_encoders/qwen3_vl.py @@ -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) diff --git a/toolkit/models/v2/text_encoders/t5.py b/toolkit/models/v2/text_encoders/t5.py new file mode 100644 index 0000000..1a5953e --- /dev/null +++ b/toolkit/models/v2/text_encoders/t5.py @@ -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"] diff --git a/toolkit/models/v2/text_encoders/umt5.py b/toolkit/models/v2/text_encoders/umt5.py new file mode 100644 index 0000000..d0b2e99 --- /dev/null +++ b/toolkit/models/v2/text_encoders/umt5.py @@ -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="", + unk_token="", + pad_token="", + _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 + ) diff --git a/toolkit/models/v2/vae/autoencoder_kl.py b/toolkit/models/v2/vae/autoencoder_kl.py new file mode 100644 index 0000000..0eb58fc --- /dev/null +++ b/toolkit/models/v2/vae/autoencoder_kl.py @@ -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" diff --git a/extensions_built_in/diffusion_models/ideogram4/src/vae.py b/toolkit/models/v2/vae/flux2_kl.py similarity index 83% rename from extensions_built_in/diffusion_models/ideogram4/src/vae.py rename to toolkit/models/v2/vae/flux2_kl.py index d03a2e2..e7f6b00 100644 --- a/extensions_built_in/diffusion_models/ideogram4/src/vae.py +++ b/toolkit/models/v2/vae/flux2_kl.py @@ -1,4 +1,13 @@ -"""Flux2 KL autoencoder.""" +"""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 @@ -22,6 +31,17 @@ class AutoEncoderParams: 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) @@ -194,30 +214,30 @@ class Encoder(nn.Module): 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]) # type: ignore[index, operator] - if len(self.down[i_level].attn) > 0: # type: ignore[arg-type] - h = ckpt.checkpoint(self.down[i_level].attn[i_block], h) # type: ignore[index, operator] + 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]) # type: ignore[index, operator] - if len(self.down[i_level].attn) > 0: # type: ignore[arg-type] - h = self.down[i_level].attn[i_block](h) # type: ignore[index, operator] + 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])) # type: ignore[operator] + hs.append(ckpt.checkpoint(self.down[i_level].downsample, hs[-1])) else: - hs.append(self.down[i_level].downsample(hs[-1])) # type: ignore[operator] + 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) # type: ignore[operator] - h = ckpt.checkpoint(self.mid.attn_1, h) # type: ignore[operator] - h = ckpt.checkpoint(self.mid.block_2, h) # type: ignore[operator] + 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) # type: ignore[operator] - h = self.mid.attn_1(h) # type: ignore[operator] - h = self.mid.block_2(h) # type: ignore[operator] + 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) @@ -300,13 +320,13 @@ class Decoder(nn.Module): # middle if torch.is_grad_enabled() and self.gradient_checkpointing: - h = ckpt.checkpoint(self.mid.block_1, h) # type: ignore[operator] - h = ckpt.checkpoint(self.mid.attn_1, h) # type: ignore[operator] - h = ckpt.checkpoint(self.mid.block_2, h) # type: ignore[operator] + 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) # type: ignore[operator] - h = self.mid.attn_1(h) # type: ignore[operator] - h = self.mid.block_2(h) # type: ignore[operator] + 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) @@ -314,18 +334,18 @@ class Decoder(nn.Module): 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) # type: ignore[index, operator] - if len(self.up[i_level].attn) > 0: # type: ignore[arg-type] - h = ckpt.checkpoint(self.up[i_level].attn[i_block], h) # type: ignore[index, operator] + 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) # type: ignore[index, operator] - if len(self.up[i_level].attn) > 0: # type: ignore[arg-type] - h = self.up[i_level].attn[i_block](h) # type: ignore[index, operator] + 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) # type: ignore[operator] + h = ckpt.checkpoint(self.up[i_level].upsample, h) else: - h = self.up[i_level].upsample(h) # type: ignore[operator] + h = self.up[i_level].upsample(h) # end h = self.norm_out(h) @@ -346,10 +366,13 @@ class AutoEncoder(nn.Module): 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=params.ch, + ch=decoder_ch, out_ch=params.out_ch, ch_mult=params.ch_mult, num_res_blocks=params.num_res_blocks, @@ -369,7 +392,7 @@ class AutoEncoder(nn.Module): self._gradient_checkpointing = False @property - def gradient_checkpointing(self) -> bool: + def gradient_checkpointing(self): return self._gradient_checkpointing @gradient_checkpointing.setter @@ -378,18 +401,52 @@ class AutoEncoder(nn.Module): 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() - @property - def device(self) -> torch.device: - return next(self.parameters()).device + def normalize(self, z): + self.bn.eval() + return self.bn(z) - @property - def dtype(self) -> torch.dtype: - return next(self.parameters()).dtype + 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 diff --git a/toolkit/models/v2/vae/qwen_image.py b/toolkit/models/v2/vae/qwen_image.py new file mode 100644 index 0000000..af6bc59 --- /dev/null +++ b/toolkit/models/v2/vae/qwen_image.py @@ -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) diff --git a/toolkit/models/wan21/wan21.py b/toolkit/models/wan21/wan21.py index 3f7f953..1fec78d 100644 --- a/toolkit/models/wan21/wan21.py +++ b/toolkit/models/wan21/wan21.py @@ -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")