From dd7074a21fd3e6424391a500654acd97e06ca0a8 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Tue, 14 Apr 2026 19:19:02 -0600 Subject: [PATCH] Fix issue with layer offloading on ernie --- .../diffusion_models/ernie_image/ernie_image.py | 3 +-- .../diffusion_models/ernie_image/transformer.py | 6 ++++++ 2 files changed, 7 insertions(+), 2 deletions(-) diff --git a/extensions_built_in/diffusion_models/ernie_image/ernie_image.py b/extensions_built_in/diffusion_models/ernie_image/ernie_image.py index c824869e..665cd6ff 100644 --- a/extensions_built_in/diffusion_models/ernie_image/ernie_image.py +++ b/extensions_built_in/diffusion_models/ernie_image/ernie_image.py @@ -110,8 +110,7 @@ class ErnieImageModel(BaseModel): self.device_torch, offload_percent=self.model_config.layer_offloading_transformer_percent, ignore_modules=[ - transformer.x_pad_token, - transformer.cap_pad_token, + transformer.x_embedder, ], ) diff --git a/extensions_built_in/diffusion_models/ernie_image/transformer.py b/extensions_built_in/diffusion_models/ernie_image/transformer.py index fbd850bb..3d27efee 100644 --- a/extensions_built_in/diffusion_models/ernie_image/transformer.py +++ b/extensions_built_in/diffusion_models/ernie_image/transformer.py @@ -338,6 +338,12 @@ class ErnieImageTransformer2DModel(ModelMixin, ConfigMixin): nn.init.zeros_(self.final_linear.weight) nn.init.zeros_(self.final_linear.bias) self.gradient_checkpointing = False + self.onload_device = None + + @property + def device(self): + # use self.x_embeddersince we ignore it in memory management + return next(self.x_embedder.parameters()).device def forward( self,