Skip to content
Draft
Show file tree
Hide file tree
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
20 changes: 19 additions & 1 deletion invokeai/app/invocations/krea2_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
import math
from contextlib import ExitStack
from pathlib import Path
from typing import Callable, Iterator, Optional
from typing import Any, Callable, Iterator, Optional

import torch
import torchvision.transforms as tv_transforms
Expand Down Expand Up @@ -46,6 +46,7 @@
from invokeai.backend.patches.layer_patcher import LayerPatcher, PatchSpec
from invokeai.backend.patches.lora_conversions.krea2_lora_constants import KREA2_LORA_TRANSFORMER_PREFIX
from invokeai.backend.patches.model_patch_raw import ModelPatchRaw
from invokeai.backend.quantization.int8_convrot import Int8ConvrotLinear
from invokeai.backend.rectified_flow.rectified_flow_inpaint_extension import RectifiedFlowInpaintExtension
from invokeai.backend.stable_diffusion.diffusers_pipeline import PipelineIntermediateState
from invokeai.backend.stable_diffusion.diffusion.conditioning_data import Krea2ConditioningInfo
Expand All @@ -55,6 +56,20 @@
KREA2_LATENT_CHANNELS = 16


def requires_sidecar_patching(transformer: Any, model_format: ModelFormat) -> bool:
"""Whether LoRA has to be applied as a sidecar rather than written into the weights.

The format alone does not answer this. A plain ``checkpoint`` Krea-2 may still be an
``int8_tensorwise`` build, whose Linears the loader replaced with ``Int8ConvrotLinear`` --
those hold their weights as int8 buffers, which a direct patch cannot write into (and which
could not represent the patched values anyway, the rotation having mixed 256 of them). So the
loaded module tree is consulted, not just the config.
"""
if model_format in (ModelFormat.GGUFQuantized,):
return True
return any(isinstance(module, Int8ConvrotLinear) for module in transformer.modules())


@invocation(
"krea2_denoise",
title="Denoise - Krea-2",
Expand Down Expand Up @@ -409,6 +424,7 @@ def _run_diffusion(self, context: InvocationContext):
)

transformer_config = context.models.get_config(self.transformer.transformer)
# Refined against the loaded module tree below, once the transformer is in hand.
model_is_quantized = transformer_config.format in (ModelFormat.GGUFQuantized,)
num_train_timesteps = scheduler.config.num_train_timesteps

Expand Down Expand Up @@ -439,6 +455,8 @@ def _run_diffusion(self, context: InvocationContext):
# SDPA for enable_gqa=True, which PyTorch only supports on the math backend — that materializes the
# full O(seq^2) score matrix (~5.7 GB per attention at 1280x720, ~40 GB at 2560x1440) and OOMs. Swap
# in a memory-efficient processor that expands the KV heads and uses the O(seq) SDPA kernel instead.
model_is_quantized = requires_sidecar_patching(transformer, transformer_config.format)

regional_prompting_state = Krea2RegionalPromptingState()
transformer.set_attn_processor(build_krea2_attention_processors(transformer, regional_prompting_state))
# The processors remain installed on the cached transformer after this invocation. Do not let them
Expand Down
2 changes: 1 addition & 1 deletion invokeai/app/invocations/minimax_h3_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,6 @@
from invokeai.app.services.session_processor.session_processor_common import CanceledException
from invokeai.app.services.shared.invocation_context import InvocationContext
from invokeai.backend.minimax_h3.denoise import denoise
from invokeai.backend.minimax_h3.int8_convrot import Int8ConvrotLinear
from invokeai.backend.minimax_h3.packing import (
MINIMAX_H3_CANVAS_MULTIPLE,
MINIMAX_H3_FPS,
Expand Down Expand Up @@ -75,6 +74,7 @@
is_minimax_h3_adaln_layer_path,
)
from invokeai.backend.patches.model_patch_raw import ModelPatchRaw
from invokeai.backend.quantization.int8_convrot import Int8ConvrotLinear
from invokeai.backend.stable_diffusion.diffusers_pipeline import PipelineIntermediateState
from invokeai.backend.stable_diffusion.diffusion.conditioning_data import MiniMaxH3ConditioningInfo
from invokeai.backend.util.devices import TorchDevice
Expand Down
144 changes: 0 additions & 144 deletions invokeai/backend/minimax_h3/int8_convrot.py

This file was deleted.

Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
import torch

from invokeai.backend.minimax_h3.int8_convrot import Int8ConvrotLinear
from invokeai.backend.model_manager.load.model_cache.torch_module_autocast.cast_to_device import cast_to_device
from invokeai.backend.model_manager.load.model_cache.torch_module_autocast.custom_modules.custom_linear import (
autocast_linear_forward_sidecar_patches,
Expand All @@ -9,6 +8,7 @@
CustomModuleMixin,
)
from invokeai.backend.patches.layers.param_shape_utils import get_param_shape
from invokeai.backend.quantization.int8_convrot import Int8ConvrotLinear


class CustomInt8ConvrotLinear(Int8ConvrotLinear, CustomModuleMixin):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@
from diffusers.models.normalization import RMSNorm as DiffusersRMSNorm

from invokeai.backend.flux.modules.layers import RMSNorm as FluxRMSNorm
from invokeai.backend.minimax_h3.int8_convrot import Int8ConvrotLinear
from invokeai.backend.model_manager.load.model_cache.torch_module_autocast.custom_modules.custom_conv1d import (
CustomConv1d,
)
Expand Down Expand Up @@ -35,6 +34,7 @@
from invokeai.backend.model_manager.load.model_cache.torch_module_autocast.custom_modules.custom_module_mixin import (
CustomModuleMixin,
)
from invokeai.backend.quantization.int8_convrot import Int8ConvrotLinear

AUTOCAST_MODULE_TYPE_MAPPING: dict[type[torch.nn.Module], type[torch.nn.Module]] = {
torch.nn.Linear: CustomLinear,
Expand Down
Loading
Loading