diff --git a/extensions_built_in/diffusion_models/anima/anima.py b/extensions_built_in/diffusion_models/anima/anima.py index 9d3772b..b22089a 100644 --- a/extensions_built_in/diffusion_models/anima/anima.py +++ b/extensions_built_in/diffusion_models/anima/anima.py @@ -271,7 +271,7 @@ class AnimaModel(BaseModel): flush() self.print_and_status_update("Quantizing Text Conditioner") - quantize(text_conditioner, weights=get_qtype(self.model_config.qtype)) + quantize(text_conditioner, weights=get_qtype(self.model_config.qtype_te)) freeze(text_conditioner) flush() diff --git a/extensions_built_in/diffusion_models/chroma/chroma_model.py b/extensions_built_in/diffusion_models/chroma/chroma_model.py index d9d2025..6a3aa9a 100644 --- a/extensions_built_in/diffusion_models/chroma/chroma_model.py +++ b/extensions_built_in/diffusion_models/chroma/chroma_model.py @@ -11,10 +11,9 @@ from toolkit.basic import flush # from toolkit.pixel_shuffle_encoder import AutoencoderPixelMixer 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 optimum.quanto import QTensor +from toolkit.util.quantize import quantize, quantize_model from .pipeline import ChromaPipeline, prepare_latent_image_ids from einops import rearrange, repeat import random @@ -67,6 +66,9 @@ class FakeCLIP(torch.nn.Module): class ChromaModel(BaseModel): arch = "chroma" + def get_transformer_block_names(self): + return ["double_blocks", "single_blocks"] + def __init__( self, device, @@ -163,13 +165,10 @@ class ChromaModel(BaseModel): transformer.config.num_single_layers = transformer.params.depth_single_blocks if self.model_config.quantize: - # patch the state dict method - patch_dequantization_on_save(transformer) - quantization_type = get_qtype(self.model_config.qtype) + # block-streaming quantize (handles dequant-on-save patching, + # excludes, ARA, and quantize_kwargs) self.print_and_status_update("Quantizing transformer") - quantize(transformer, weights=quantization_type, - **self.model_config.quantize_kwargs) - freeze(transformer) + quantize_model(self, transformer) transformer.to(self.device_torch) else: transformer.to(self.device_torch, dtype=dtype) @@ -389,19 +388,16 @@ class ChromaModel(BaseModel): return self.text_encoder[1].encoder.block[0].layer[0].SelfAttention.q.weight.requires_grad def save_model(self, output_path, meta, save_dtype): + # comfy-format single-file save via the mixin (chroma's class keys ARE + # the original layout); handles torchao/Ostris dequant, not just quanto if not output_path.endswith(".safetensors"): - output_path = output_path + ".safetensors" - # only save the unet + output_path = output_path + ".safetensors" transformer: Chroma = unwrap_model(self.model) - state_dict = transformer.state_dict() - save_dict = {} - for k, v in state_dict.items(): - if isinstance(v, QTensor): - v = v.dequantize() - save_dict[k] = v.clone().to('cpu', dtype=save_dtype) - - meta = get_meta_for_safetensors(meta, name='chroma') - save_file(save_dict, output_path, metadata=meta) + transformer.save_model( + output_path, + dtype=save_dtype, + metadata=get_meta_for_safetensors(meta, name="chroma"), + ) def get_loss_target(self, *args, **kwargs): noise = kwargs.get('noise') 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 d95dddc..2fcc390 100644 --- a/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py +++ b/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py @@ -10,10 +10,9 @@ from toolkit.basic import flush # from toolkit.pixel_shuffle_encoder import AutoencoderPixelMixer 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 optimum.quanto import QTensor +from toolkit.util.quantize import quantize, quantize_model from .pipeline import ChromaPipeline, prepare_latent_image_ids from einops import rearrange, repeat import random @@ -37,36 +36,16 @@ scheduler_config = { "use_dynamic_shifting": True } -class FakeConfig: - # for diffusers compatability - def __init__(self): - self.attention_head_dim = 128 - self.guidance_embeds = True - self.in_channels = 64 - self.joint_attention_dim = 4096 - self.num_attention_heads = 24 - self.num_layers = 19 - self.num_single_layers = 38 - self.patch_size = 1 - -class FakeCLIP(torch.nn.Module): - def __init__(self, device='cuda'): - super().__init__() - self.dtype = torch.bfloat16 - # the pipeline derives its execution device from this attribute; - # nn.Module.to() does not update it - self.device = device - self.text_model = None - self.tokenizer = None - self.model_max_length = 77 - - def forward(self, *args, **kwargs): - return torch.zeros(1, 1, 1).to(self.device) +# shared with the base chroma model (identical stubs) +from .chroma_model import FakeCLIP, FakeConfig class ChromaRadianceModel(BaseModel): arch = "chroma_radiance" + def get_transformer_block_names(self): + return ["double_blocks", "single_blocks"] + def __init__( self, device, @@ -165,13 +144,10 @@ class ChromaRadianceModel(BaseModel): transformer.config.num_single_layers = transformer.params.depth_single_blocks if self.model_config.quantize: - # patch the state dict method - patch_dequantization_on_save(transformer) - quantization_type = get_qtype(self.model_config.qtype) + # block-streaming quantize (handles dequant-on-save patching, + # excludes, ARA, and quantize_kwargs) self.print_and_status_update("Quantizing transformer") - quantize(transformer, weights=quantization_type, - **self.model_config.quantize_kwargs) - freeze(transformer) + quantize_model(self, transformer) transformer.to(self.device_torch) else: transformer.to(self.device_torch, dtype=dtype) @@ -373,19 +349,16 @@ class ChromaRadianceModel(BaseModel): return False def save_model(self, output_path, meta, save_dtype): + # comfy-format single-file save via the mixin (chroma's class keys ARE + # the original layout); handles torchao/Ostris dequant, not just quanto if not output_path.endswith(".safetensors"): - output_path = output_path + ".safetensors" - # only save the unet + output_path = output_path + ".safetensors" transformer: Chroma = unwrap_model(self.model) - state_dict = transformer.state_dict() - save_dict = {} - for k, v in state_dict.items(): - if isinstance(v, QTensor): - v = v.dequantize() - save_dict[k] = v.clone().to('cpu', dtype=save_dtype) - - meta = get_meta_for_safetensors(meta, name='chroma') - save_file(save_dict, output_path, metadata=meta) + transformer.save_model( + output_path, + dtype=save_dtype, + metadata=get_meta_for_safetensors(meta, name="chroma"), + ) def get_loss_target(self, *args, **kwargs): noise = kwargs.get('noise') 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 2f3795c..018b97b 100644 --- a/extensions_built_in/diffusion_models/f_light/f_light.py +++ b/extensions_built_in/diffusion_models/f_light/f_light.py @@ -11,10 +11,9 @@ from toolkit.models.v2.vae.autoencoder_kl import KLVAE from toolkit.basic import flush 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 optimum.quanto import QTensor +from toolkit.util.quantize import quantize, quantize_model from .src import FLitePipeline, DiT if TYPE_CHECKING: @@ -34,6 +33,9 @@ scheduler_config = { class FLiteModel(BaseModel): arch = "f-lite" + def get_transformer_block_names(self): + return ["blocks"] + def __init__( self, device, @@ -79,13 +81,10 @@ class FLiteModel(BaseModel): transformer.to(self.quantize_device, dtype=dtype) if self.model_config.quantize: - # patch the state dict method - patch_dequantization_on_save(transformer) - quantization_type = get_qtype(self.model_config.qtype) + # block-streaming quantize (handles dequant-on-save patching, + # excludes, ARA, and quantize_kwargs) self.print_and_status_update("Quantizing transformer") - quantize(transformer, weights=quantization_type, - **self.model_config.quantize_kwargs) - freeze(transformer) + quantize_model(self, transformer) transformer.to(self.device_torch) else: transformer.to(self.device_torch, dtype=dtype) diff --git a/extensions_built_in/diffusion_models/flux2/flux2_model.py b/extensions_built_in/diffusion_models/flux2/flux2_model.py index 20aedf0..b61f119 100644 --- a/extensions_built_in/diffusion_models/flux2/flux2_model.py +++ b/extensions_built_in/diffusion_models/flux2/flux2_model.py @@ -111,7 +111,7 @@ class Flux2Model(BaseModel): if self.model_config.quantize_te: self.print_and_status_update("Quantizing Mistral") - quantize(text_encoder, weights=get_qtype(self.model_config.qtype)) + quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) freeze(text_encoder) flush() @@ -136,9 +136,6 @@ class Flux2Model(BaseModel): transformer_path = model_path self.print_and_status_update("Loading transformer") - with torch.device("meta"): - transformer = Flux2(self.get_flux2_params()) - # use local path if provided if os.path.exists(os.path.join(transformer_path, self.flux2_te_filename)): transformer_path = os.path.join(transformer_path, self.flux2_te_filename) @@ -152,12 +149,9 @@ class Flux2Model(BaseModel): ) transformer_state_dict = load_file(transformer_path, device="cpu") - - # cast to dtype - for key in transformer_state_dict: - transformer_state_dict[key] = transformer_state_dict[key].to(dtype) - - transformer.load_state_dict(transformer_state_dict, assign=True) + transformer = Flux2.load_from_state_dict( + transformer_state_dict, dtype, config=self.get_flux2_params() + ) if self.model_config.quantize: # patch the state dict method diff --git a/extensions_built_in/diffusion_models/flux2/src/model.py b/extensions_built_in/diffusion_models/flux2/src/model.py index cffae9f..5c04503 100644 --- a/extensions_built_in/diffusion_models/flux2/src/model.py +++ b/extensions_built_in/diffusion_models/flux2/src/model.py @@ -1,4 +1,6 @@ import torch + +from toolkit.models.v2._mixin import OstrisModelMixin from einops import rearrange from torch import Tensor, nn import torch.utils.checkpoint as ckpt @@ -54,7 +56,7 @@ class FakeConfig: self.patch_size = 1 -class Flux2(nn.Module): +class Flux2(nn.Module, OstrisModelMixin): def __init__(self, params: Flux2Params): super().__init__() self.config = FakeConfig() 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 927477f..535e9fb 100644 --- a/extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py +++ b/extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py @@ -17,11 +17,10 @@ from toolkit.basic import flush from toolkit.prompt_utils import PromptEmbeds from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler from toolkit.models.flux import add_model_gpu_splitter_to_flux, bypass_flux_guidance, restore_flux_guidance -from toolkit.dequantize import patch_dequantization_on_save from toolkit.accelerator import get_accelerator, unwrap_model -from optimum.quanto import freeze, QTensor +from optimum.quanto import QTensor from toolkit.util.mask import generate_random_mask, random_dialate_mask -from toolkit.util.quantize import quantize, get_qtype +from toolkit.util.quantize import quantize, quantize_model from einops import rearrange, repeat import random import torch.nn.functional as F @@ -44,6 +43,9 @@ scheduler_config = { class FluxKontextModel(BaseModel): arch = "flux_kontext" + def get_transformer_block_names(self): + return ["transformer_blocks", "single_transformer_blocks"] + def __init__( self, device, @@ -95,13 +97,10 @@ class FluxKontextModel(BaseModel): transformer.to(self.quantize_device, dtype=dtype) if self.model_config.quantize: - # patch the state dict method - patch_dequantization_on_save(transformer) - quantization_type = get_qtype(self.model_config.qtype) + # block-streaming quantize (handles dequant-on-save patching, + # excludes, ARA, and quantize_kwargs) self.print_and_status_update("Quantizing transformer") - quantize(transformer, weights=quantization_type, - **self.model_config.quantize_kwargs) - freeze(transformer) + quantize_model(self, transformer) transformer.to(self.device_torch) else: transformer.to(self.device_torch, dtype=dtype) diff --git a/extensions_built_in/diffusion_models/ideogram4/ideogram4.py b/extensions_built_in/diffusion_models/ideogram4/ideogram4.py index 2048bdd..d84c6ab 100644 --- a/extensions_built_in/diffusion_models/ideogram4/ideogram4.py +++ b/extensions_built_in/diffusion_models/ideogram4/ideogram4.py @@ -238,9 +238,6 @@ class Ideogram4Model(BaseModel): self.print_and_status_update("Loading transformer") transformer_config = Ideogram4Config() - with torch.device("meta"): - transformer = Ideogram4Transformer2DModel(transformer_config) - self.print_and_status_update(" - fetching transformer weights") state_dict = _load_component_state_dict( base, "transformer", "diffusion_pytorch_model" @@ -250,7 +247,9 @@ class Ideogram4Model(BaseModel): state_dict, dtype, self.device_torch, self.model_config.low_vram ) self.print_and_status_update(" - loading transformer state dict") - transformer.load_state_dict(state_dict, assign=True) + transformer = Ideogram4Transformer2DModel.load_from_state_dict( + state_dict, dtype, config=transformer_config + ) del state_dict flush() diff --git a/extensions_built_in/diffusion_models/ideogram4/src/transformer.py b/extensions_built_in/diffusion_models/ideogram4/src/transformer.py index 3173bda..4abc5c8 100644 --- a/extensions_built_in/diffusion_models/ideogram4/src/transformer.py +++ b/extensions_built_in/diffusion_models/ideogram4/src/transformer.py @@ -12,6 +12,8 @@ import math from dataclasses import dataclass import torch + +from toolkit.models.v2._mixin import OstrisModelMixin import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint @@ -355,7 +357,7 @@ class Ideogram4FinalLayer(nn.Module): return self.linear(self.norm_final(x) * scale) -class Ideogram4Transformer2DModel(nn.Module): +class Ideogram4Transformer2DModel(nn.Module, OstrisModelMixin): """Ideogram 4 flow-matching transformer.""" def __init__(self, config: Ideogram4Config) -> None: diff --git a/extensions_built_in/diffusion_models/krea2/krea2.py b/extensions_built_in/diffusion_models/krea2/krea2.py index b0fa0a6..90462d9 100644 --- a/extensions_built_in/diffusion_models/krea2/krea2.py +++ b/extensions_built_in/diffusion_models/krea2/krea2.py @@ -217,21 +217,15 @@ class Krea2Model(QwenImageVAEHolderMixin, BaseModel): mmdit_kwargs.update(self.model_config.model_kwargs.get("mmdit_config", {})) config = SingleMMDiTConfig(**mmdit_kwargs) - # Build on meta, then materialize straight from the checkpoint. - with torch.device("meta"): - transformer = SingleStreamDiT(config) - self.print_and_status_update(" - fetching transformer weights") state_dict = _load_mmdit_state_dict( self.model_config.name_or_path, self.model_config.model_kwargs.get("checkpoint_filename", None), ) - state_dict = { - k: (v.to(dtype) if v.is_floating_point() else v) - for k, v in state_dict.items() - } self.print_and_status_update(" - loading transformer state dict") - transformer.load_state_dict(state_dict, strict=True, assign=True) + transformer = SingleStreamDiT.load_from_state_dict( + state_dict, dtype, config=config + ) del state_dict flush() return transformer diff --git a/extensions_built_in/diffusion_models/krea2/src/mmdit.py b/extensions_built_in/diffusion_models/krea2/src/mmdit.py index bb6619a..5e7cfcb 100644 --- a/extensions_built_in/diffusion_models/krea2/src/mmdit.py +++ b/extensions_built_in/diffusion_models/krea2/src/mmdit.py @@ -20,6 +20,8 @@ import math from dataclasses import dataclass import torch + +from toolkit.models.v2._mixin import OstrisModelMixin import torch.nn as nn import torch.nn.functional as F from einops import rearrange @@ -408,7 +410,7 @@ class SingleStreamBlock(nn.Module): return x -class SingleStreamDiT(nn.Module): +class SingleStreamDiT(nn.Module, OstrisModelMixin): def __init__(self, config: SingleMMDiTConfig): super().__init__() self.config = config diff --git a/extensions_built_in/diffusion_models/ltx2/ltx2.py b/extensions_built_in/diffusion_models/ltx2/ltx2.py index 2821e23..ddd6f3f 100644 --- a/extensions_built_in/diffusion_models/ltx2/ltx2.py +++ b/extensions_built_in/diffusion_models/ltx2/ltx2.py @@ -1279,7 +1279,7 @@ class LTX25Model(LTX2Model): from the state dict. Works unchanged for bf16 checkpoints, where no quant markers exist and everything strict-loads.""" from toolkit.util.comfy_quant_import import import_comfy_quantized_layers - from toolkit.util.ostris_quant import OstrisLinear + from toolkit.models.v2._mixin import OstrisModelMixin state_dict, num_quantized = import_comfy_quantized_layers( module, state_dict, orig_dtype=self.torch_dtype @@ -1288,32 +1288,8 @@ class LTX25Model(LTX2Model): self.print_and_status_update( f" - attached {num_quantized} pre-quantized ConvRot layers to {name}" ) - result = module.load_state_dict(state_dict, assign=True, strict=False) - # quantized linears hold their weight as backend buffers and had their - # bias assigned by the importer, so both report as "missing" here - quantized_param_keys = set() - for mod_name, m in module.named_modules(): - if isinstance(m, OstrisLinear): - quantized_param_keys.add(f"{mod_name}.weight") - if m.bias is not None: - quantized_param_keys.add(f"{mod_name}.bias") - bad_missing = [k for k in result.missing_keys if k not in quantized_param_keys] - if bad_missing or result.unexpected_keys: - raise ValueError( - f"LTX-2.5 {name} load mismatch: missing {bad_missing[:8]}, " - f"unexpected {result.unexpected_keys[:8]}" - ) - # nothing may be left on the meta device (e.g. a bias the importer - # should have filled) - leftover_meta = [ - param_name - for param_name, p in module.named_parameters() - if p.is_meta - ] - if leftover_meta: - raise ValueError( - f"LTX-2.5 {name} load left meta parameters: {leftover_meta[:8]}" - ) + # whitelist for quantized weights + leftover-meta check + OstrisModelMixin._load_state_dict_with_quantized(module, state_dict) return num_quantized def _load_gemma4_text_encoder(self, te_path: str, te_state_dict: dict): diff --git a/extensions_built_in/diffusion_models/mageflow/mageflow.py b/extensions_built_in/diffusion_models/mageflow/mageflow.py index 49b3530..1f11192 100644 --- a/extensions_built_in/diffusion_models/mageflow/mageflow.py +++ b/extensions_built_in/diffusion_models/mageflow/mageflow.py @@ -178,22 +178,13 @@ class MageFlowModel(BaseModel): structure.update(self.model_config.model_kwargs.get("transformer_config", {})) params = MageFlowParams(**structure) - # Build on meta, then materialize straight from the checkpoint. - with torch.device("meta"): - transformer = MageFlow(params) - self.print_and_status_update(" - fetching transformer weights") state_dict = load_file( self._get_model_file("transformer/diffusion_pytorch_model.safetensors") ) - state_dict = { - k: (v.to(dtype) if v.is_floating_point() else v) - for k, v in state_dict.items() - } self.print_and_status_update(" - loading transformer state dict") - transformer.load_state_dict(state_dict, strict=True, assign=True) - # The RoPE tables are plain tensors (not buffers/params), so the meta - # init left them unmaterialized — rebuild them for real. + transformer = MageFlow.load_from_state_dict(state_dict, dtype, config=params) + # rebuild the plain-tensor RoPE tables at full precision transformer.reset_rope() del state_dict flush() diff --git a/extensions_built_in/diffusion_models/mageflow/src/transformer.py b/extensions_built_in/diffusion_models/mageflow/src/transformer.py index 6b0d247..7ce1a76 100644 --- a/extensions_built_in/diffusion_models/mageflow/src/transformer.py +++ b/extensions_built_in/diffusion_models/mageflow/src/transformer.py @@ -28,6 +28,8 @@ from dataclasses import dataclass from typing import Any import torch + +from toolkit.models.v2._mixin import OstrisModelMixin import torch.nn as nn from torch import Tensor from torch.utils.checkpoint import checkpoint @@ -681,7 +683,7 @@ class MageFlowParams: patch_size: int = 1 -class MageFlow(nn.Module): +class MageFlow(nn.Module, OstrisModelMixin): def __init__(self, params: MageFlowParams): super().__init__() self.params = params diff --git a/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py index ab0d7fa..05a0d4a 100644 --- a/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py +++ b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py @@ -373,46 +373,12 @@ class MinimaxH3Model(BaseModel): self.invert_assistant_lora = False def _load_transformer(self) -> MiniMaxH3Transformer: - dtype = self.torch_dtype dit_path = self._resolve_comfy_file(self._dit_component()) self.print_and_status_update(f"Loading transformer from {dit_path}") - state_dict = load_file(dit_path) - - params = MiniMaxH3TransformerParams() - table = state_dict.get("adaln_t_table", None) - if table is not None: - # pruned checkpoint: factored timestep table instead of the MLP - params.adaln_t_table_size = table.shape[0] - params.time_embed_dim = table.shape[1] - - with torch.device("meta"): - transformer = MiniMaxH3Transformer(params) - - # attach the pre-quantized (int8 ConvRot) linears onto the toolkit's - # quantization backends; the rest loads at its stored precision (the - # checkpoint's bf16/fp16/fp32 mix is deliberate) - state_dict, num_quantized = import_comfy_quantized_layers( - transformer, state_dict, orig_dtype=dtype - ) - if num_quantized: - self.print_and_status_update( - f" - attached {num_quantized} pre-quantized ConvRot layers" - ) - result = transformer.load_state_dict(state_dict, assign=True, strict=False) - quantized_weight_keys = { - f"{name}.weight" - for name, m in transformer.named_modules() - if isinstance(m, OstrisLinear) - } - bad_missing = [k for k in result.missing_keys if k not in quantized_weight_keys] - if bad_missing or result.unexpected_keys: - raise ValueError( - f"MiniMax-H3 transformer load mismatch: missing {bad_missing[:8]}, " - f"unexpected {result.unexpected_keys[:8]}" - ) - del state_dict - flush() - return transformer + # the mixin single-file path: config sniffed from the checkpoint + # (adaln_t_table), pre-quantized ConvRot linears attached, everything + # else at its stored precision (the bf16/fp16/fp32 mix is deliberate) + return MiniMaxH3Transformer.load_model(dit_path, dtype=self.torch_dtype) def _load_text_encoder(self): from accelerate import init_empty_weights diff --git a/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py b/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py index f958268..809ff03 100644 --- a/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py +++ b/extensions_built_in/diffusion_models/minimax_h3/src/transformer.py @@ -30,6 +30,8 @@ import math from dataclasses import dataclass from typing import Optional, Tuple +from toolkit.models.v2._mixin import OstrisModelMixin + import torch import torch.nn.functional as F from torch import nn @@ -356,7 +358,31 @@ class MiniMaxH3FinalLayer(nn.Module): return self.video_out(h), self.audio_out(h) -class MiniMaxH3Transformer(nn.Module): +class MiniMaxH3Transformer(nn.Module, OstrisModelMixin): + # comfy checkpoints carry a deliberate bf16/fp16/fp32 mix + aitk_cast_on_load = False + + @classmethod + def aitk_config_from_state_dict(cls, state_dict): + params = MiniMaxH3TransformerParams() + table = state_dict.get("adaln_t_table", None) + if table is not None: + # pruned checkpoint: factored timestep table instead of the MLP + params.adaln_t_table_size = table.shape[0] + params.time_embed_dim = table.shape[1] + return params + + @classmethod + def aitk_from_config(cls, config): + from accelerate import init_empty_weights + + with init_empty_weights(include_buffers=False): + return cls(config) + + @classmethod + def get_transformer_block_names(cls): + return ["blocks"] + def __init__(self, params: Optional[MiniMaxH3TransformerParams] = None): super().__init__() if params is None: diff --git a/extensions_built_in/diffusion_models/omnigen2/__init__.py b/extensions_built_in/diffusion_models/omnigen2/__init__.py index 77741be..09b00f9 100644 --- a/extensions_built_in/diffusion_models/omnigen2/__init__.py +++ b/extensions_built_in/diffusion_models/omnigen2/__init__.py @@ -13,7 +13,7 @@ from toolkit.samplers.custom_flowmatch_sampler import ( ) from toolkit.accelerator import unwrap_model from optimum.quanto import freeze -from toolkit.util.quantize import quantize, get_qtype +from toolkit.util.quantize import quantize, get_qtype, quantize_model from .src.pipelines.omnigen2.pipeline_omnigen2 import OmniGen2Pipeline from .src.models.transformers import OmniGen2Transformer2DModel from .src.models.transformers.repo import OmniGen2RotaryPosEmbed @@ -105,9 +105,9 @@ class OmniGen2Model(BaseModel): if self.model_config.quantize: self.print_and_status_update("Quantizing transformer") - quantization_type = get_qtype(self.model_config.qtype) - quantize(transformer, weights=quantization_type) - freeze(transformer) + quantize_model(self, transformer) + if not self.low_vram: + transformer.to(self.device_torch) if self.low_vram: # unload it for now 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 3a26897..97ccd52 100644 --- a/extensions_built_in/diffusion_models/wan22/wan22_14b_model.py +++ b/extensions_built_in/diffusion_models/wan22/wan22_14b_model.py @@ -22,6 +22,7 @@ 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 +from toolkit.metadata import get_meta_for_safetensors from toolkit.models.wan21.wan21 import Wan21 from .wan22_5b_model import ( scheduler_config, @@ -431,20 +432,20 @@ class Wan2214bModel(Wan21): return False def save_model(self, output_path, meta, save_dtype): + # comfy-format single-file saves, one per DiT (comfy convention: + # separate high/low noise files) transformer_combo: DualWanTransformer3DModel = unwrap_model(self.model) - transformer_combo.transformer_1.save_pretrained( - save_directory=os.path.join(output_path, "transformer"), - safe_serialization=True, + base = output_path + if base.endswith(".safetensors"): + base = base[: -len(".safetensors")] + metadata = get_meta_for_safetensors(meta, name=self.arch) + transformer_combo.transformer_1.save_model( + f"{base}_high_noise.safetensors", dtype=save_dtype, metadata=metadata ) - transformer_combo.transformer_2.save_pretrained( - save_directory=os.path.join(output_path, "transformer_2"), - safe_serialization=True, + transformer_combo.transformer_2.save_model( + f"{base}_low_noise.safetensors", dtype=save_dtype, metadata=metadata ) - meta_path = os.path.join(output_path, "aitk_meta.yaml") - with open(meta_path, "w") as f: - yaml.dump(meta, f) - def save_lora( self, state_dict: Dict[str, torch.Tensor], 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 b4a3bb3..d9af932 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 @@ -92,10 +92,6 @@ class ZetaChromaModel(BaseModel): transformer_state_dict = load_file(transformer_path, device="cpu") - # cast to dtype - for key in transformer_state_dict: - transformer_state_dict[key] = transformer_state_dict[key].to(dtype) - # Auto-detect use_x0 from checkpoint use_x0 = "__x0__" in transformer_state_dict @@ -107,10 +103,9 @@ class ZetaChromaModel(BaseModel): use_x0=use_x0, ) - with torch.device("meta"): - transformer = ZImageDCT(model_params) - - transformer.load_state_dict(transformer_state_dict, assign=True) + transformer = ZImageDCT.load_from_state_dict( + transformer_state_dict, dtype, config=model_params + ) del transformer_state_dict transformer.to(self.quantize_device, dtype=dtype) diff --git a/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_transformer.py b/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_transformer.py index e2db51d..03995b1 100644 --- a/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_transformer.py +++ b/extensions_built_in/diffusion_models/zeta_chroma/zeta_chroma_transformer.py @@ -6,6 +6,8 @@ from typing import List, Optional import math import torch + +from toolkit.models.v2._mixin import OstrisModelMixin import torch.nn as nn import torch.nn.functional as F from torch import Tensor @@ -449,7 +451,7 @@ class SimpleMLPAdaLN(nn.Module): return self.final_layer(x) -class ZImageDCT(nn.Module): +class ZImageDCT(nn.Module, OstrisModelMixin): def __init__(self, params: ZImageDCTParams): super().__init__() self.config = FakeConfig() diff --git a/toolkit/models/v2/PLANNING.md b/toolkit/models/v2/PLANNING.md index 7b17bef..74096aa 100644 --- a/toolkit/models/v2/PLANNING.md +++ b/toolkit/models/v2/PLANNING.md @@ -207,8 +207,11 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord variants inherit - [x] nucleus_image — `v2/diffusion_models/nucleus_image.py`, TE stanza collapsed to prepare_text_encoder -- [ ] krea2, ideogram4, mageflow — TE/VAE migrated; their custom local DiT - classes still to be rebased onto the mixin +- [x] krea2, ideogram4, mageflow — DiTs on the mixin via the `config=` + passthrough (holder builds config from model_kwargs/holder state, mixin + does build/markers/whitelist/casting; plain nn.Module classes now work + with the default builder). krea2 + ideogram4 verified by harness; + mageflow untestable while its repo 404s - [x] chroma, chroma_radiance — both vendored Chroma classes now carry `OstrisModelMixin` with the block-count sniff moved into a new `aitk_config_from_state_dict` hook (mixin now supports checkpoint-derived @@ -217,14 +220,23 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord depends on holder state (patch_size), not the checkpoint - [x] flux_kontext — `v2/diffusion_models/flux.py` (FluxTransformer2DModel); whole model now loads through v2 (transformer, T5, CLIP, KLVAE) -- [ ] flux2 — TE/VAE partially migrated (flux2_kl); custom DiT still local. - krea2/mageflow/ideogram4/zeta_chroma DiTs stay model-specific: their - configs come from model_kwargs / holder state, so the mixin adds nothing - until the comfy-weights flip (Phase 2) -- [ ] minimax_h3 (+ ref2va), ltx2 family — already on the shared resolver + - comfy_quant_import; the full mixin port waits for Phase 2, when the - mixin's single-file precision policy (stored-precision loading, fp32-key - protection) is settled to match their deliberate behavior +- [x] flux2 + zeta_chroma DiTs on the mixin via `config=` passthrough + (flux2_klein_4b verified by harness). Every arch's DiT now loads + through the mixin except: ltx2 family (one-file→two-modules split, + helper delegated), anima (diffusers modular pipeline), ace_step + (bundled single-file loader), z_image_l2p's local subclass, and the + grandfathered legacy stable_diffusion_model archs +- [x] minimax_h3 (+ ref2va) transformer ported to the mixin: config sniffed + from the checkpoint via aitk_config_from_state_dict (adaln_t_table), + marker attach + stored-precision load via the new + `aitk_cast_on_load = False` knob. Verified on the real pruned convrot + file: 200 ConvRot linears, pruned table detected, fp32/fp16/bf16 mix + preserved, no meta leftovers. Its TE stays custom (50-layer truncation + + key_map). ltx2.5's `_load_quantized_module` now delegates to the + mixin's whitelist/meta helper (~30 lines deleted); its full port is + blocked on the one-comfy-file → transformer+connectors split, which + doesn't fit the per-class single-file shape — revisit with the live + server's component model - [x] wan21 / wan22 family — `v2/diffusion_models/wan.py` (WanTransformer3DModel, both wan22 dual loads included) + `v2/text_encoders/umt5.py` (UMT5TextEncoder + PatchedT5Tokenizer; @@ -307,9 +319,16 @@ Decisions: fp8 weight / scalar scale_weight, e.g. every wan *_fp8_scaled file): imports onto the float8 backend; scale_input (activation quant) is dropped, matmuls run dequantized. +- [x] wan comfy-format saves: `convert_state_dict_on_save` inverts diffusers' + rename table (base/t2v/i2v; vace/animate excluded — their reverse + mappings collide). Round-trip verified on both real comfy files + (exact; the 2.1 file's legacy model.diffusion_model. prefix drops per + the modern convention) and with real weights (load → save → 825-key + original-layout file → reload bit-equal). wan21 + wan22_5b save one + comfy file; wan22_14b saves the comfy-standard _high_noise/_low_noise + pair instead of two diffusers folders. - [ ] Wire remaining archs' candidate lists (chroma/others as their key - conversions are verified per file); wan comfy-format save needs the - inverse key mapping (defer with the other save flips) + conversions are verified per file) - [x] Fused-layout quantized attach for diffusers-split classes: `split_fused_quantized_keys` / `fuse_split_quantized_keys` (comfy_quant_import) do exact out-dim row surgery on quantized comfy @@ -333,6 +352,14 @@ Decisions: bit-exact fused qkv weights/scales/markers, and the reload's quantized forward is bit-identical — toolkit saves are byte-compatible with ComfyUI. +- [x] chroma + chroma_radiance saves flipped to the mixin (their class keys + ARE the original layout) — also fixes their quanto-only dequant bug + (torchao/Ostris weights now dequantize on save). Tiny-model round trip + verified. Save flips so far: z_image, qwen_image, wan21, wan22_5b, + wan22_14b (dual files), chroma ×2. +- [ ] flux_kontext comfy wiring deferred: its Comfy-Org repo ships a single + legacy-fp8 file in fused BFL layout — needs the flux fused-split + conversion (split_fused_quantized_keys pattern + BFL↔diffusers maps) - [ ] Flip the remaining per-arch `save_model` overrides as each arch's save-side key conversion is in place - [ ] Publish/verify comfy repacks per model as they flip @@ -354,6 +381,13 @@ Decisions: - [x] Missing weights skip rather than fail: default is HF_HUB_OFFLINE=1 and hub/file errors classify as SKIP; `--allow-download` opts into fetching. (GPU + local-weights test, not CI-portable.) +- [x] Final certification sweep (2026-08-27, post-polish): 14/14 runnable + archs PASS — comfy-source loads (zimage convrot8, qwen fp8, wan ×2 + + fp8 umt5 TE), all ported holder-config DiTs, and the migrated + quantize_model paths (chroma, flux_kontext, f_light block-streamed) in + one run; mageflow remains the upstream 404 skip. One regression caught + and fixed: qwen's _load_single_file override needed the new config + kwarg. - [x] Full sweep run 2026-08-27: 14/15 PASS (zimage, qwen_image, krea2, boogu_image, ernie_image, ideogram4, hidream_o1, anima, wan21, wan22_5b, chroma, flux_kontext, flux2_klein_4b, ltx2.3 — the quantized 22B ltx @@ -378,20 +412,28 @@ Decisions: ## TODO / look at later -- [ ] Quantize-path consolidation quirks: `quantize_kwargs` is honored only by the - raw `quantize()` call sites and silently dropped by `quantize_model()`; the - ARA path inside `quantize_model` hardcodes `uint8`. Decide the unified - behavior when consolidating. +- [x] Quantize consolidation: `quantize_model` now honors `quantize_kwargs` + (blocks + extras) and tolerates missing block names; the chroma ×2, + flux_kontext, f_light, omnigen2 raw-quantize sites migrated onto it + (gaining block streaming, excludes, ARA, dequant-on-save patching) with + holder block names added. Remaining raw sites are legacy/extension + (flex2, cogview4, stable_diffusion_model). The ARA uint8 hardcode + stands — revisit if a non-uint8 ARA base is ever wanted. +- [x] Last known `qtype_te` bugs fixed (flux2's Mistral TE, anima's + text_conditioner) — 9/9 sites from the survey now correct outside the + grandfathered legacy monolith (cogview4/legacy SD remain as-is). - [x] wan comfy-TE resolved for real: UMT5TextEncoder carries comfy candidates (fp8_e4m3fn_scaled via the legacy importer, fp16), files already in transformers key layout (spiece blob dropped, tied embed_tokens materialized). Verified: wan21 samples with the local comfy fp8 TE. The loaders/umt5.py `comfy_files` param stays as a no-op shim for old callers. -- [ ] Registry hardening: error (don't fall back to SD1) on unknown arch; lazy +- [x] Registry hardening: unknown archs now raise with the known-arch list + (legacy monolith archs whitelisted via LEGACY_ARCHS). Still open: lazy per-arch imports; single source of truth shared with the UI's `options.tsx` model list. -- [ ] Fake/stub components: consolidate on `toolkit/models/FakeVAE.py` / - `toolkit/unloader.py`, delete local copies. +- [x] Stub dedup where identical: chroma_radiance imports FakeCLIP/FakeConfig + from chroma_model. The other Fake* copies (hidream_o1, flux2, zeta) + carry model-specific values — left in place. - [ ] Vendored upstream code (hidream/src, omnigen2/src, ltx2 converter's private comfy-quant parser): dedupe against toolkit utils where practical. diff --git a/toolkit/models/v2/_mixin.py b/toolkit/models/v2/_mixin.py index 33b9fa9..e4d58de 100644 --- a/toolkit/models/v2/_mixin.py +++ b/toolkit/models/v2/_mixin.py @@ -67,6 +67,12 @@ class OstrisModelMixin: aitk_processor_repo: Optional[str] = None aitk_processor_subfolder: Optional[str] = None + # single-file loads without quant markers cast tensors to the requested + # dtype; classes whose checkpoints carry a deliberate precision mix (e.g. + # fp32 norms next to bf16 weights) set this False to load at stored + # precision instead + aitk_cast_on_load: bool = True + # ---- state set by the loader / quantizer ---- aitk_is_quantized: bool = False aitk_qtype: Optional[str] = None @@ -119,7 +125,10 @@ class OstrisModelMixin: from accelerate import init_empty_weights with init_empty_weights(include_buffers=False): - return cls.from_config(config) + if hasattr(cls, "from_config"): + return cls.from_config(config) + # plain nn.Module classes take their params/config object directly + return cls(config) # ------------------------------------------------------------------ # loading @@ -135,6 +144,7 @@ class OstrisModelMixin: quantize_device: Optional[torch.device] = None, exclude_quant_modules: Optional[List[str]] = None, config_path: Optional[str] = None, + config=None, subfolder: Optional[str] = None, use_comfy_weights: bool = True, **kwargs, @@ -180,7 +190,11 @@ class OstrisModelMixin: if name_or_path.endswith(".safetensors"): file_path = cls._resolve_single_file(name_or_path) model = cls._load_single_file( - file_path, dtype=dtype, config_path=config_path, subfolder=subfolder + file_path, + dtype=dtype, + config_path=config_path, + config=config, + subfolder=subfolder, ) else: if os.path.isdir(name_or_path): @@ -243,11 +257,16 @@ class OstrisModelMixin: file_path: str, dtype: torch.dtype, config_path: Optional[str] = None, + config=None, subfolder: Optional[str] = None, ): state_dict = load_file(file_path) return cls.load_from_state_dict( - state_dict, dtype, config_path=config_path, subfolder=subfolder + state_dict, + dtype, + config_path=config_path, + config=config, + subfolder=subfolder, ) @classmethod @@ -263,16 +282,20 @@ class OstrisModelMixin: state_dict: Dict[str, torch.Tensor], dtype: torch.dtype, config_path: Optional[str] = None, + config=None, subfolder: Optional[str] = None, ): """Build the model and load an already-read single-file state dict (the tail of the single-file path; also callable directly for - checkpoints read from non-safetensors sources).""" + checkpoints read from non-safetensors sources). ``config``, when + given, is used directly (for models whose config comes from the + holder, e.g. model_kwargs-driven archs).""" state_dict = cls.convert_state_dict_on_load(state_dict) has_quant_markers = "scaled_fp8" in state_dict or any( k.endswith(".comfy_quant") for k in state_dict ) - config = cls.aitk_config_from_state_dict(state_dict) + if config is None: + config = cls.aitk_config_from_state_dict(state_dict) if config is None: config = cls._load_single_file_config(config_path, subfolder) model = cls.aitk_from_config(config) @@ -288,11 +311,15 @@ class OstrisModelMixin: ) cls._load_state_dict_with_quantized(model, state_dict) model.aitk_is_quantized = True - else: + elif cls.aitk_cast_on_load: for key, value in state_dict.items(): - state_dict[key] = value.to(dtype=dtype) + if value.is_floating_point(): + state_dict[key] = value.to(dtype=dtype) model.load_state_dict(state_dict, assign=True) model.to(dtype=dtype) + else: + # stored-precision load (the checkpoint's dtype mix is deliberate) + model.load_state_dict(state_dict, assign=True) del state_dict flush() return model diff --git a/toolkit/models/v2/diffusion_models/qwen_image.py b/toolkit/models/v2/diffusion_models/qwen_image.py index 3c4e689..56e3f3a 100644 --- a/toolkit/models/v2/diffusion_models/qwen_image.py +++ b/toolkit/models/v2/diffusion_models/qwen_image.py @@ -29,7 +29,9 @@ class QwenImageTransformer2DModel( return ["transformer_blocks"] @classmethod - def _load_single_file(cls, file_path, dtype, config_path=None, subfolder=None): + def _load_single_file( + cls, file_path, dtype, config_path=None, config=None, subfolder=None + ): from safetensors import safe_open with safe_open(file_path, framework="pt") as f: @@ -38,7 +40,11 @@ class QwenImageTransformer2DModel( # comfy prequantized checkpoint (diffusers key layout): the mixin # path attaches the quantized layers return super()._load_single_file( - file_path, dtype, config_path=config_path, subfolder=subfolder + file_path, + dtype, + config_path=config_path, + config=config, + subfolder=subfolder, ) # other single-file checkpoints carry diffusers or original key # layouts; diffusers' single-file machinery owns that conversion diff --git a/toolkit/models/v2/diffusion_models/wan.py b/toolkit/models/v2/diffusion_models/wan.py index 98ddd1c..dfa5958 100644 --- a/toolkit/models/v2/diffusion_models/wan.py +++ b/toolkit/models/v2/diffusion_models/wan.py @@ -84,3 +84,54 @@ class WanTransformer3DModel(DiffusersWanTransformer3DModel, OstrisModelMixin): ) return convert_wan_transformer_to_diffusers(dict(state_dict)) + + # inverse of diffusers' rename table, for the base/t2v/i2v variants (no + # vace/animate: their extra mappings collide in reverse). Order matters: + # longest/most-specific first, and the norm2/norm3 swap uses a placeholder. + _SAVE_RENAMES = [ + ("condition_embedder.time_embedder.linear_1", "time_embedding.0"), + ("condition_embedder.time_embedder.linear_2", "time_embedding.2"), + ("condition_embedder.text_embedder.linear_1", "text_embedding.0"), + ("condition_embedder.text_embedder.linear_2", "text_embedding.2"), + ("condition_embedder.time_proj", "time_projection.1"), + ("condition_embedder.image_embedder.norm1", "img_emb.proj.0"), + ("condition_embedder.image_embedder.ff.net.0.proj", "img_emb.proj.1"), + ("condition_embedder.image_embedder.ff.net.2", "img_emb.proj.3"), + ("condition_embedder.image_embedder.norm2", "img_emb.proj.4"), + ("ffn.net.0.proj", "ffn.0"), + ("ffn.net.2", "ffn.2"), + (".norm_added_k.", ".norm_k_img."), + (".add_k_proj.", ".k_img."), + (".add_v_proj.", ".v_img."), + (".to_out.0.", ".o."), + (".to_q.", ".q."), + (".to_k.", ".k."), + (".to_v.", ".v."), + ("attn2", "cross_attn"), + ("attn1", "self_attn"), + # norm2 <-> norm3 swap back + ("norm3", "norm__placeholder"), + ("norm2", "norm3"), + ("norm__placeholder", "norm2"), + ("proj_out", "head.head"), + ] + + @classmethod + def convert_state_dict_on_save(cls, state_dict): + """Diffusers layout back to the original/comfy wan key layout + (rename-only, so quantized weight/scale/marker keys ride along).""" + if any(".self_attn." in k or ".cross_attn." in k for k in state_dict): + return state_dict # already original layout + new_sd = {} + for key, value in state_dict.items(): + k = key + # scale_shift_table: blocks keep the name as `modulation`, the + # top-level one belongs to the output head + if k == "scale_shift_table": + k = "head.modulation" + elif k.endswith(".scale_shift_table"): + k = k[: -len("scale_shift_table")] + "modulation" + for src, dst in cls._SAVE_RENAMES: + k = k.replace(src, dst) + new_sd[k] = value + return new_sd diff --git a/toolkit/models/wan21/wan21.py b/toolkit/models/wan21/wan21.py index 1fec78d..5e07efc 100644 --- a/toolkit/models/wan21/wan21.py +++ b/toolkit/models/wan21/wan21.py @@ -44,6 +44,7 @@ from typing import Any, Callable, Dict, List, Optional, Union from toolkit.models.wan21.wan_lora_convert import convert_to_diffusers, convert_to_original from toolkit.util.quantize import quantize_model from toolkit.models.v2.text_encoders.umt5 import UMT5TextEncoder +from toolkit.metadata import get_meta_for_safetensors # for generation only? scheduler_configUniPC = { @@ -680,17 +681,16 @@ class Wan21(BaseModel): return False def save_model(self, output_path, meta, save_dtype): - # only save the unet - transformer: Wan21 = unwrap_model(self.model) - transformer.save_pretrained( - save_directory=os.path.join(output_path, 'transformer'), - safe_serialization=True, + # comfy-format single-file save (original wan key layout) + transformer = unwrap_model(self.model) + if not output_path.endswith(".safetensors"): + output_path += ".safetensors" + transformer.save_model( + output_path, + dtype=save_dtype, + metadata=get_meta_for_safetensors(meta, name=self.arch), ) - meta_path = os.path.join(output_path, 'aitk_meta.yaml') - with open(meta_path, 'w') as f: - yaml.dump(meta, f) - def get_loss_target(self, *args, **kwargs): noise = kwargs.get('noise') batch = kwargs.get('batch') diff --git a/toolkit/util/get_model.py b/toolkit/util/get_model.py index 545175d..3cd9e79 100644 --- a/toolkit/util/get_model.py +++ b/toolkit/util/get_model.py @@ -41,10 +41,34 @@ def get_all_models() -> List[BaseModel]: return all_model_classes +# archs the legacy StableDiffusion monolith still serves (see the arch +# normalization in toolkit/config_modules.py) +LEGACY_ARCHS = { + "sd1", + "sd2", + "sd3", + "sdxl", + "pixart", + "pixart_sigma", + "auraflow", + "flux", + "lumina2", + "vega", + "ssd", +} + + def get_model_class(config: ModelConfig): all_models = get_all_models() for ModelClass in all_models: if ModelClass.arch == config.arch: return ModelClass - # default to the legacy model - return StableDiffusion + if config.arch in LEGACY_ARCHS: + return StableDiffusion + # a typo'd or unregistered arch used to silently fall back to SD1; error + # instead (a broken extension import also lands here — its error was + # printed during get_all_models) + known = sorted({m.arch for m in all_models if m.arch} | LEGACY_ARCHS) + raise ValueError( + f"Unknown model arch {config.arch!r}. Known archs: {', '.join(known)}" + ) diff --git a/toolkit/util/quantize.py b/toolkit/util/quantize.py index b187616..db54f58 100644 --- a/toolkit/util/quantize.py +++ b/toolkit/util/quantize.py @@ -429,9 +429,10 @@ def quantize_model( # quantize model the original way without an accuracy recovery adapter # move and quantize only certain pieces at a time. quantization_type = get_qtype(base_model.model_config.qtype) + quantize_kwargs = base_model.model_config.quantize_kwargs or {} # all_blocks = list(model_to_quantize.transformer_blocks) all_blocks: List[torch.nn.Module] = [] - transformer_block_names = base_model.get_transformer_block_names() + transformer_block_names = base_model.get_transformer_block_names() or [] for name in transformer_block_names: # name may be a dotted path for models that nest their blocks # (e.g. hidream_o1's "model.language_model.layers"). @@ -456,7 +457,12 @@ def quantize_model( block.to(base_model.device_torch, dtype=base_model.torch_dtype, non_blocking=True) # exclude patterns with a leading wildcard (e.g. "*adaln_proj*") # also apply inside blocks, where names are block-relative - quantize(block, weights=quantization_type, exclude=exclude_modules) + quantize( + block, + weights=quantization_type, + exclude=exclude_modules, + **quantize_kwargs, + ) freeze(block) # NOT non_blocking: an async D2H allocates the cpu destination in pinned # memory, which the caching host allocator keeps forever (with power-of-2 @@ -472,5 +478,10 @@ def quantize_model( # device without having to move the transformer blocks to the device first base_model.print_and_status_update(" - quantizing extras") # model_to_quantize.to(base_model.device_torch, dtype=base_model.torch_dtype) - quantize(model_to_quantize, weights=quantization_type, exclude=exclude_modules) + quantize( + model_to_quantize, + weights=quantization_type, + exclude=exclude_modules, + **quantize_kwargs, + ) freeze(model_to_quantize)