From e20b42e84aeecbefd8cae565fa3b51216bf23d39 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Tue, 23 Jun 2026 10:43:36 -0600 Subject: [PATCH] Add offloading support for krea2 --- .../diffusion_models/krea2/krea2.py | 33 ++++++++++++++++++- 1 file changed, 32 insertions(+), 1 deletion(-) diff --git a/extensions_built_in/diffusion_models/krea2/krea2.py b/extensions_built_in/diffusion_models/krea2/krea2.py index fa68819e..521affd3 100644 --- a/extensions_built_in/diffusion_models/krea2/krea2.py +++ b/extensions_built_in/diffusion_models/krea2/krea2.py @@ -39,8 +39,14 @@ from toolkit.samplers.custom_flowmatch_sampler import ( from toolkit.accelerator import unwrap_model from toolkit.metadata import get_meta_for_safetensors from toolkit.util.quantize import quantize, get_qtype, quantize_model +from toolkit.memory_management import MemoryManager -from .src.mmdit import SingleStreamDiT, SingleMMDiTConfig +from .src.mmdit import ( + DoubleSharedModulation, + SimpleModulation, + SingleMMDiTConfig, + SingleStreamDiT, +) from .src.text_encoder import encode_krea_prompt, SELECT_LAYERS from .src.pipeline import Krea2Pipeline, pad_text_features, predict_velocity @@ -234,6 +240,21 @@ class Krea2Model(BaseModel): quantize_model(self, transformer) flush() + if ( + self.model_config.layer_offloading + and self.model_config.layer_offloading_transformer_percent > 0 + ): + MemoryManager.attach( + transformer, + self.device_torch, + offload_percent=self.model_config.layer_offloading_transformer_percent, + ignore_modules=[ + module + for module in transformer.modules() + if isinstance(module, (SimpleModulation, DoubleSharedModulation)) + ], + ) + if self.model_config.low_vram: self.print_and_status_update("Moving transformer to CPU") transformer.to("cpu") @@ -248,6 +269,16 @@ class Krea2Model(BaseModel): quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te)) freeze(text_encoder) flush() + if ( + self.model_config.layer_offloading + and self.model_config.layer_offloading_text_encoder_percent > 0 + ): + MemoryManager.attach( + text_encoder, + self.device_torch, + offload_percent=self.model_config.layer_offloading_text_encoder_percent, + ) + if self.model_config.low_vram: self.print_and_status_update("Moving text encoder to CPU") text_encoder.to("cpu")