Models v2 - phase 0

This commit is contained in:
Jaret Burkett
2026-08-27 10:51:43 -06:00
parent 5497a001cb
commit e8d9cf6d35
29 changed files with 583 additions and 401 deletions

View File

@@ -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

View File

@@ -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()
}

View File

@@ -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"

View File

@@ -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"

View File

@@ -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

View File

@@ -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()
}

View File

@@ -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"

View File

@@ -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:

View File

@@ -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"

View File

@@ -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

View File

@@ -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()
}

View File

@@ -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}"
from toolkit.models.v2.resolver import (
repo_id_from_name_or_path,
resolve_comfy_file,
)
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}"
)
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

View File

@@ -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"

View File

@@ -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()
}
lora_keys_use_comfy_prefix = True
class MinimaxH3Ref2VAModel(MinimaxH3Model):
"""Reference-to-video (ref2va): the control images ride along as reference

View File

@@ -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

View File

@@ -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"

View File

@@ -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()
}

View File

@@ -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:

View File

@@ -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

View File

@@ -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

View File

@@ -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'):

View File

@@ -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.

View File

@@ -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)
has_quant_markers = any(k.endswith(".comfy_quant") for k in state_dict)
model = cls.aitk_from_config(config)
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 = cls.aitk_from_config(config)
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("."):

View File

@@ -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

View File

@@ -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,
)

View File