diff --git a/src/pet/pet_model.py b/src/pet/pet_model.py index a865623..bcc97b3 100644 --- a/src/pet/pet_model.py +++ b/src/pet/pet_model.py @@ -134,6 +134,13 @@ class PETModel(torch.nn.Module): f"trainable params: {trainable_params} || all params: {all_param} || trainable%: {100 * trainable_params / all_param}" ) + def __getattr__(self, name: str): + """Forward missing attributes to the wrapped module.""" + try: + return super().__getattr__(name) # defer to nn.Module's logic + except AttributeError: + return getattr(self.base_model, name) + class PETModelForSequenceClassification(PETModel): """ diff --git a/src/pet/tuners/lora.py b/src/pet/tuners/lora.py index 9c3899a..d0e9c71 100644 --- a/src/pet/tuners/lora.py +++ b/src/pet/tuners/lora.py @@ -116,3 +116,10 @@ class LoRAModel(torch.nn.Module): def forward(self, *args, **kwargs): return self.model(*args, **kwargs) + + def __getattr__(self, name: str): + """Forward missing attributes to the wrapped module.""" + try: + return super().__getattr__(name) # defer to nn.Module's logic + except AttributeError: + return getattr(self.model, name)