Fix issue with layer offloading on ernie
This commit is contained in:
parent
e74bc9ac7b
commit
dd7074a21f
|
|
@ -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,
|
||||
],
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Reference in New Issue