From e8d9cf6d35976d5905c0578c1b0825a6eaa32305 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Thu, 27 Aug 2026 10:51:43 -0600 Subject: [PATCH] Models v2 - phase 0 --- .../audio_models/base_audio_model.py | 16 +- .../boogu_image/boogu_image.py | 11 +- .../diffusion_models/chroma/chroma_model.py | 15 +- .../chroma/chroma_radiance_model.py | 15 +- .../ernie_image/ernie_image.py | 13 +- .../example_model/example_model.py | 19 +- .../diffusion_models/f_light/f_light.py | 15 +- .../diffusion_models/flux2/flux2_model.py | 14 +- .../diffusion_models/hidream/hidream_model.py | 15 +- .../diffusion_models/ideogram4/ideogram4.py | 13 +- .../diffusion_models/krea2/krea2.py | 11 +- .../diffusion_models/ltx2/ltx2.py | 110 ++------- .../diffusion_models/mageflow/mageflow.py | 13 +- .../diffusion_models/minimax_h3/minimax_h3.py | 93 ++----- .../nucleus_image/nucleus_image_model.py | 13 +- .../diffusion_models/omnigen2/__init__.py | 16 +- .../prx_pixel_t2i/prx_pixel_t2i.py | 11 +- .../diffusion_models/qwen_image/qwen_image.py | 14 +- .../diffusion_models/z_image/z_image.py | 17 +- .../zeta_chroma/zeta_chroma_model.py | 14 +- toolkit/models/base_model.py | 15 +- toolkit/models/v2/PLANNING.md | 227 ++++++++++++++++++ toolkit/models/v2/_mixin.py | 151 +++++++++++- .../models/v2/diffusion_models/__init__.py | 0 .../v2/{ => diffusion_models}/z_image.py | 4 +- toolkit/models/v2/resolver.py | 129 ++++++++++ toolkit/models/v2/text_encoders/__init__.py | 0 toolkit/models/v2/vae/__init__.py | 0 toolkit/models/v2/vision_encoders/__init__.py | 0 29 files changed, 583 insertions(+), 401 deletions(-) create mode 100644 toolkit/models/v2/PLANNING.md create mode 100644 toolkit/models/v2/diffusion_models/__init__.py rename toolkit/models/v2/{ => diffusion_models}/z_image.py (98%) create mode 100644 toolkit/models/v2/resolver.py create mode 100644 toolkit/models/v2/text_encoders/__init__.py create mode 100644 toolkit/models/v2/vae/__init__.py create mode 100644 toolkit/models/v2/vision_encoders/__init__.py diff --git a/extensions_built_in/audio_models/base_audio_model.py b/extensions_built_in/audio_models/base_audio_model.py index 6860a1a..f47b8f6 100644 --- a/extensions_built_in/audio_models/base_audio_model.py +++ b/extensions_built_in/audio_models/base_audio_model.py @@ -69,21 +69,7 @@ class BaseAudioModel(BaseModel): "save_model is not implemented for this model. Use the pipeline directly instead." ) - def convert_lora_weights_before_save(self, state_dict): - # currently starte with transformer. but needs to start with diffusion_model. for comfyui - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd - - def convert_lora_weights_before_load(self, state_dict): - # saved as diffusion_model. but needs to be transformer. for ai-toolkit - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True def encode_images(self, image_list: torch.Tensor, device=None, dtype=None): # make it more obvious for audio models 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 fe5b75b..8241561 100644 --- a/extensions_built_in/diffusion_models/boogu_image/boogu_image.py +++ b/extensions_built_in/diffusion_models/boogu_image/boogu_image.py @@ -446,14 +446,5 @@ class BooguImageModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["double_stream_layers", "single_stream_layers"] - def convert_lora_weights_before_save(self, state_dict): - return { - k.replace("transformer.", "diffusion_model."): v - for k, v in state_dict.items() - } + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - return { - k.replace("diffusion_model.", "transformer."): v - for k, v in state_dict.items() - } diff --git a/extensions_built_in/diffusion_models/chroma/chroma_model.py b/extensions_built_in/diffusion_models/chroma/chroma_model.py index 236d950..15da572 100644 --- a/extensions_built_in/diffusion_models/chroma/chroma_model.py +++ b/extensions_built_in/diffusion_models/chroma/chroma_model.py @@ -443,21 +443,8 @@ class ChromaModel(BaseModel): batch = kwargs.get('batch') return (noise - batch.latents).detach() - def convert_lora_weights_before_save(self, state_dict): - # currently starte with transformer. but needs to start with diffusion_model. for comfyui - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - # saved as diffusion_model. but needs to be transformer. for ai-toolkit - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd def get_base_model_version(self): return "chroma" 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 333600e..e5a79fe 100644 --- a/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py +++ b/extensions_built_in/diffusion_models/chroma/chroma_radiance_model.py @@ -425,21 +425,8 @@ class ChromaRadianceModel(BaseModel): batch = kwargs.get('batch') return (noise - batch.latents).detach() - def convert_lora_weights_before_save(self, state_dict): - # currently starte with transformer. but needs to start with diffusion_model. for comfyui - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - # saved as diffusion_model. but needs to be transformer. for ai-toolkit - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd def get_base_model_version(self): return "chroma_radiance" 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 dba1204..bb296f6 100644 --- a/extensions_built_in/diffusion_models/ernie_image/ernie_image.py +++ b/extensions_built_in/diffusion_models/ernie_image/ernie_image.py @@ -377,16 +377,5 @@ class ErnieImageModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["layers"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd diff --git a/extensions_built_in/diffusion_models/example_model/example_model.py b/extensions_built_in/diffusion_models/example_model/example_model.py index a813f39..51b973e 100644 --- a/extensions_built_in/diffusion_models/example_model/example_model.py +++ b/extensions_built_in/diffusion_models/example_model/example_model.py @@ -504,19 +504,8 @@ class ExampleModel(BaseModel): attribute in src/model.py.""" return ["blocks"] - def convert_lora_weights_before_save(self, state_dict): - """Map internal LoRA keys to the ecosystem-standard naming right before - the .safetensors is written. Most modern models ship LoRAs with a - ``diffusion_model.`` prefix (ComfyUI convention); internally ai-toolkit - uses ``transformer.``.""" - return { - k.replace("transformer.", "diffusion_model."): v - for k, v in state_dict.items() - } + # LoRA keys save with the ecosystem-standard ``diffusion_model.`` prefix + # (ComfyUI convention) and load back to the internal ``transformer.`` + # prefix; see BaseModel.convert_lora_weights_before_save/load + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - """Inverse of the above, applied when resuming from a saved LoRA.""" - return { - k.replace("diffusion_model.", "transformer."): v - for k, v in state_dict.items() - } 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 2fc4f5a..813e6c3 100644 --- a/extensions_built_in/diffusion_models/f_light/f_light.py +++ b/extensions_built_in/diffusion_models/f_light/f_light.py @@ -270,21 +270,8 @@ class FLiteModel(BaseModel): # return (noise - batch.latents).detach() return (batch.latents - noise).detach() - def convert_lora_weights_before_save(self, state_dict): - # currently starte with transformer. but needs to start with diffusion_model. for comfyui - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - # saved as diffusion_model. but needs to be transformer. for ai-toolkit - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd def get_base_model_version(self): return "f-lite" diff --git a/extensions_built_in/diffusion_models/flux2/flux2_model.py b/extensions_built_in/diffusion_models/flux2/flux2_model.py index e22a1cd..e39e468 100644 --- a/extensions_built_in/diffusion_models/flux2/flux2_model.py +++ b/extensions_built_in/diffusion_models/flux2/flux2_model.py @@ -505,19 +505,7 @@ class Flux2Model(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["double_blocks", "single_blocks"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd - - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True def encode_images(self, image_list: List[torch.Tensor], device=None, dtype=None): if device is None: diff --git a/extensions_built_in/diffusion_models/hidream/hidream_model.py b/extensions_built_in/diffusion_models/hidream/hidream_model.py index 7bba831..71921aa 100644 --- a/extensions_built_in/diffusion_models/hidream/hidream_model.py +++ b/extensions_built_in/diffusion_models/hidream/hidream_model.py @@ -432,21 +432,8 @@ class HidreamModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ['double_stream_blocks', 'single_stream_blocks'] - def convert_lora_weights_before_save(self, state_dict): - # currently starte with transformer. but needs to start with diffusion_model. for comfyui - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - # saved as diffusion_model. but needs to be transformer. for ai-toolkit - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd def get_base_model_version(self): return "hidream_i1" diff --git a/extensions_built_in/diffusion_models/ideogram4/ideogram4.py b/extensions_built_in/diffusion_models/ideogram4/ideogram4.py index e42a4a2..390f7a4 100644 --- a/extensions_built_in/diffusion_models/ideogram4/ideogram4.py +++ b/extensions_built_in/diffusion_models/ideogram4/ideogram4.py @@ -620,16 +620,5 @@ class Ideogram4Model(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["layers"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd diff --git a/extensions_built_in/diffusion_models/krea2/krea2.py b/extensions_built_in/diffusion_models/krea2/krea2.py index 128fdb0..22d75a8 100644 --- a/extensions_built_in/diffusion_models/krea2/krea2.py +++ b/extensions_built_in/diffusion_models/krea2/krea2.py @@ -867,14 +867,5 @@ class Krea2Model(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["blocks"] - def convert_lora_weights_before_save(self, state_dict): - return { - k.replace("transformer.", "diffusion_model."): v - for k, v in state_dict.items() - } + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - return { - k.replace("diffusion_model.", "transformer."): v - for k, v in state_dict.items() - } diff --git a/extensions_built_in/diffusion_models/ltx2/ltx2.py b/extensions_built_in/diffusion_models/ltx2/ltx2.py index bbf5e57..2821e23 100644 --- a/extensions_built_in/diffusion_models/ltx2/ltx2.py +++ b/extensions_built_in/diffusion_models/ltx2/ltx2.py @@ -1191,23 +1191,17 @@ class LTX2Model(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["transformer_blocks"] + lora_keys_use_comfy_prefix = True + def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - new_sd = convert_lora_diffusers_to_original(new_sd, version=self.ltx_version) - return new_sd + state_dict = super().convert_lora_weights_before_save(state_dict) + return convert_lora_diffusers_to_original(state_dict, version=self.ltx_version) def convert_lora_weights_before_load(self, state_dict): state_dict = convert_lora_original_to_diffusers( state_dict, version=self.ltx_version ) - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd + return super().convert_lora_weights_before_load(state_dict) class LTX23Model(LTX2Model): @@ -1239,88 +1233,30 @@ class LTX25Model(LTX2Model): ltx_te_path = None # ------------------------------------------------------------------ - # ComfyUI-style file resolution (mirrors MinimaxH3Model) + # ComfyUI-style file resolution (toolkit/models/v2/resolver.py) # ------------------------------------------------------------------ - @staticmethod - def _find_file_recursive(root_dir: str, filename: str) -> Optional[str]: - if not os.path.isdir(root_dir): - return None - for dirpath, dirnames, filenames in os.walk(root_dir): - dirnames.sort() - if filename in filenames: - return os.path.join(dirpath, filename) - return None - def _resolve_comfy_file(self, component: str) -> str: - """Find a weight file at its local location, or download it there - when (and only when) it is missing. - - Search order: model_kwargs override, the repo-relative path under - MODELS_PATH (diffusion_models/, text_encoders/, vae/), the bare - filename at the root, any subfolder of the component's category - folder (recursive), then the hub — downloaded to the repo-relative - path under MODELS_PATH. - """ - override = self.model_config.model_kwargs.get(f"{component}_path", None) - if override is not None: - if not os.path.exists(override): - raise FileNotFoundError( - f"model_kwargs.{component}_path does not exist: {override}" - ) - return override - - rel_path = COMFY_LTX25_FILES[component] - filename = os.path.basename(rel_path) - category = os.path.dirname(rel_path) - for rel in (rel_path, filename): - candidate = os.path.join(MODELS_PATH, rel) - if os.path.exists(candidate): - return candidate - found = self._find_file_recursive(os.path.join(MODELS_PATH, category), filename) - if found is not None: - return found - - repo_id = COMFY_LTX25_REPO - name_or_path = self.model_config.name_or_path - if ( - name_or_path - and not os.path.exists(name_or_path) - and not name_or_path.endswith(".safetensors") - and "/" in name_or_path - ): - repo_id = name_or_path - self.print_and_status_update( - f"Downloading {rel_path} from {repo_id} into {MODELS_PATH}" + from toolkit.models.v2.resolver import ( + repo_id_from_name_or_path, + resolve_comfy_file, ) - return huggingface_hub.hf_hub_download( - repo_id=repo_id, filename=rel_path, token=HF_TOKEN, local_dir=MODELS_PATH + + return resolve_comfy_file( + COMFY_LTX25_FILES[component], + repo_id=repo_id_from_name_or_path( + self.model_config.name_or_path, COMFY_LTX25_REPO + ), + override_path=self.model_config.model_kwargs.get( + f"{component}_path", None + ), + hf_token=HF_TOKEN, + status_fn=self.print_and_status_update, ) def _resolve_named_file(self, path: str, component: str) -> str: - """Resolve an explicit .safetensors path: local file, models-folder - file, or an 'org/repo/path/file.safetensors' hub path (downloaded - into the models folder).""" - if os.path.exists(path): - return path - splits = path.split("/") - if len(splits) < 3: - raise ValueError( - f"Invalid {component} path: {path}. Must be a local file or " - "'repo_id/repo/filename.safetensors' to download from the Hugging Face Hub." - ) - rel_path = "/".join(splits[2:]) - for candidate in ( - os.path.join(MODELS_PATH, rel_path), - os.path.join(MODELS_PATH, splits[-1]), - ): - if os.path.exists(candidate): - return candidate - return huggingface_hub.hf_hub_download( - repo_id="/".join(splits[:2]), - filename=rel_path, - token=HF_TOKEN, - local_dir=MODELS_PATH, - ) + from toolkit.models.v2.resolver import resolve_named_file + + return resolve_named_file(path, component=component, hf_token=HF_TOKEN) def _resolve_dit_path(self) -> str: name_or_path = self.model_config.name_or_path diff --git a/extensions_built_in/diffusion_models/mageflow/mageflow.py b/extensions_built_in/diffusion_models/mageflow/mageflow.py index 0d15ffd..589f9f2 100644 --- a/extensions_built_in/diffusion_models/mageflow/mageflow.py +++ b/extensions_built_in/diffusion_models/mageflow/mageflow.py @@ -629,18 +629,7 @@ class MageFlowModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["transformer_blocks"] - def convert_lora_weights_before_save(self, state_dict): - return { - k.replace("transformer.", "diffusion_model."): v - for k, v in state_dict.items() - } - - def convert_lora_weights_before_load(self, state_dict): - return { - k.replace("diffusion_model.", "transformer."): v - for k, v in state_dict.items() - } - + lora_keys_use_comfy_prefix = True class MageFlowEditModel(MageFlowModel): arch = "mageflow_edit" 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 1ec8619..ab0d7fa 100644 --- a/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py +++ b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py @@ -52,6 +52,11 @@ from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.memory_management import MemoryManager from toolkit.metadata import get_meta_for_safetensors from toolkit.models.base_model import BaseModel +from toolkit.models.v2.resolver import ( + find_file_recursive, + repo_id_from_name_or_path, + resolve_comfy_file, +) from toolkit.paths import MODELS_PATH from toolkit.util.comfy_quant_import import import_comfy_quantized_layers from toolkit.util.ostris_quant import OstrisLinear @@ -233,64 +238,23 @@ class MinimaxH3Model(BaseModel): # ------------------------------------------------------------------ # Loading # ------------------------------------------------------------------ - @staticmethod - def _find_file_recursive(root_dir: str, filename: str) -> Optional[str]: - """First (breadth-stable, sorted) match of ``filename`` anywhere under - ``root_dir``.""" - if not os.path.isdir(root_dir): - return None - for dirpath, dirnames, filenames in os.walk(root_dir): - dirnames.sort() - if filename in filenames: - return os.path.join(dirpath, filename) - return None - def _resolve_comfy_file(self, component: str) -> str: - """Find a weight file at its local location, or download it there - when (and only when) it is missing. - - Search order: model_kwargs override, the repo-relative path under - MODELS_PATH (diffusion_models/, text_encoders/, vae/), the bare - filename at the root, any subfolder of the component's category - folder (recursive — e.g. diffusion_models/my_custom_sub/), the same - spots under name_or_path when it is a local folder, then the hub — - downloaded to the repo-relative path under MODELS_PATH. - """ - override = self.model_config.model_kwargs.get(f"{component}_path", None) - if override is not None: - if not os.path.exists(override): - raise FileNotFoundError( - f"model_kwargs.{component}_path does not exist: {override}" - ) - return override - - rel_path = COMFY_FILES[component] - filename = os.path.basename(rel_path) - category = os.path.dirname(rel_path) - roots = [MODELS_PATH] + """Find a weight file at its local location (model_kwargs override, + comfy layout under MODELS_PATH or a local name_or_path dir), or + download it there when (and only when) it is missing — see + toolkit/models/v2/resolver.py for the search order.""" name_or_path = self.model_config.name_or_path - if name_or_path and os.path.isdir(name_or_path): - roots.append(name_or_path) - for root in roots: - for rel in (rel_path, filename): - candidate = os.path.join(root, rel) - if os.path.exists(candidate): - return candidate - for root in roots: - found = self._find_file_recursive(os.path.join(root, category), filename) - if found is not None: - return found - - import huggingface_hub - - repo_id = COMFY_REPO - if name_or_path and not os.path.exists(name_or_path) and "/" in name_or_path: - repo_id = name_or_path - self.print_and_status_update( - f"Downloading {rel_path} from {repo_id} into {MODELS_PATH}" + extra_roots = ( + [name_or_path] if name_or_path and os.path.isdir(name_or_path) else [] ) - return huggingface_hub.hf_hub_download( - repo_id=repo_id, filename=rel_path, local_dir=MODELS_PATH + return resolve_comfy_file( + COMFY_FILES[component], + repo_id=repo_id_from_name_or_path(name_or_path, COMFY_REPO), + override_path=self.model_config.model_kwargs.get( + f"{component}_path", None + ), + extra_roots=extra_roots, + status_fn=self.print_and_status_update, ) def _dit_component(self) -> str: @@ -322,7 +286,7 @@ class MinimaxH3Model(BaseModel): lora_path = self.model_config.assistant_lora_path if not os.path.exists(lora_path): filename = os.path.basename(lora_path) - found = self._find_file_recursive( + found = find_file_recursive( os.path.join(MODELS_PATH, "loras"), filename ) if found is not None: @@ -1226,20 +1190,9 @@ class MinimaxH3Model(BaseModel): "*adaln_proj*", ] - def convert_lora_weights_before_save(self, state_dict): - # ComfyUI's MiniMax-H3 keys are the original checkpoint keys, so the - # standard diffusion_model prefix maps directly - return { - k.replace("transformer.", "diffusion_model."): v - for k, v in state_dict.items() - } - - def convert_lora_weights_before_load(self, state_dict): - return { - k.replace("diffusion_model.", "transformer."): v - for k, v in state_dict.items() - } - + # ComfyUI's MiniMax-H3 keys are the original checkpoint keys, so the + # standard diffusion_model prefix maps directly + lora_keys_use_comfy_prefix = True class MinimaxH3Ref2VAModel(MinimaxH3Model): """Reference-to-video (ref2va): the control images ride along as reference 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 572fda8..54ec7ff 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 @@ -405,16 +405,5 @@ class NucleusImageModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["transformer_blocks"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd diff --git a/extensions_built_in/diffusion_models/omnigen2/__init__.py b/extensions_built_in/diffusion_models/omnigen2/__init__.py index 77ce910..edb10bd 100644 --- a/extensions_built_in/diffusion_models/omnigen2/__init__.py +++ b/extensions_built_in/diffusion_models/omnigen2/__init__.py @@ -343,21 +343,7 @@ class OmniGen2Model(BaseModel): return ["noise_refiner", "context_refiner", "ref_image_refiner", "layers"] return ["noise_refiner", "context_refiner", "layers"] - def convert_lora_weights_before_save(self, state_dict): - # currently starte with transformer. but needs to start with diffusion_model. for comfyui - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd - - def convert_lora_weights_before_load(self, state_dict): - # saved as diffusion_model. but needs to be transformer. for ai-toolkit - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True def get_base_model_version(self): return "omnigen2" 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 fb44ded..b6f2137 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 @@ -332,14 +332,5 @@ class PRXPixelT2IModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["blocks"] - def convert_lora_weights_before_save(self, state_dict): - return { - k.replace("transformer.", "diffusion_model."): v - for k, v in state_dict.items() - } + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - return { - k.replace("diffusion_model.", "transformer."): v - for k, v in state_dict.items() - } 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 c03af90..5199321 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py @@ -415,19 +415,7 @@ class QwenImageModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["transformer_blocks"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd - - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True def encode_images(self, image_list: List[torch.Tensor], device=None, dtype=None): if device is None: 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 bf13983..324b76a 100644 --- a/extensions_built_in/diffusion_models/z_image/z_image.py +++ b/extensions_built_in/diffusion_models/z_image/z_image.py @@ -31,8 +31,8 @@ try: from diffusers import ZImagePipeline # our subclass of the diffusers transformer with the universal loading / - # quantization mixin (see toolkit/models/classes/_mixin.py) - from toolkit.models.v2.z_image import ZImageTransformer2DModel + # quantization mixin (see toolkit/models/v2/_mixin.py) + from toolkit.models.v2.diffusion_models.z_image import ZImageTransformer2DModel except ImportError: raise ImportError( "Diffusers is out of date. Update diffusers to the latest version by doing pip uninstall diffusers and then pip install -r requirements.txt" @@ -464,16 +464,5 @@ class ZImageModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["layers"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd 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 ca8b491..a4f3c0c 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 @@ -27,7 +27,6 @@ from .zeta_chroma_transformer import ZImageDCT, ZImageDCTParams, vae_flatten, va from .zeta_chroma_pipeline import ZetaChromaPipeline - scheduler_config = { "num_train_timesteps": 1000, "use_dynamic_shifting": False, @@ -368,16 +367,5 @@ class ZetaChromaModel(BaseModel): def get_transformer_block_names(self) -> Optional[List[str]]: return ["layers"] - def convert_lora_weights_before_save(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("transformer.", "diffusion_model.") - new_sd[new_key] = value - return new_sd + lora_keys_use_comfy_prefix = True - def convert_lora_weights_before_load(self, state_dict): - new_sd = {} - for key, value in state_dict.items(): - new_key = key.replace("diffusion_model.", "transformer.") - new_sd[new_key] = value - return new_sd diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index 9aadeea..ba4a1d7 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -98,6 +98,9 @@ UNET_IN_CHANNELS = 4 # Stable Diffusion の in_channels は 4 で固定。XLも class BaseModel: # override these in child classes arch = None + # rename LoRA keys transformer. <-> diffusion_model. (the ComfyUI-standard + # prefix) on save/load + lora_keys_use_comfy_prefix = False def __init__( self, @@ -1627,10 +1630,20 @@ class BaseModel: 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: + return { + k.replace("transformer.", "diffusion_model."): v + for k, v in state_dict.items() + } return state_dict - + def convert_lora_weights_before_load(self, state_dict): # can be overridden in child classes to convert weights before loading + if self.lora_keys_use_comfy_prefix: + return { + k.replace("diffusion_model.", "transformer."): v + for k, v in state_dict.items() + } return state_dict def condition_noisy_latents(self, latents: torch.Tensor, batch:'DataLoaderBatchDTO'): diff --git a/toolkit/models/v2/PLANNING.md b/toolkit/models/v2/PLANNING.md new file mode 100644 index 0000000..796723f --- /dev/null +++ b/toolkit/models/v2/PLANNING.md @@ -0,0 +1,227 @@ +# v2 Model Module Restructure — Planning + +## Goal + +Every component the toolkit loads (DiTs/transformers, unets, text encoders, vision +encoders, VAEs, audio VAEs) becomes a class extending one base module in +`toolkit/models/v2`. `BaseModel` (`toolkit/models/base_model.py`) stays as the +multimodal holder that each arch in `extensions_built_in/diffusion_models` extends — +that layer is good. The layer below it is what gets unified: one loading entry point, +one quantization path, one save path, shared component definitions instead of +per-model-folder copies. + +End state this enables: + +- **Live server with model hot-swap**: a resident process where, when a generation or + training run requests a different model, the unused components are dropped and the + new ones loaded. Shared component classes (same TE/VAE reused across archs) make + component-level reuse possible instead of full teardown/reload. +- **Model loading test suite**: a test that loads each registered arch one at a time + and runs an inference pass. Every model type gets added to this suite as it is + migrated (see Testing below). +- **Comfy-aligned weights**: weights live in the ComfyUI folder layout under + `MODELS_PATH` (shareable with a ComfyUI install), download there when missing, and + saves are comfy-format. Eventually defaults move to comfy / our own prequantized + releases for everything. + +## Current state (survey 2026-08-27) + +Three generations of loading conventions coexist: + +1. Legacy monolith `toolkit/stable_diffusion_model.py` (`is_flux` / `is_v3` branches); + still the silent fallback in `toolkit/util/get_model.py` when an arch string + doesn't match. +2. `BaseModel` subclass per arch (35 registered classes), each with a hand-written + `load_model()` / `save_model()`. +3. `toolkit/models/v2/_mixin.py` (`OstrisModelMixin`) — the intended fix, currently + used by one model (`v2/z_image.py` → z_image extension). + +### Duplication highlights + +- BFL KL autoencoder: full copies in `flux2/src/autoencoder.py` and + `ideogram4/src/vae.py` (header says "Flux2 KL autoencoder"). +- Qwen3-VL text encoder loaded independently in qwen_image, nucleus_image, krea2, + ideogram4, mageflow, minimax_h3 — from several different repo sources; krea2 + hand-patches the vision tower locally. +- `Qwen3ForCausalLM` TE load: 3 verbatim line-for-line copies + (`z_image/z_image.py:257`, `z_image/z_image_l2p_model.py:471`, + `zeta_chroma/zeta_chroma_model.py:146`). +- Flux1 VAE + T5 + CLIP trio loaded 4x from 2 different repos (chroma ×2, + flux_kontext, legacy SD path). +- Comfy-file resolver copy-pasted: `minimax_h3/minimax_h3.py:248` + (`_resolve_comfy_file`) → `ltx2/ltx2.py:1254` ("mirrors MinimaxH3Model"). +- `AutoencoderKLQwenImage` latents mean/std handling triplicated (qwen_image, + nucleus_image, krea2). +- `transformer.` ↔ `diffusion_model.` LoRA key rename copy-pasted ~15x in + `convert_lora_weights_before_save/load` overrides. +- Fake CLIP/TE/config stubs redefined in ~5 places + (canonical: `toolkit/models/FakeVAE.py`, `toolkit/unloader.py`). + +### Inconsistency highlights + +- **Quantization, 5 paths**: `quantize_model()` (block-streaming, ARA-aware — ~21 + users), raw `quantize()` (~10 users, no block streaming/excludes, but the only path + honoring `quantize_kwargs`), hidream's hand-rolled block loop, the v2 mixin's own + `quantize_`, and bare TE quantization everywhere. +- **Known bugs**: ~9 sites quantize the TE with `qtype` instead of `qtype_te` + (chroma ×2, flux2, flux_kontext, cogview4, wan21, legacy SD, ...); + `toolkit/models/loaders/umt5.py` accepts a `comfy_files` param it never uses, so + wan21's comfy-TE path is a silent no-op. +- **Saving, 4 incompatible styles**: diffusers `save_pretrained` folders, flat + safetensors, safetensors-inside-diffusers-folder hybrids, and z_image's + loaded-format-dependent branch. Dequant-on-save done 3 ways; the + `isinstance(v, QTensor)` variant (chroma, flux2, boogu_image, ideogram4) misses + torchao and Ostris weights entirely. Every `save_pretrained` override ignores its + `save_dtype` argument. +- Registry: linear scan in `toolkit/util/get_model.py`, silent SD1 fallback on a + typo'd arch, eager import of every model file at startup. Second unsynchronized + registry in `ui/src/app/jobs/new/options.tsx`. + +## Decisions (locked in) + +1. **Save format = ComfyUI format.** Single-file safetensors in comfy key layout. + Must support saving quantized — primarily convrot8 and nvfp4 (comfy_quant marker + format, see `toolkit/util/comfy_quant_import.py`) — and plain bf16, all in comfy + format. Loading stays backwards compatible: diffusers dirs, transformers repos, + and single files all still digest through `load_model`; only saving standardizes + on comfy. +2. **v2 folder layout mirrors the comfy save path structure**: + + ``` + toolkit/models/v2/ + _mixin.py # base module (OstrisModelMixin, evolving) + resolver.py # comfy-layout weight resolution (lift from minimax_h3) + diffusion_models/ # one file per DiT/unet family + text_encoders/ # qwen3_vl.py, qwen3.py, t5.py, clip.py, gemma.py, ... + vae/ # flux_kl.py, qwen_image.py, wan.py, audio VAEs, ... + vision_encoders/ + ``` + +3. **Method names win from `BaseModel`**: `get_transformer_block_names` and + `get_quantization_exclude_modules`. The mixin's `get_quantization_block_names` + gets renamed to match; resolve the classmethod-vs-instance-method mismatch while + doing so. +4. **Loading policy**: + - Per-model special handling is allowed via the hook methods. + - If `name_or_path` is a diffusers/transformers source, load it with + diffusers/transformers for now. **Step 1 is migrating every model to the v2 + module format and loader without breaking anything** — same weights, same + sources, same results. + - Each model declares a `comfy_weight_names` dict keyed per standard + `name_or_path`. If the user points at a local folder or a non-standard repo, + load it as-is. If `name_or_path` is the standard repo and we have matching + comfy weight names, load those instead when any of them exist (locally under + `MODELS_PATH` in comfy layout, or downloadable to there). + - Eventually the default flips to comfy weights / our own prequantized releases + for everything. + +## Base module: what `OstrisModelMixin` still needs + +The mixin already handles: diffusers dir / hub repo / local single file / +`org/repo/file.safetensors`, key-conversion hooks on load and save, overridable +backend hooks for transformers-lib models, block-wise quantize. + +To add: + +- [x] **Comfy weight spec + resolver.** `aitk_comfy_repo` / `aitk_comfy_weight_names` + class attrs + `find_comfy_weights` (local-only until Phase 2); resolution chain + generalized from `minimax_h3._resolve_comfy_file` into `v2/resolver.py`: + explicit override → `MODELS_PATH` at the repo-relative comfy path → flat at + root → recursive walk of the category folder → hub download **to the + repo-relative path** (folder stays shareable with ComfyUI, no duplicate + downloads). +- [x] **Automatic prequantized import.** Single-file path sniffs `comfy_quant` + markers and routes through `import_comfy_quantized_layers` before + `load_state_dict`, including the OstrisLinear missing-key whitelist that + minimax_h3 and ltx2 each hand-rolled. +- [x] **One save path.** `save_model(path, dtype)`: dequantize via + `dequantize_if_quantized` (honors dtype), run `convert_state_dict_on_save`, + write single-file comfy-layout safetensors. (Quantized-storage saves — + convrot8 / nvfp4 with comfy_quant markers — land with Phase 2; diffusers-folder + save as an explicit flag still to add.) +- [x] **Tokenizer/processor declaration** for text encoders + (`aitk_tokenizer_repo`/`aitk_processor_repo` + `load_tokenizer`/`load_processor`). +- [x] Rename quantization hooks to the `BaseModel` spellings (decision 3): + `get_transformer_block_names` (classmethod on the module). + +## Migration steps + +Track progress here; check items off as they land. + +### Phase 0 — foundation (done 2026-08-27) +- [x] Evolve `_mixin.py` per the list above (comfy spec local-only until Phase 2) +- [x] Create `v2/diffusion_models/`, `v2/text_encoders/`, `v2/vae/`, + `v2/vision_encoders/`; move `v2/z_image.py` → `v2/diffusion_models/z_image.py` +- [x] Lift the comfy resolver out of minimax_h3 into `v2/resolver.py`; point + minimax_h3 and ltx2 at it (delete their copies) +- [x] `BaseModel` default `convert_lora_weights_before_save/load` doing the + `transformer.` ↔ `diffusion_model.` rename, gated on the class attr + `lora_keys_use_comfy_prefix` (default False, so passthrough models keep + their behavior); the ~18 identical overrides replaced with the flag. + Custom conversions (anima, hidream_o1, ltx2, wan21) keep their overrides; + ltx2's now composes with the flag via super(). + +### Phase 1 — migrate all models to v2 modules, no behavior change +Every arch's components become v2 classes; if `name_or_path` is diffusers, it still +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) +- [ ] Per-model fixes folded in as each migrates: `qtype_te` bug, dequant-on-save + (`dequantize_if_quantized` everywhere), raw-`quantize()` → `quantize_model()` + +### Phase 2 — comfy weights become the preferred source +- [ ] Wire `comfy_weight_names` per model; standard-repo `name_or_path` + existing + comfy weights → load comfy +- [ ] Comfy-format save (bf16 + convrot8/nvfp4 quantized) as the default + full-weight save +- [ ] Publish/verify comfy repacks per model as they flip + +### Phase 3 — live server +- [ ] Component-level identity (which TE/VAE instances are shared between archs) so + a model switch drops only what the next run doesn't need +- [ ] Resident process: request comes in → diff requested components vs loaded → + unload/load the difference +- [ ] Legacy `stable_diffusion_model.py` archs: grandfather or port last + +## Testing + +- [ ] `testing/` (or `tests/`) harness: for each migrated arch, load the model via + its v2 modules and run one small inference pass (single low-step sample; video + models at minimum frame count). One arch at a time, full unload between archs. +- [ ] Weights resolved through the normal resolver against `MODELS_PATH` + (GPU + local-weights test, not CI-portable at first; skip archs whose weights + are absent rather than failing). +- [ ] Round-trip test per model: load → save comfy format → reload from the save → + outputs match (bf16) / load cleanly (quantized saves). +- [ ] Each newly migrated model adds its test in the same PR as its migration. + +## 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. +- [ ] `toolkit/models/loaders/umt5.py` dead `comfy_files` param (wan21 comfy-TE + no-op) — fix when wan migrates. +- [ ] Registry hardening: error (don't fall back to SD1) on unknown arch; 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. +- [ ] 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 6fb71ec..4e82640 100644 --- a/toolkit/models/v2/_mixin.py +++ b/toolkit/models/v2/_mixin.py @@ -9,7 +9,11 @@ digests any of: - a HuggingFace repo id ("org/repo") - a local single .safetensors file in the model's original key layout - a remote single .safetensors file ("org/repo/path/file.safetensors") -plus optional automatic quantization of the loaded weights (`qtype=...`). +plus optional automatic quantization of the loaded weights (`qtype=...`). A +single-file checkpoint carrying ComfyUI `comfy_quant` markers loads with its +pre-quantized layers attached to the toolkit's quantization backends +(convrot8 / nvfp4) automatically. `save_model` writes a single-file +.safetensors back in the original (comfy) key layout. Model specific behavior lives in small overridable hooks (key conversion, config source, block names, backend loading), so subclassing for a new model usually means @@ -47,6 +51,19 @@ class OstrisModelMixin: # .safetensors file and no config_path is given aitk_config_repo: Optional[str] = None + # ---- comfy weight sources (Phase 2 flips the loading default to these) ---- + # hub repo the comfy-format weight files are published in + aitk_comfy_repo: Optional[str] = None + # standard name_or_path (hub repo id) -> repo-relative comfy weight file + # (ComfyUI folder layout: diffusion_models/, text_encoders/, vae/, ...) + aitk_comfy_weight_names: Dict[str, str] = {} + + # ---- tokenizer/processor source, for text-encoder modules ---- + aitk_tokenizer_repo: Optional[str] = None + aitk_tokenizer_subfolder: Optional[str] = None + aitk_processor_repo: Optional[str] = None + aitk_processor_subfolder: Optional[str] = None + # ---- state set by the loader / quantizer ---- aitk_is_quantized: bool = False aitk_qtype: Optional[str] = None @@ -68,7 +85,7 @@ class OstrisModelMixin: return state_dict @classmethod - def get_quantization_block_names(cls) -> Optional[List[str]]: + def get_transformer_block_names(cls) -> Optional[List[str]]: """Names (dotted paths allowed) of the repeated block lists to quantize one block at a time so the whole model never has to sit on the gpu at once.""" return None @@ -200,16 +217,132 @@ class OstrisModelMixin: state_dict = load_file(file_path) state_dict = cls.convert_state_dict_on_load(state_dict) - for key, value in state_dict.items(): - state_dict[key] = value.to(dtype=dtype) - + has_quant_markers = any(k.endswith(".comfy_quant") for k in state_dict) model = cls.aitk_from_config(config) - model.load_state_dict(state_dict, assign=True) - model.to(dtype=dtype) + + if has_quant_markers: + # pre-quantized comfy checkpoint: attach the quantized linears onto + # the toolkit's backends, load the rest at its stored precision + # (the checkpoint's bf16/fp16/fp32 mix is deliberate) + from toolkit.util.comfy_quant_import import import_comfy_quantized_layers + + state_dict, num_quantized = import_comfy_quantized_layers( + model, state_dict, orig_dtype=dtype + ) + cls._load_state_dict_with_quantized(model, state_dict) + model.aitk_is_quantized = True + else: + for key, value in state_dict.items(): + state_dict[key] = value.to(dtype=dtype) + model.load_state_dict(state_dict, assign=True) + model.to(dtype=dtype) del state_dict flush() return model + @staticmethod + def _load_state_dict_with_quantized(model, state_dict): + """Load a state dict onto a (meta-built) model whose quantized linears + were already attached by import_comfy_quantized_layers: their weight + (and importer-assigned bias) legitimately report as missing keys.""" + from toolkit.util.ostris_quant import OstrisLinear + + result = model.load_state_dict(state_dict, assign=True, strict=False) + quantized_param_keys = set() + for name, m in model.named_modules(): + if isinstance(m, OstrisLinear): + quantized_param_keys.add(f"{name}.weight") + if m.bias is not None: + quantized_param_keys.add(f"{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"{type(model).__name__} load mismatch: missing {bad_missing[:8]}, " + f"unexpected {result.unexpected_keys[:8]}" + ) + leftover_meta = [n for n, p in model.named_parameters() if p.is_meta] + if leftover_meta: + raise ValueError( + f"{type(model).__name__} load left meta parameters: " + f"{leftover_meta[:8]}" + ) + + # ------------------------------------------------------------------ + # comfy weight sources + # ------------------------------------------------------------------ + + @classmethod + def find_comfy_weights(cls, name_or_path: str) -> Optional[str]: + """Local comfy-format weight file registered for a standard + ``name_or_path``, or None. Never downloads — Phase 2 flips the loading + default to comfy sources; until then callers opt in explicitly.""" + from toolkit.models.v2.resolver import resolve_comfy_file + + rel_path = cls.aitk_comfy_weight_names.get(name_or_path) + if rel_path is None: + return None + return resolve_comfy_file( + rel_path, repo_id=cls.aitk_comfy_repo, local_only=True + ) + + # ------------------------------------------------------------------ + # tokenizer / processor (text-encoder modules) + # ------------------------------------------------------------------ + + @classmethod + def load_tokenizer(cls, **kwargs): + from transformers import AutoTokenizer + + if cls.aitk_tokenizer_repo is None: + raise ValueError(f"{cls.__name__} does not declare aitk_tokenizer_repo") + return AutoTokenizer.from_pretrained( + cls.aitk_tokenizer_repo, subfolder=cls.aitk_tokenizer_subfolder, **kwargs + ) + + @classmethod + def load_processor(cls, **kwargs): + from transformers import AutoProcessor + + if cls.aitk_processor_repo is None: + raise ValueError(f"{cls.__name__} does not declare aitk_processor_repo") + return AutoProcessor.from_pretrained( + cls.aitk_processor_repo, subfolder=cls.aitk_processor_subfolder, **kwargs + ) + + # ------------------------------------------------------------------ + # saving + # ------------------------------------------------------------------ + + @torch.no_grad() + def save_model( + self, + output_path: str, + dtype: Optional[torch.dtype] = None, + metadata: Optional[Dict[str, str]] = None, + ): + """Save as a single-file .safetensors in the model's original (comfy) + key layout, via convert_state_dict_on_save. Quantized weights are + dequantized to full precision; dtype, when given, casts the floating + point tensors. (Quantized-storage saves — comfy_quant markers — come + with Phase 2.)""" + from safetensors.torch import save_file + from toolkit.util.quantize import dequantize_if_quantized + + state_dict = {} + for key, value in self.state_dict().items(): + value = dequantize_if_quantized(value) + if dtype is not None and value.is_floating_point(): + value = value.to(dtype=dtype) + state_dict[key] = value.detach().to("cpu").contiguous() + state_dict = self.convert_state_dict_on_save(state_dict) + + parent = os.path.dirname(output_path) + if parent: + os.makedirs(parent, exist_ok=True) + save_file(state_dict, output_path, metadata=metadata) + del state_dict + flush() + # ------------------------------------------------------------------ # quantization # ------------------------------------------------------------------ @@ -222,7 +355,7 @@ class OstrisModelMixin: exclude: Optional[List[str]] = None, ): """Quantize the model weights in place. When device is given, the repeated - blocks (get_quantization_block_names) are moved there one at a time for the + blocks (get_transformer_block_names) are moved there one at a time for the quantization math and returned to their original device, so the whole model never has to fit on the gpu in full precision.""" from optimum.quanto import freeze @@ -238,7 +371,7 @@ class OstrisModelMixin: ) blocks: List[torch.nn.Module] = [] - for name in self.get_quantization_block_names() or []: + for name in self.get_transformer_block_names() or []: # name may be a dotted path for models that nest their blocks block_list = self for part in name.split("."): diff --git a/toolkit/models/v2/diffusion_models/__init__.py b/toolkit/models/v2/diffusion_models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/toolkit/models/v2/z_image.py b/toolkit/models/v2/diffusion_models/z_image.py similarity index 98% rename from toolkit/models/v2/z_image.py rename to toolkit/models/v2/diffusion_models/z_image.py index 90df66d..5ee32c4 100644 --- a/toolkit/models/v2/z_image.py +++ b/toolkit/models/v2/diffusion_models/z_image.py @@ -3,7 +3,7 @@ from diffusers.models.transformers import ( ZImageTransformer2DModel as DiffusersZImageTransformer2DModel, ) -from ._mixin import OstrisModelMixin +from .._mixin import OstrisModelMixin class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMixin): @@ -12,7 +12,7 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix aitk_config_repo = "Tongyi-MAI/Z-Image-Turbo" @classmethod - def get_quantization_block_names(cls): + def get_transformer_block_names(cls): return ["layers"] @classmethod diff --git a/toolkit/models/v2/resolver.py b/toolkit/models/v2/resolver.py new file mode 100644 index 0000000..fda8ff0 --- /dev/null +++ b/toolkit/models/v2/resolver.py @@ -0,0 +1,129 @@ +"""ComfyUI-layout weight file resolution. + +Weight files live under MODELS_PATH in ComfyUI's folder layout +(diffusion_models/, text_encoders/, vae/, ...) so the folder is shareable with +a ComfyUI install. Files are used in place when present and downloaded to +exactly their repo-relative location only when missing, so nothing is ever +duplicated on re-run. + +Lifted from the minimax_h3 / ltx2.5 model implementations; those now call +into here. +""" + +import os +from typing import Callable, Iterable, Optional + +from toolkit.paths import MODELS_PATH + + +def find_file_recursive(root_dir: str, filename: str) -> Optional[str]: + """First (breadth-stable, sorted) match of ``filename`` anywhere under + ``root_dir``.""" + if not os.path.isdir(root_dir): + return None + for dirpath, dirnames, filenames in os.walk(root_dir): + dirnames.sort() + if filename in filenames: + return os.path.join(dirpath, filename) + return None + + +def repo_id_from_name_or_path( + name_or_path: Optional[str], default: str +) -> str: + """Treat a hub-style ``name_or_path`` ("org/repo") as a replacement comfy + repo; anything local (or an explicit .safetensors file) keeps the + default.""" + if ( + name_or_path + and not os.path.exists(name_or_path) + and not name_or_path.endswith(".safetensors") + and "/" in name_or_path + ): + return name_or_path + return default + + +def resolve_comfy_file( + rel_path: str, + repo_id: str, + override_path: Optional[str] = None, + extra_roots: Optional[Iterable[str]] = None, + hf_token: Optional[str] = None, + status_fn: Optional[Callable[[str], None]] = None, + local_only: bool = False, +) -> Optional[str]: + """Find a weight file at its local location, or download it there when + (and only when) it is missing. + + Search order: ``override_path`` (must exist), the repo-relative path under + MODELS_PATH (and each of ``extra_roots``), the bare filename at each root, + any subfolder of the category folder (recursive — e.g. + diffusion_models/my_custom_sub/), then the hub — downloaded to the + repo-relative path under MODELS_PATH. With ``local_only`` the hub is never + touched and a miss returns None. + """ + if override_path is not None: + if not os.path.exists(override_path): + raise FileNotFoundError( + f"Override path for {rel_path} does not exist: {override_path}" + ) + return override_path + + filename = os.path.basename(rel_path) + category = os.path.dirname(rel_path) + roots = [MODELS_PATH] + [r for r in (extra_roots or []) if os.path.isdir(r)] + for root in roots: + for rel in (rel_path, filename): + candidate = os.path.join(root, rel) + if os.path.exists(candidate): + return candidate + for root in roots: + found = find_file_recursive(os.path.join(root, category), filename) + if found is not None: + return found + + if local_only: + return None + + import huggingface_hub + + if status_fn is not None: + status_fn(f"Downloading {rel_path} from {repo_id} into {MODELS_PATH}") + return huggingface_hub.hf_hub_download( + repo_id=repo_id, filename=rel_path, token=hf_token, local_dir=MODELS_PATH + ) + + +def resolve_named_file( + path: str, + component: str = "model", + hf_token: Optional[str] = None, +) -> str: + """Resolve an explicit .safetensors reference: a local file, a file already + under MODELS_PATH, or an 'org/repo/path/file.safetensors' hub path + (downloaded into the models folder at its repo-relative path).""" + if os.path.exists(path): + return path + splits = path.split("/") + if len(splits) < 3: + raise ValueError( + f"Invalid {component} path: {path}. Must be a local file or " + "'org/repo/filename.safetensors' to download from the Hugging Face Hub." + ) + rel_path = "/".join(splits[2:]) + for candidate in ( + os.path.join(MODELS_PATH, rel_path), + os.path.join(MODELS_PATH, splits[-1]), + ): + if os.path.exists(candidate): + return candidate + + import huggingface_hub + + return huggingface_hub.hf_hub_download( + repo_id="/".join(splits[:2]), + filename=rel_path, + token=hf_token, + local_dir=MODELS_PATH, + ) diff --git a/toolkit/models/v2/text_encoders/__init__.py b/toolkit/models/v2/text_encoders/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/toolkit/models/v2/vae/__init__.py b/toolkit/models/v2/vae/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/toolkit/models/v2/vision_encoders/__init__.py b/toolkit/models/v2/vision_encoders/__init__.py new file mode 100644 index 0000000..e69de29