From 82767772fc041ecf0d516404b4c87f56474f778c Mon Sep 17 00:00:00 2001 From: Alexander Eichhorn Date: Sun, 2 Aug 2026 08:50:37 +0200 Subject: [PATCH] feat(architectures): declare variant enums per architecture and guard the unions One list of eleven variant enums is written by hand four times -- AnyVariant, the variant_type_adapter subscript, that adapter's runtime argument (all three in taxonomy.py) and ModelRecordChanges.variant. Nothing checked that they agreed. Nothing checked the invariant they all rest on either: variant *values* must be globally unique, because build_common_fields resolves a bare variant string without knowing the base. That invariant has already cost something -- it is why Krea2VariantType.Turbo is `krea2_turbo` rather than `turbo` -- and until now it lived only in a docstring. VariantFacet is keyed by model type, not by base alone. The plan assumed one variant enum per base. The code disagrees, and the mapping here was derived from Config_Base.CONFIG_CLASSES rather than read off by hand: - Wan mains carry WanVariantType, Wan LoRAs carry WanLoRAVariantType. Same base, different enum, and not interchangeable -- an A14B LoRA against a TI2V-5B main crashes in the layer patcher on a tensor shape. - PiDDecoderVariantType is one enum shared across five bases. - ClipVariantType and Qwen3VariantType sit on base=Any configs. Any is a sentinel the registry refuses to register, so they cannot be declared at all. The last point is why the completeness check carries an explicit BASE_AGNOSTIC_VARIANT_ENUMS allowlist: without it the check would silently cover 9 of 11 enums while looking total. A second test asserts that nothing in that allowlist is used by a real architecture, so it cannot become a dumping ground. Four architectures -- CogView4, ERNIE-Image, Ideogram 4, Anima -- model no variants and omit the facet. An empty facet would be indistinguishable from an absent one. configs/factory.py is deliberately untouched. Base-aware resolution would still need a fallback for the base=Any models, so it would move the trap rather than close it, at the cost of touching the identification path. The invariant is pinned by a test instead, in two forms: enum values are unique, and variant_type_adapter.validate_strings returns the right enum type for every value -- the second being the behavioural version of the first, through the exact call factory.py makes. test_every_facet_declaration_matches_the_config_classes derives the expected mapping from the config classes and checks both directions, so a declaration the configs do not back fails just as loudly as a missing one. Verified against both mistakes: pointing Wan LoRAs at the main enum, and declaring a model type no config has a variant field for. Deriving rather than reading also corrected the mapping: sdxl-refiner mains carry ModelVariantType, which a hand-written list would have missed. configs/main.py's from_base signature still lists only six of the enums and omits QwenImageVariantType. Left alone -- that annotation is on the way out in the capabilities PR, which derives from_base from the registry. openapi.json, schema.ts and invocation-context.json are unchanged. Co-Authored-By: Claude Opus 5 (1M context) --- invokeai/backend/architectures/__init__.py | 8 + invokeai/backend/architectures/defs/flux.py | 14 +- invokeai/backend/architectures/defs/flux2.py | 16 +- invokeai/backend/architectures/defs/krea_2.py | 6 +- .../backend/architectures/defs/qwen_image.py | 14 +- invokeai/backend/architectures/defs/sd_1.py | 4 +- invokeai/backend/architectures/defs/sd_2.py | 4 +- invokeai/backend/architectures/defs/sd_3.py | 5 +- invokeai/backend/architectures/defs/sdxl.py | 14 +- .../architectures/defs/sdxl_refiner.py | 4 +- invokeai/backend/architectures/defs/wan.py | 17 +- .../backend/architectures/defs/z_image.py | 9 +- .../backend/architectures/facets/variant.py | 69 ++++++ tests/backend/architectures/test_variants.py | 212 ++++++++++++++++++ 14 files changed, 385 insertions(+), 11 deletions(-) create mode 100644 invokeai/backend/architectures/facets/variant.py create mode 100644 tests/backend/architectures/test_variants.py diff --git a/invokeai/backend/architectures/__init__.py b/invokeai/backend/architectures/__init__.py index 10d95f81758..68ea52460db 100644 --- a/invokeai/backend/architectures/__init__.py +++ b/invokeai/backend/architectures/__init__.py @@ -39,6 +39,11 @@ resolve_latent_space, ) from invokeai.backend.architectures.facets.unet import UNetDownscaleFacet, get_max_unet_downscale +from invokeai.backend.architectures.facets.variant import ( + VariantFacet, + declared_variant_enums, + get_variant_enum, +) from invokeai.backend.architectures.registry import ( ArchitectureError, defs_module_path, @@ -56,12 +61,15 @@ "LatentSpace", "LatentSpaceFacet", "UNetDownscaleFacet", + "VariantFacet", + "declared_variant_enums", "defs_module_path", "facets_of", "generative_bases", "get", "get_latent_space", "get_max_unet_downscale", + "get_variant_enum", "register", "require", "resolve_latent_space", diff --git a/invokeai/backend/architectures/defs/flux.py b/invokeai/backend/architectures/defs/flux.py index 043cbf94fbe..9637be2cd42 100644 --- a/invokeai/backend/architectures/defs/flux.py +++ b/invokeai/backend/architectures/defs/flux.py @@ -1,8 +1,20 @@ from invokeai.backend.architectures.facets.latent_space import FLUX_16, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import ( + BaseModelType, + FluxVariantType, + ModelType, + PiDDecoderVariantType, +) register( BaseModelType.Flux, LatentSpaceFacet(FLUX_16), + VariantFacet( + { + ModelType.Main: FluxVariantType, + ModelType.PiDDecoder: PiDDecoderVariantType, + } + ), ) diff --git a/invokeai/backend/architectures/defs/flux2.py b/invokeai/backend/architectures/defs/flux2.py index e55b0f37879..b5b9e19c084 100644 --- a/invokeai/backend/architectures/defs/flux2.py +++ b/invokeai/backend/architectures/defs/flux2.py @@ -1,8 +1,22 @@ from invokeai.backend.architectures.facets.latent_space import FLUX2_32, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import ( + BaseModelType, + Flux2VariantType, + ModelType, + PiDDecoderVariantType, +) register( BaseModelType.Flux2, LatentSpaceFacet(FLUX2_32), + VariantFacet( + { + # FLUX.2 LoRAs are labelled with the same variant enum as the mains they target. + ModelType.Main: Flux2VariantType, + ModelType.LoRA: Flux2VariantType, + ModelType.PiDDecoder: PiDDecoderVariantType, + } + ), ) diff --git a/invokeai/backend/architectures/defs/krea_2.py b/invokeai/backend/architectures/defs/krea_2.py index c0b7524ddeb..9660eb7ce0b 100644 --- a/invokeai/backend/architectures/defs/krea_2.py +++ b/invokeai/backend/architectures/defs/krea_2.py @@ -1,10 +1,14 @@ from invokeai.backend.architectures.facets.latent_space import WAN21_16, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import BaseModelType, Krea2VariantType, ModelType register( BaseModelType.Krea2, # Krea-2 decodes with the Qwen-Image VAE, which is the Wan 2.1 VAE (16 latent channels), so it # shares the preview factors. LatentSpaceFacet(WAN21_16), + # The values are krea2_turbo / krea2_base rather than turbo / base, because variant strings are + # resolved without base context in configs/factory.py and so must be globally unique. + VariantFacet({ModelType.Main: Krea2VariantType}), ) diff --git a/invokeai/backend/architectures/defs/qwen_image.py b/invokeai/backend/architectures/defs/qwen_image.py index 7ea0c30c4e8..d60a9e96d36 100644 --- a/invokeai/backend/architectures/defs/qwen_image.py +++ b/invokeai/backend/architectures/defs/qwen_image.py @@ -1,9 +1,21 @@ from invokeai.backend.architectures.facets.latent_space import WAN21_16, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import ( + BaseModelType, + ModelType, + PiDDecoderVariantType, + QwenImageVariantType, +) register( BaseModelType.QwenImage, # Qwen-Image decodes with the 16-channel Wan 2.1 VAE. LatentSpaceFacet(WAN21_16), + VariantFacet( + { + ModelType.Main: QwenImageVariantType, + ModelType.PiDDecoder: PiDDecoderVariantType, + } + ), ) diff --git a/invokeai/backend/architectures/defs/sd_1.py b/invokeai/backend/architectures/defs/sd_1.py index 70b90d45062..2685e3f1704 100644 --- a/invokeai/backend/architectures/defs/sd_1.py +++ b/invokeai/backend/architectures/defs/sd_1.py @@ -1,10 +1,12 @@ from invokeai.backend.architectures.facets.latent_space import SD15_4, LatentSpaceFacet from invokeai.backend.architectures.facets.unet import UNetDownscaleFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import BaseModelType, ModelType, ModelVariantType register( BaseModelType.StableDiffusion1, LatentSpaceFacet(SD15_4), UNetDownscaleFacet(max_unet_downscale=8), + VariantFacet({ModelType.Main: ModelVariantType}), ) diff --git a/invokeai/backend/architectures/defs/sd_2.py b/invokeai/backend/architectures/defs/sd_2.py index 737b45daea3..3d8d8aa8b29 100644 --- a/invokeai/backend/architectures/defs/sd_2.py +++ b/invokeai/backend/architectures/defs/sd_2.py @@ -1,9 +1,11 @@ from invokeai.backend.architectures.facets.latent_space import SD15_4, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import BaseModelType, ModelType, ModelVariantType register( BaseModelType.StableDiffusion2, # SD2 shares SD1's 4-channel latent space and preview factors. LatentSpaceFacet(SD15_4), + VariantFacet({ModelType.Main: ModelVariantType}), ) diff --git a/invokeai/backend/architectures/defs/sd_3.py b/invokeai/backend/architectures/defs/sd_3.py index f7fa3b24988..a093efbdd99 100644 --- a/invokeai/backend/architectures/defs/sd_3.py +++ b/invokeai/backend/architectures/defs/sd_3.py @@ -1,8 +1,11 @@ from invokeai.backend.architectures.facets.latent_space import SD3_16, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import BaseModelType, ModelType, PiDDecoderVariantType register( BaseModelType.StableDiffusion3, LatentSpaceFacet(SD3_16), + # SD3 mains carry no variant; only its PiD decoder checkpoints do. + VariantFacet({ModelType.PiDDecoder: PiDDecoderVariantType}), ) diff --git a/invokeai/backend/architectures/defs/sdxl.py b/invokeai/backend/architectures/defs/sdxl.py index c5f9d96ff97..76fe80b836e 100644 --- a/invokeai/backend/architectures/defs/sdxl.py +++ b/invokeai/backend/architectures/defs/sdxl.py @@ -1,10 +1,22 @@ from invokeai.backend.architectures.facets.latent_space import SDXL_4, LatentSpaceFacet from invokeai.backend.architectures.facets.unet import UNetDownscaleFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import ( + BaseModelType, + ModelType, + ModelVariantType, + PiDDecoderVariantType, +) register( BaseModelType.StableDiffusionXL, LatentSpaceFacet(SDXL_4), UNetDownscaleFacet(max_unet_downscale=4), + VariantFacet( + { + ModelType.Main: ModelVariantType, + ModelType.PiDDecoder: PiDDecoderVariantType, + } + ), ) diff --git a/invokeai/backend/architectures/defs/sdxl_refiner.py b/invokeai/backend/architectures/defs/sdxl_refiner.py index 877413bf42a..4ec480d0e21 100644 --- a/invokeai/backend/architectures/defs/sdxl_refiner.py +++ b/invokeai/backend/architectures/defs/sdxl_refiner.py @@ -1,9 +1,11 @@ from invokeai.backend.architectures.facets.latent_space import SDXL_4, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import BaseModelType, ModelType, ModelVariantType register( BaseModelType.StableDiffusionXLRefiner, # The refiner shares SDXL's latent space, smooth matrix included. LatentSpaceFacet(SDXL_4), + VariantFacet({ModelType.Main: ModelVariantType}), ) diff --git a/invokeai/backend/architectures/defs/wan.py b/invokeai/backend/architectures/defs/wan.py index 0e9795ab4e7..d5f41a0acb9 100644 --- a/invokeai/backend/architectures/defs/wan.py +++ b/invokeai/backend/architectures/defs/wan.py @@ -1,6 +1,12 @@ from invokeai.backend.architectures.facets.latent_space import WAN21_16, WAN22_48, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import ( + BaseModelType, + ModelType, + WanLoRAVariantType, + WanVariantType, +) register( BaseModelType.Wan, @@ -8,4 +14,13 @@ # VAE at 8x spatial; TI2V-5B uses the 48-channel Wan2.2-VAE at 16x. The latent channel count # uniquely identifies the variant, which is how `LatentSpaceFacet.resolve()` tells them apart. LatentSpaceFacet(WAN21_16, alternates=(WAN22_48,)), + # The one architecture whose LoRAs carry a different variant enum from its mains. They are not + # interchangeable: an A14B LoRA (inner_dim=5120) against a TI2V-5B main (3072) crashes in the + # layer patcher, which is why the LoRA enum exists separately at all. + VariantFacet( + { + ModelType.Main: WanVariantType, + ModelType.LoRA: WanLoRAVariantType, + } + ), ) diff --git a/invokeai/backend/architectures/defs/z_image.py b/invokeai/backend/architectures/defs/z_image.py index 39e879a1f42..4ad0d0a0fa4 100644 --- a/invokeai/backend/architectures/defs/z_image.py +++ b/invokeai/backend/architectures/defs/z_image.py @@ -1,9 +1,16 @@ from invokeai.backend.architectures.facets.latent_space import FLUX_16, LatentSpaceFacet +from invokeai.backend.architectures.facets.variant import VariantFacet from invokeai.backend.architectures.registry import register -from invokeai.backend.model_manager.taxonomy import BaseModelType +from invokeai.backend.model_manager.taxonomy import BaseModelType, ModelType, ZImageVariantType register( BaseModelType.ZImage, # Z-Image uses a FLUX-compatible VAE with 16 latent channels. LatentSpaceFacet(FLUX_16), + VariantFacet( + { + ModelType.Main: ZImageVariantType, + ModelType.LoRA: ZImageVariantType, + } + ), ) diff --git a/invokeai/backend/architectures/facets/variant.py b/invokeai/backend/architectures/facets/variant.py new file mode 100644 index 00000000000..8ae945a1b7d --- /dev/null +++ b/invokeai/backend/architectures/facets/variant.py @@ -0,0 +1,69 @@ +"""Which variant enum labels an architecture's models, keyed by model type. + +Base alone does not determine it, which is why this facet carries a mapping rather than a single +enum: + +- Wan mains carry `WanVariantType`, Wan LoRAs carry `WanLoRAVariantType`. Same base, different + enum, and they are not interchangeable -- applying an A14B LoRA to a TI2V-5B main crashes in the + layer patcher. +- The PiD decoder's resolution presets are one enum shared across five bases. + +Two variant enums cannot live here at all: `ClipVariantType` and `Qwen3VariantType` sit on +`base=Any` configs, and `Any` is a sentinel the registry refuses to register. They are named +explicitly in `tests/backend/architectures/test_variants.py` so the completeness check against +`AnyVariant` stays total rather than quietly partial. + +This facet is deliberately *not* wired into `configs/factory.py`. That module validates a bare +variant string against `variant_type_adapter` without passing the base (`build_common_fields`), so +variant values have to be globally unique -- which is why `Krea2VariantType.Turbo` is +``krea2_turbo`` rather than ``turbo``. Making resolution base-aware would still need a fallback for +the `base=Any` models, so it would move the trap rather than close it, at the cost of touching the +identification path. The invariant is pinned by a test instead. +""" + +from collections.abc import Mapping +from dataclasses import dataclass +from enum import Enum + +from invokeai.backend.architectures.facet import Facet +from invokeai.backend.architectures.registry import generative_bases, get +from invokeai.backend.model_manager.taxonomy import BaseModelType, ModelType + + +@dataclass(frozen=True) +class VariantFacet(Facet): + """The variant enums an architecture's models are labelled with, by model type. + + Optional: four registered architectures (CogView4, ERNIE-Image, Ideogram 4, Anima) model no + variants at all, and declaring an empty facet would be indistinguishable from declaring nothing. + """ + + by_model_type: Mapping[ModelType, type[Enum]] + + def __post_init__(self) -> None: + if not self.by_model_type: + raise ValueError( + "VariantFacet declares no model types. An architecture without variants omits the facet entirely." + ) + + +def get_variant_enum(base: BaseModelType, model_type: ModelType) -> type[Enum] | None: + """The variant enum for this base and model type, or None if that combination has no variants.""" + facet = get(base, VariantFacet) + if facet is None: + return None + return facet.by_model_type.get(model_type) + + +def declared_variant_enums() -> frozenset[type[Enum]]: + """Every variant enum declared by any registered architecture. + + The registry side of the CI guards that keep `AnyVariant`, `variant_type_adapter` and + `ModelRecordChanges.variant` -- four hand-maintained copies of one list -- from drifting apart. + """ + enums: set[type[Enum]] = set() + for base in generative_bases(): + facet = get(base, VariantFacet) + if facet is not None: + enums.update(facet.by_model_type.values()) + return frozenset(enums) diff --git a/tests/backend/architectures/test_variants.py b/tests/backend/architectures/test_variants.py new file mode 100644 index 00000000000..b67bbbefd85 --- /dev/null +++ b/tests/backend/architectures/test_variants.py @@ -0,0 +1,212 @@ +"""CI guards for the variant plumbing. + +One list of variant enums is written by hand four times -- `AnyVariant`, the `variant_type_adapter` +subscript, the adapter's runtime argument (all three in `taxonomy.py`) and +`ModelRecordChanges.variant`. Nothing checked that they agreed, and nothing checked the invariant +they all rest on: that variant *values* are globally unique, because `configs/factory.py` resolves a +bare variant string without knowing the base. + +That invariant is real and has already cost something -- it is why `Krea2VariantType.Turbo` is +``krea2_turbo`` rather than ``turbo``. Until now it lived only in a docstring. +""" + +import typing +from enum import Enum + +import pytest + +from invokeai.backend.architectures import declared_variant_enums, get_variant_enum +from invokeai.backend.architectures.registry import _NOT_ARCHITECTURES +from invokeai.backend.model_manager import taxonomy +from invokeai.backend.model_manager.configs.base import Config_Base +from invokeai.backend.model_manager.configs.factory import AnyModelConfig # noqa: F401 (registers every config class) +from invokeai.backend.model_manager.taxonomy import ( + AnyVariant, + BaseModelType, + ClipVariantType, + ModelType, + Qwen3VariantType, + variant_type_adapter, +) + +BASE_AGNOSTIC_VARIANT_ENUMS = frozenset({ClipVariantType, Qwen3VariantType}) +"""Variant enums that cannot be declared by any architecture. + +Both sit on `base=Any` configs -- CLIP embedders and Qwen3 text encoders are components shared +across architectures, not architectures. `Any` is a sentinel the registry refuses to register, so +these are named here to keep the completeness check against `AnyVariant` total rather than quietly +partial. +""" + + +def _variant_enums_in_taxonomy() -> set[type[Enum]]: + """Every `*VariantType` enum defined in taxonomy.py. + + Discovered by name rather than listed, so that adding one and forgetting `AnyVariant` fails + here. `ModelRepoVariant` is deliberately not matched: it is the fp16/fp32 repo flavour, an + unrelated concept that is not part of `AnyVariant`. + """ + return { + obj + for name, obj in vars(taxonomy).items() + if name.endswith("VariantType") and isinstance(obj, type) and issubclass(obj, Enum) + } + + +def _enums_in_annotation(annotation: object) -> set[type[Enum]]: + """Every Enum class reachable from a type annotation, through Optional/Union/Literal.""" + origin = typing.get_origin(annotation) + if origin is typing.Literal: + return {type(arg) for arg in typing.get_args(annotation) if isinstance(arg, Enum)} + if origin is not None: + return set().union(*(_enums_in_annotation(arg) for arg in typing.get_args(annotation)), set()) + if isinstance(annotation, type) and issubclass(annotation, Enum): + return {annotation} + return set() + + +def _config_classes_with_a_variant() -> dict[tuple[BaseModelType, ModelType], set[type[Enum]]]: + """The (base, type) -> variant enum map as the model config classes actually declare it.""" + rows: dict[tuple[BaseModelType, ModelType], set[type[Enum]]] = {} + for config_class in Config_Base.CONFIG_CLASSES: + field = config_class.model_fields.get("variant") + if field is None: + continue + enums = _enums_in_annotation(field.annotation) + if not enums: + continue + key = (config_class.model_fields["base"].default, config_class.model_fields["type"].default) + rows.setdefault(key, set()).update(enums) + return rows + + +# --- the invariant configs/factory.py rests on ---------------------------------------------------- + + +def test_variant_values_are_globally_unique() -> None: + """`build_common_fields` validates a bare variant string against the union of every variant enum + without passing the base (configs/factory.py), so two enums sharing a value would silently + resolve to whichever the union tried first. + """ + seen: dict[str, str] = {} + collisions: list[str] = [] + for enum_class in sorted(_variant_enums_in_taxonomy(), key=lambda e: e.__name__): + for member in enum_class: + owner = seen.setdefault(member.value, f"{enum_class.__name__}.{member.name}") + if owner != f"{enum_class.__name__}.{member.name}": + collisions.append(f"{member.value!r}: {owner} vs {enum_class.__name__}.{member.name}") + + assert collisions == [], ( + "variant values must be globally unique -- see the note on Krea2VariantType.Turbo, which is " + f"'krea2_turbo' for exactly this reason. Collisions: {collisions}" + ) + + +def test_the_type_adapter_resolves_every_value_to_its_own_enum() -> None: + """The behavioural form of the check above, straight through the call factory.py makes.""" + wrong: list[str] = [] + for enum_class in sorted(_variant_enums_in_taxonomy(), key=lambda e: e.__name__): + for member in enum_class: + resolved = variant_type_adapter.validate_strings(member.value) + if type(resolved) is not enum_class: + wrong.append(f"{member.value!r} -> {type(resolved).__name__}, expected {enum_class.__name__}") + + assert wrong == [] + + +# --- the four hand-maintained copies of one list -------------------------------------------------- + + +def test_any_variant_covers_every_variant_enum_in_the_taxonomy() -> None: + assert set(typing.get_args(AnyVariant)) == _variant_enums_in_taxonomy() + + +def test_model_record_changes_covers_every_variant_enum() -> None: + from invokeai.app.services.model_records.model_records_base import ModelRecordChanges + + annotation = ModelRecordChanges.model_fields["variant"].annotation + + assert _enums_in_annotation(annotation) == _variant_enums_in_taxonomy() + + +def test_the_registry_and_the_allowlist_together_are_any_variant() -> None: + """Every variant enum is either declared by an architecture or explicitly base-agnostic. + + This is what makes the facet load-bearing rather than decorative: a new variant enum has to be + accounted for on one side or the other. + """ + assert declared_variant_enums() | BASE_AGNOSTIC_VARIANT_ENUMS == set(typing.get_args(AnyVariant)) + + +def test_the_allowlist_holds_only_genuinely_base_agnostic_enums() -> None: + """Guards against parking an architecture's enum in the allowlist to make the check pass.""" + for enum_class in BASE_AGNOSTIC_VARIANT_ENUMS: + bases = {base for (base, _type), enums in _config_classes_with_a_variant().items() if enum_class in enums} + assert bases <= set(_NOT_ARCHITECTURES), f"{enum_class.__name__} is used by real architectures: {bases}" + + +# --- the facet against reality -------------------------------------------------------------------- + + +def test_every_facet_declaration_matches_the_config_classes() -> None: + """Derived from `Config_Base.CONFIG_CLASSES`, so a wrong declaration cannot pass unnoticed. + + Both directions: a (base, type) the configs give a variant must be declared, and a declaration + the configs do not back must not exist. + """ + from_configs = { + (base, model_type): enums + for (base, model_type), enums in _config_classes_with_a_variant().items() + if base not in _NOT_ARCHITECTURES + } + + problems: list[str] = [] + for (base, model_type), enums in sorted( + from_configs.items(), key=lambda item: (item[0][0].value, item[0][1].value) + ): + assert len(enums) == 1, f"{base.value} x {model_type.value} declares several variant enums: {enums}" + expected = next(iter(enums)) + declared = get_variant_enum(base, model_type) + if declared is not expected: + name = declared.__name__ if declared else "nothing" + problems.append(f"{base.value} x {model_type.value}: registry says {name}, configs say {expected.__name__}") + + for base in BaseModelType: + if base in _NOT_ARCHITECTURES: + continue + for model_type in ModelType: + declared = get_variant_enum(base, model_type) + if declared is not None and (base, model_type) not in from_configs: + problems.append( + f"{base.value} x {model_type.value}: registry declares {declared.__name__}, " + f"but no config class for that combination has a variant field" + ) + + assert not problems, "VariantFacet declarations disagree with the model config classes:\n " + "\n ".join(problems) + + +@pytest.mark.parametrize( + ("base", "model_type", "expected"), + [ + # The two rows that a base-keyed facet could not express, spelled out. + (BaseModelType.Wan, ModelType.Main, "WanVariantType"), + (BaseModelType.Wan, ModelType.LoRA, "WanLoRAVariantType"), + # One enum, five bases. + (BaseModelType.Flux, ModelType.PiDDecoder, "PiDDecoderVariantType"), + (BaseModelType.StableDiffusion3, ModelType.PiDDecoder, "PiDDecoderVariantType"), + ], +) +def test_the_cases_that_motivate_the_model_type_dimension( + base: BaseModelType, model_type: ModelType, expected: str +) -> None: + declared = get_variant_enum(base, model_type) + + assert declared is not None + assert declared.__name__ == expected + + +def test_an_architecture_without_variants_declares_nothing() -> None: + # CogView4, ERNIE-Image, Ideogram 4 and Anima model no variants at all. An empty facet would be + # indistinguishable from an absent one, so they omit it. + for base in (BaseModelType.CogView4, BaseModelType.ErnieImage, BaseModelType.Ideogram4, BaseModelType.Anima): + assert get_variant_enum(base, ModelType.Main) is None