From b55b9a56439e418bbcb290600f234a9b315980c6 Mon Sep 17 00:00:00 2001 From: 51616 Date: Wed, 8 Jan 2025 19:33:08 +0000 Subject: [PATCH] workaround to work with deepspeed (avoid dtype cast of hypernet) --- hyperlora/intx_sft.py | 2 ++ hyperlora/model_loading.py | 4 ++-- hyperlora/modeling_utils.py | 10 +++++++++- 3 files changed, 13 insertions(+), 3 deletions(-) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index a52d16b..99eb191 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -404,6 +404,8 @@ def main(output_dir): # HACK: see transformers/trainer.py for liger-kernel patch # slows down training speed w/ short inputs # might improve/decrease training speed w/ longer inputs + # HACK: see transformers/trainer_seq2seq.py for supressing + # "Trainer.tokenizer is now deprecated. You should use Trainer.processing_class instead." wandb.init( project=os.getenv("WANDB_PROJECT"), diff --git a/hyperlora/model_loading.py b/hyperlora/model_loading.py index 79842b0..465f516 100644 --- a/hyperlora/model_loading.py +++ b/hyperlora/model_loading.py @@ -18,7 +18,7 @@ def get_model_and_tokenizer( peft_config=None, model_kwargs=None, tokenizer_kwargs=None, - device="cuda:0", + device="cuda", dtype=torch.bfloat16, ): model = get_model( @@ -106,7 +106,7 @@ def get_model( use_flash_attn=True, peft_config=None, model_kwargs=None, - device="cuda:0", + device="cuda", dtype=torch.bfloat16, ): model_init_kwargs = dict( diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index 92a8cca..467706c 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -153,7 +153,7 @@ class Perceiver(nn.Module): ctx_features: Float[Tensor, "bs seq_len feature_dim"], ctx_attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None, ): - x = self.perceiver(ctx_features, ctx_attn_mask).logits # .last_hidden_state + x = self.perceiver(ctx_features, ctx_attn_mask).logits x = rearrange( x, "bs (n_layers n_modules) d -> bs n_layers n_modules d", @@ -550,6 +550,14 @@ class ModulatedPretrainedModel(nn.Module): get_base_model(self.base_model), self.ctx_encoder_args.layer_idx ) + def to(self, *args, **kwargs): + # workaround to avoid the hypernet being wrapped by DeepSpeed + self.base_model = self.base_model.to(*args, **kwargs) + self.ctx_encoder = self.ctx_encoder.to(*args, **kwargs) + # self.hypernet = self.hypernet.to(*args, **kwargs) + self.hypernet.to(torch.float32) + return self + # Delegate to base_model @property def config(self):