Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 25 additions & 6 deletions invokeai/backend/model_manager/load/model_loaders/z_image.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down