Models v2 - phase 0
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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'):
|
||||
|
||||
227
toolkit/models/v2/PLANNING.md
Normal file
227
toolkit/models/v2/PLANNING.md
Normal 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.
|
||||
@@ -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("."):
|
||||
|
||||
0
toolkit/models/v2/diffusion_models/__init__.py
Normal file
0
toolkit/models/v2/diffusion_models/__init__.py
Normal 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
|
||||
129
toolkit/models/v2/resolver.py
Normal file
129
toolkit/models/v2/resolver.py
Normal 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,
|
||||
)
|
||||
0
toolkit/models/v2/text_encoders/__init__.py
Normal file
0
toolkit/models/v2/text_encoders/__init__.py
Normal file
0
toolkit/models/v2/vae/__init__.py
Normal file
0
toolkit/models/v2/vae/__init__.py
Normal file
0
toolkit/models/v2/vision_encoders/__init__.py
Normal file
0
toolkit/models/v2/vision_encoders/__init__.py
Normal file
Reference in New Issue
Block a user