From 8920caeb1b3986d9c172a5e8cdbd518a5a80f465 Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Thu, 1 Dec 2022 15:28:45 +0530 Subject: [PATCH] fix --- src/pet/pet_model.py | 7 +++++++ src/pet/tuners/lora.py | 7 +++++++ 2 files changed, 14 insertions(+) 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)