From f92f1a0443d23569311db641376e2726b351ca61 Mon Sep 17 00:00:00 2001 From: Alexander Eichhorn Date: Fri, 27 Feb 2026 01:03:39 +0100 Subject: [PATCH] Filter non-transformer keys from Z-Image checkpoint state dicts Merged Z-Image checkpoints (e.g. models with LoRAs baked in) may bundle text encoder weights (text_encoders.*) or other non-transformer keys alongside the transformer weights. These cause load_state_dict() to fail with strict=True. Instead of disabling strict mode, explicitly whitelist valid ZImageTransformer2DModel key prefixes and discard everything else. Also moves RAM allocation after filtering so it doesn't over-allocate for discarded keys. --- .../load/model_loaders/z_image.py | 31 +++++++++++++++---- 1 file changed, 25 insertions(+), 6 deletions(-) diff --git a/invokeai/backend/model_manager/load/model_loaders/z_image.py b/invokeai/backend/model_manager/load/model_loaders/z_image.py index aadced8f569..c381e02718d 100644 --- a/invokeai/backend/model_manager/load/model_loaders/z_image.py +++ b/invokeai/backend/model_manager/load/model_loaders/z_image.py @@ -253,16 +253,35 @@ def _load_from_singlefile( target_device = TorchDevice.choose_torch_device() model_dtype = TorchDevice.choose_bfloat16_safe_dtype(target_device) + # Filter out keys that don't belong to the ZImageTransformer2DModel. + # Merged checkpoints (e.g. LoRA-baked models) may bundle text encoder weights + # (text_encoders.*) or other non-transformer keys alongside the transformer weights. + # Also filter FP8 quantization metadata (scale_weight, scaled_fp8). + valid_prefixes = ( + "all_x_embedder.", + "all_final_layer.", + "layers.", + "noise_refiner.", + "context_refiner.", + "t_embedder.", + "cap_embedder.", + "rope_embedder.", + ) + valid_exact = {"x_pad_token", "cap_pad_token"} + keys_to_remove = [ + k + for k in sd.keys() + if not (k.startswith(valid_prefixes) or k in valid_exact) + or k.endswith(".scale_weight") + or k == "scaled_fp8" + ] + for k in keys_to_remove: + del sd[k] + # Handle memory management and dtype conversion new_sd_size = sum([ten.nelement() * model_dtype.itemsize for ten in sd.values()]) self._ram_cache.make_room(new_sd_size) - # Filter out FP8 scale_weight and scaled_fp8 metadata keys - # These are quantization metadata that shouldn't be loaded into the model - keys_to_remove = [k for k in sd.keys() if k.endswith(".scale_weight") or k == "scaled_fp8"] - for k in keys_to_remove: - del sd[k] - # Convert to target dtype for k in sd.keys(): sd[k] = sd[k].to(model_dtype)