From 7d9d47319377b12df02fb50d00d11f39d342344e Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 20 Aug 2025 20:36:51 -0700 Subject: [PATCH] Make Lumina2 compatible with EasyCache --- comfy/ldm/lumina/model.py | 10 +++++++++- comfy_extras/nodes_easycache.py | 2 ++ 2 files changed, 11 insertions(+), 1 deletion(-) diff --git a/comfy/ldm/lumina/model.py b/comfy/ldm/lumina/model.py index f8dc4d7db..e08ed817d 100644 --- a/comfy/ldm/lumina/model.py +++ b/comfy/ldm/lumina/model.py @@ -11,6 +11,7 @@ import comfy.ldm.common_dit from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder from comfy.ldm.modules.attention import optimized_attention_masked from comfy.ldm.flux.layers import EmbedND +import comfy.patcher_extension def modulate(x, scale): @@ -590,8 +591,15 @@ class NextDiT(nn.Module): return padded_full_embed, mask, img_sizes, l_effective_cap_len, freqs_cis - # def forward(self, x, t, cap_feats, cap_mask): def forward(self, x, timesteps, context, num_tokens, attention_mask=None, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, kwargs.get("transformer_options", {})) + ).execute(x, timesteps, context, num_tokens, attention_mask, **kwargs) + + # def forward(self, x, t, cap_feats, cap_mask): + def _forward(self, x, timesteps, context, num_tokens, attention_mask=None, **kwargs): t = 1.0 - timesteps cap_feats = context cap_mask = attention_mask diff --git a/comfy_extras/nodes_easycache.py b/comfy_extras/nodes_easycache.py index fd3c87926..80e7e8fc3 100644 --- a/comfy_extras/nodes_easycache.py +++ b/comfy_extras/nodes_easycache.py @@ -13,6 +13,8 @@ def easycache_forward_wrapper(executor, *args, **kwargs): # get values from args x: torch.Tensor = args[0] transformer_options: dict[str] = args[-1] + if not isinstance(transformer_options, dict): + transformer_options = kwargs.get("transformer_options") easycache: EasyCacheHolder = transformer_options["easycache"] sigmas = transformer_options["sigmas"] uuids = transformer_options["uuids"]