From e63f47ca520a95ddc9a91ec479efa6451472842c Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Thu, 1 Dec 2022 21:15:01 +0530 Subject: [PATCH] fix --- src/pet/pet_model.py | 4 ++-- src/pet/utils/other.py | 12 ++++++++++++ 2 files changed, 14 insertions(+), 2 deletions(-) diff --git a/src/pet/pet_model.py b/src/pet/pet_model.py index bcc97b3..fff3dd3 100644 --- a/src/pet/pet_model.py +++ b/src/pet/pet_model.py @@ -46,10 +46,10 @@ class PETModel(torch.nn.Module): num_transformer_submodules = 0 transformer_backbone = None for name, module in self.base_model.named_children(): + for param in module.parameters(): + param.requires_grad = False if isinstance(module, PreTrainedModel): # Make sure to freeze Tranformers model - for param in module.parameters(): - param.requires_grad = False if transformer_backbone is None: transformer_backbone = module self.transformer_backbone_name = name diff --git a/src/pet/utils/other.py b/src/pet/utils/other.py index c6274db..a13b4f6 100644 --- a/src/pet/utils/other.py +++ b/src/pet/utils/other.py @@ -44,3 +44,15 @@ def _set_trainable(model): param.requires_grad = True else: param.requires_grad = False + + +# def fsdp_auto_wrap_policy(): +# from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy, lambda_auto_wrap_policy, _or_policy + +# def lambda_policy(module): +# if len(module.named_children()) != 0 and +# return True +# return False + + +# pass