fix(model): load transformer from all-in-one checkpoints in native loader

All-in-one exports bundle the text encoder and VAE alongside the
transformer under LDM-style family prefixes (cond_stage_model.,
first_stage_model., text_encoders., vae.). The native loader treated
those keys as a mixed-prefix error and rejected the file. Drop known
companion families before prefix detection and log what was skipped;
keys matching neither a transformer prefix nor a known family still
raise. TE and VAE keep coming from the base repo or their own overrides.
This commit is contained in:
CalamitousFelicitousness
2026-07-08 21:03:49 +01:00
parent 3476125374
commit 4c289e8c9d
2 changed files with 126 additions and 6 deletions
+58 -6
View File
@@ -19,15 +19,19 @@ from diffusers' ``SINGLE_FILE_LOADABLE_CLASSES`` table.
Algorithm:
1. Read the safetensors state dict (.gguf and .pth are rejected up front).
2. Detect and strip one of the spec's known prefixes (raises on mixed prefixes).
3. Check forbidden markers (catches structural mismatches like Cosmos 1.0 keys
2. Drop keys belonging to known companion component families (all-in-one
exports bundle the text encoder and VAE under prefixes like
``cond_stage_model.`` / ``first_stage_model.``; sdnext sources those
components elsewhere).
3. Detect and strip one of the spec's known prefixes (raises on mixed prefixes).
4. Check forbidden markers (catches structural mismatches like Cosmos 1.0 keys
in a Cosmos 2.0 loader).
4. Partition off sibling component keys (e.g. Anima's bundled ``llm_adapter.*``).
5. Run the spec's converter if present (else pass through unchanged).
6. Fetch ``<subfolder>/config.json`` from the base repo, instantiate via
5. Partition off sibling component keys (e.g. Anima's bundled ``llm_adapter.*``).
6. Run the spec's converter if present (else pass through unchanged).
7. Fetch ``<subfolder>/config.json`` from the base repo, instantiate via
``cls.from_config``, ``load_state_dict(strict=False)``, validate, dtype-cast,
quantize, and offload-place.
7. Repeat the build for each populated sibling (no converter, no quant by
8. Repeat the build for each populated sibling (no converter, no quant by
default; sibling weights are read raw from the bundled file).
Returns ``(transformer, sibling_components_dict)``. The dict is empty for
@@ -58,6 +62,13 @@ DEFAULT_ACCEPTABLE_MISSING: tuple[str, ...] = (
"pos_embedder.",
"learnable_pos_embed.",
)
DEFAULT_IGNORED_PREFIXES: tuple[str, ...] = (
"cond_stage_model.",
"conditioner.",
"first_stage_model.",
"text_encoders.",
"vae.",
)
class OverrideArchMismatch(Exception):
@@ -106,11 +117,17 @@ class TransformerSpec:
``acceptable_missing`` (buffer-only keys left at their init), these are
zero-filled on load so the branch stays a no-op, matching a base model
that ships the branch dormant (output projection all-zeros).
``ignored_prefixes`` names companion component families (text encoder,
VAE) that all-in-one exports bundle alongside the transformer; their keys
are dropped before prefix detection rather than treated as a malformed
file.
"""
cls: type
subfolder: str = "transformer"
prefixes: tuple[str, ...] = DEFAULT_PREFIXES
ignored_prefixes: tuple[str, ...] = DEFAULT_IGNORED_PREFIXES
converter: Callable[[dict], dict] | None = None
siblings: dict[str, SiblingSpec] = field(default_factory=dict)
acceptable_missing: tuple[str, ...] = DEFAULT_ACCEPTABLE_MISSING
@@ -246,6 +263,7 @@ def load(
quant_type = model_quant.get_quant_type(quant_args)
state_dict = sd_models.read_state_dict(local_file, what="transformer")
state_dict = drop_companion_keys(state_dict, spec.ignored_prefixes, spec.cls.__name__)
state_dict, detected_prefix = strip_prefix(state_dict, spec.prefixes, spec.cls.__name__)
check_forbidden_markers(state_dict, spec.forbidden_markers, spec.cls.__name__, local_file)
transformer_sd, sibling_sds = partition_siblings(state_dict, spec.siblings)
@@ -309,6 +327,40 @@ def load(
return transformer, loaded_siblings
def drop_companion_keys(state_dict: dict, ignored_prefixes: tuple[str, ...], type_name: str) -> dict:
"""Remove keys belonging to known non-transformer component families.
All-in-one exports bundle the text encoder and VAE alongside the
transformer under LDM/ComfyUI-style family prefixes (``cond_stage_model.``,
``first_stage_model.``, ``text_encoders.``, ``vae.``). Only the transformer
(plus any spec-declared inline siblings) is wanted here; sdnext sources the
other components from the base repo or their own override dropdowns.
Dropped families are logged so the skip is visible in the load log.
Raises ValueError when nothing remains after filtering, meaning the
selected file holds no transformer at all.
"""
if not ignored_prefixes:
return state_dict
dropped: dict[str, int] = {}
kept: dict = {}
for key, value in state_dict.items():
prefix = next((p for p in ignored_prefixes if key.startswith(p)), None)
if prefix is None:
kept[key] = value
else:
dropped[prefix] = dropped.get(prefix, 0) + 1
if not dropped:
return state_dict
counts = " ".join(f"{prefix.rstrip('.')}={count}" for prefix, count in dropped.items())
if not kept:
raise ValueError(
f"Load model: type={type_name} native_transformer has no transformer keys ({counts})"
)
log.info(f"Load model: type={type_name} native_transformer skipping bundled components: {counts}")
return kept
def strip_prefix(state_dict: dict, prefixes: tuple[str, ...], type_name: str) -> tuple[dict, str]:
"""Detect and uniformly strip the most common known prefix from every key.
+68
View File
@@ -4,6 +4,7 @@ Offline unit tests for pipelines.native_transformer.
Covers the pure helpers that own per-arch knob handling:
- ``drop_companion_keys`` for filtering bundled TE/VAE families out of all-in-one files
- ``strip_prefix`` for single/multi prefix detection and mixed-prefix rejection
- ``partition_siblings`` for inline-sibling key partitioning
- ``check_forbidden_markers`` for structural-mismatch rejection
@@ -94,6 +95,61 @@ def run_test(cat: str, fn):
traceback.print_exc()
# ============================================================
# drop_companion_keys
# ============================================================
def test_drop_companion_keys_no_companions_pass_through():
sd = {'net.blocks.0.weight': 1, 'net.blocks.1.weight': 2}
out = nt.drop_companion_keys(sd, nt.DEFAULT_IGNORED_PREFIXES, 'Test')
assert out == sd, 'files without companion keys must pass through unchanged'
def test_drop_companion_keys_filters_all_in_one_layout():
sd = {
'net.blocks.0.weight': 1,
'net.llm_adapter.proj.weight': 2,
'cond_stage_model.qwen3_06b.transformer.model.embed_tokens.weight': 3,
'first_stage_model.decoder.conv1.weight': 4,
'vae.decoder.conv_in.weight': 5,
'text_encoders.clip_l.weight': 6,
}
out = nt.drop_companion_keys(sd, nt.DEFAULT_IGNORED_PREFIXES, 'Test')
assert set(out.keys()) == {'net.blocks.0.weight', 'net.llm_adapter.proj.weight'}
def test_drop_companion_keys_raises_when_nothing_left():
sd = {
'cond_stage_model.te.weight': 1,
'first_stage_model.decoder.weight': 2,
}
try:
nt.drop_companion_keys(sd, nt.DEFAULT_IGNORED_PREFIXES, 'Test')
raise AssertionError('expected ValueError')
except ValueError as e:
assert 'no transformer keys' in str(e)
def test_drop_companion_keys_empty_prefix_tuple_no_op():
sd = {'cond_stage_model.te.weight': 1}
out = nt.drop_companion_keys(sd, (), 'Test')
assert out == sd, 'empty ignored_prefixes must disable filtering'
def test_drop_then_strip_all_in_one_layout():
"""Companion filtering must clear the way for normal prefix detection."""
sd = {
'net.blocks.0.weight': 1,
'net.final_layer.weight': 2,
'cond_stage_model.te.weight': 3,
'first_stage_model.decoder.weight': 4,
}
out = nt.drop_companion_keys(sd, nt.DEFAULT_IGNORED_PREFIXES, 'Test')
out, prefix = nt.strip_prefix(out, nt.DEFAULT_PREFIXES, 'Test')
assert prefix == 'net.'
assert set(out.keys()) == {'blocks.0.weight', 'final_layer.weight'}
# ============================================================
# strip_prefix
# ============================================================
@@ -380,6 +436,7 @@ def test_transformer_spec_defaults():
assert spec.converter is None
assert spec.siblings == {}
assert spec.acceptable_missing == ('rope.', 'pos_embedder.', 'learnable_pos_embed.')
assert spec.ignored_prefixes == ('cond_stage_model.', 'conditioner.', 'first_stage_model.', 'text_encoders.', 'vae.')
assert spec.forbidden_markers == ()
@@ -779,6 +836,17 @@ def test_load_converter_crash_raises_mismatch():
# ============================================================
def run_all():
log.warning('=== drop_companion_keys ===')
cat = category('companion')
for fn in [
test_drop_companion_keys_no_companions_pass_through,
test_drop_companion_keys_filters_all_in_one_layout,
test_drop_companion_keys_raises_when_nothing_left,
test_drop_companion_keys_empty_prefix_tuple_no_op,
test_drop_then_strip_all_in_one_layout,
]:
run_test(cat, fn)
log.warning('=== strip_prefix ===')
cat = category('strip')
for fn in [