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,