From 273b35378f7ddebe1bcde251a156ad7576f02ee5 Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 13 Jan 2025 16:45:26 +0000 Subject: [PATCH] custom ctx_encoder! --- configs/fw_qa_tiny.yaml | 2 ++ hyperlora/configs.py | 17 ++++++++++++++--- hyperlora/data_utils.py | 23 ++++++++++++++++++++--- hyperlora/intx_sft.py | 33 ++++++++++++++++++++++++++++----- hyperlora/model_loading.py | 17 ++++++++++++++--- hyperlora/modeling_utils.py | 19 ++++++++++++++++--- hyperlora/training_utils.py | 12 ++++++------ 7 files changed, 100 insertions(+), 23 deletions(-) diff --git a/configs/fw_qa_tiny.yaml b/configs/fw_qa_tiny.yaml index f3d359d..ed49e67 100644 --- a/configs/fw_qa_tiny.yaml +++ b/configs/fw_qa_tiny.yaml @@ -27,6 +27,8 @@ neftune_noise_alpha: 1 weight_decay: 0.01 warmup_ratio: 0.1 +dataloader_prefetch_factor: 16 +dataloader_num_workers: 16 # LoRA lora_r: 8 lora_dropout: 0.05 diff --git a/hyperlora/configs.py b/hyperlora/configs.py index 3cbc430..8526615 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -252,6 +252,10 @@ class CtxTrainingArguments: default=2**13, metadata={"help": "Maximum base length for training."}, ) + max_ctx_len: Optional[int] = field( + default=2**13, + metadata={"help": "Maximum context length for training."}, + ) max_new_tokens: Optional[int] = field( default=2**10, metadata={"help": "Maximum new tokens for generation-based evaluation."}, @@ -320,9 +324,16 @@ class HypernetArguments: @dataclass class CtxEncoderArguments: - layer_idx: int = field( - default=4, - metadata={"help": "Layer index for context encoder."}, + ctx_encoder_model_name_or_path: str = field( + default=None, + metadata={"help": "Context encoder model name or path."}, + ) + layer_idx: Optional[int] = field( + default=None, + metadata={ + "help": "Layer index for context encoder. " + "Default to L//4 where L is the number of layers of the ctx model" + }, ) diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index 0061725..f70add5 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -171,6 +171,8 @@ def get_tokenized_dataset( split: str, tokenizer: PreTrainedTokenizerBase, tokenizer_kwargs: dict[str, Any], + ctx_tokenizer: PreTrainedTokenizerBase, + ctx_tokenizer_kwargs: dict[str, Any], add_ctx_to_chat: bool, add_repeat_prompt: bool, add_negative_prompt: bool, @@ -199,13 +201,27 @@ def get_tokenized_dataset( if add_repeat_prompt and "context_numbers" not in ds_name: ds = ds.map(add_repeat_prompt_fn, batched=True, batch_size=None) tokenized_ds = construct_and_tokenize_ctx_qa( - tokenizer, tokenizer_kwargs, add_ctx_to_chat, use_kl_loss, need_ctx_ids, ds + tokenizer, + tokenizer_kwargs, + ctx_tokenizer, + ctx_tokenizer_kwargs, + add_ctx_to_chat, + use_kl_loss, + need_ctx_ids, + ds, ) return tokenized_ds def construct_and_tokenize_ctx_qa( - tokenizer, tokenizer_kwargs, add_ctx_to_chat, use_kl_loss, need_ctx_ids, ds + tokenizer, + tokenizer_kwargs, + ctx_tokenizer, + ctx_tokenizer_kwargs, + add_ctx_to_chat, + use_kl_loss, + need_ctx_ids, + ds, ): # for sft + chat_model, we need to convert the dataset to chat format # add "messages" field @@ -229,6 +245,7 @@ def construct_and_tokenize_ctx_qa( # for use_kl_loss, we need "chat_ids" and "chat_attn_mask" if use_kl_loss: + raise NotImplementedError("KL loss deprecated") tokenized_ds = tokenized_ds.map( convert_ctx_prompt_response_to_messages, fn_kwargs={"add_ctx_to_chat": True}, @@ -255,7 +272,7 @@ def construct_and_tokenize_ctx_qa( # tokenize the ctx_text to get ctx_ids and ctx_attn_mask tokenized_ds = tokenized_ds.map( tokenize_ctx_text, - fn_kwargs={"tokenizer": tokenizer}, + fn_kwargs={"tokenizer": ctx_tokenizer}, batched=True, num_proc=16, ) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index a7ff2fb..e5896cd 100755 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -262,16 +262,33 @@ def main(): requires_grad=ctx_args.exp_setup == ExperimentSetup.FULL_FINETUNE, peft_config=get_lora_config(model_name, **vars(lora_args)), ) + if ctx_encoder_args.ctx_encoder_model_name_or_path is not None: + ctx_encoder_model, ctx_tokenizer = get_model_and_tokenizer( + ctx_encoder_args.ctx_encoder_model_name_or_path, + train=False, + requires_grad=False, + ) + else: + ctx_encoder_model = model + ctx_tokenizer = tokenizer if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA: logger.info("Using HyperLoRA") - hypernet_config = get_hypernet_config(model, hypernet_args, aggregator_args) + hypernet_config = get_hypernet_config( + model, ctx_encoder_model, hypernet_args, aggregator_args + ) # hypernet = HyperLoRA( # get_hypernet_config(model, hypernet_args, aggregator_args), # model, # ).to(model.device) # ctx_encoder = EarlyExit(get_base_model(model), ctx_encoder_args.layer_idx) + if ctx_encoder_args.layer_idx is None: + ctx_encoder_args.layer_idx = ctx_encoder_model.config.num_hidden_layers // 4 + logger.info( + f"Using the first {ctx_encoder_args.layer_idx} layers" + " as the context encoder" + ) model = ModulatedPretrainedModel( model, hypernet_config, @@ -293,10 +310,13 @@ def main(): add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel) tokenizer_kwargs = {"max_length": ctx_args.max_base_len} + ctx_tokenizer_kwargs = {"max_length": ctx_args.max_ctx_len} # not used for now _get_tokenized_dataset = partial( get_tokenized_dataset, tokenizer=tokenizer, tokenizer_kwargs=tokenizer_kwargs, + ctx_tokenizer=ctx_tokenizer, + ctx_tokenizer_kwargs=ctx_tokenizer_kwargs, add_ctx_to_chat=add_ctx_to_chat, add_repeat_prompt=ctx_args.add_repeat_prompt, add_negative_prompt=ctx_args.add_negative_prompt, @@ -471,6 +491,9 @@ def main(): if isinstance(model, ModulatedPretrainedModel): logger.info("Applying liger-kernel to ModulatedPretrainedModel") _apply_liger_kernel_to_instance(model=model.base_model.base_model.model) + if ctx_encoder_args.ctx_encoder_model_name_or_path is not None: + logger.info("Applying liger-kernel to ctx_encoder_model") + _apply_liger_kernel_to_instance(model=ctx_encoder_model) elif isinstance(model, PeftModel): logger.info("Applying liger-kernel to PeftModel") _apply_liger_kernel_to_instance(model=model.base_model.model) @@ -485,21 +508,21 @@ def main(): train_model( model, - tokenizer, + # tokenizer, training_args, train_ds, val_ds, test_ds, partial(train_collator, tokenizer=tokenizer), - partial(generation_collator, tokenizer=tokenizer), + # partial(generation_collator, tokenizer=tokenizer), compute_metrics=partial( compute_metrics, evaluator=Evaluator( [compute_per_token_acc, compute_prefix_matching, compute_perplexity] ), ), - max_new_tokens=ctx_args.max_new_tokens, - gen_per_device_eval_batch_size=ctx_args.gen_per_device_eval_batch_size, + # max_new_tokens=ctx_args.max_new_tokens, + # gen_per_device_eval_batch_size=ctx_args.gen_per_device_eval_batch_size, ) logger.info(f"Training run finished and saved to {output_dir}") diff --git a/hyperlora/model_loading.py b/hyperlora/model_loading.py index 465f516..d19ab42 100644 --- a/hyperlora/model_loading.py +++ b/hyperlora/model_loading.py @@ -5,7 +5,12 @@ import torch from peft import LoraConfig, PeftConfig, PeftModel, VeraConfig from peft import get_peft_config as _get_peft_config from peft.utils import PeftType -from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer +from transformers import ( + AutoModel, + AutoModelForCausalLM, + AutoTokenizer, + MllamaForConditionalGeneration, +) logger = logging.getLogger() @@ -118,15 +123,21 @@ def get_model( # load_in_4bit=True, # load_in_8bit=True, ) + is_vision_model = "Llama" in model_name_or_path and "Vision" in model_name_or_path if model_kwargs is not None: model_init_kwargs.update(model_kwargs) if use_flash_attn: model_init_kwargs["attn_implementation"] = "flash_attention_2" + if is_vision_model: + model_init_kwargs["attn_implementation"] = "sdpa" # for training disable cache - if train: + if train and not is_vision_model: model_init_kwargs["use_cache"] = False logger.debug(f"Model init kwargs: {model_init_kwargs}") - model = AutoModelForCausalLM.from_pretrained(**model_init_kwargs) + if not is_vision_model: + model = AutoModelForCausalLM.from_pretrained(**model_init_kwargs) + else: + model = MllamaForConditionalGeneration.from_pretrained(**model_init_kwargs) if peft_config is not None: model = PeftModel(model, peft_config) model.train(train) diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index 80a34e5..c6ce671 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -70,11 +70,14 @@ class AggregatorConfig: def get_aggregator_config( - model: PreTrainedModel, output_size: int, aggregator_args: AggregatorArguments + model: PreTrainedModel, + ctx_encoder_model: PreTrainedModel, + output_size: int, + aggregator_args: AggregatorArguments, ): lora_config = model.peft_config["default"] return AggregatorConfig( - feature_size=model.config.hidden_size, + feature_size=ctx_encoder_model.config.hidden_size, output_size=output_size, num_layers=get_num_layers(model), num_modules=len(lora_config.target_modules), @@ -97,6 +100,7 @@ class HypernetConfig: def get_hypernet_config( model: PreTrainedModel, + ctx_encoder_model: PreTrainedModel, hypernet_args: HypernetArguments, aggregator_args: AggregatorArguments, ): @@ -110,6 +114,7 @@ def get_hypernet_config( feature_sizes=get_peft_in_out_features(model, peft_config=lora_config), aggregator_config=get_aggregator_config( model, + ctx_encoder_model, hypernet_args.latent_size, aggregator_args, ), @@ -577,8 +582,16 @@ class ModulatedPretrainedModel(nn.Module): def _init_model(self): self.hypernet = HyperLoRA(self.hypernet_config).to(self.device) # TODO: allow ctx_encoder to be other models + if self.ctx_encoder_args.ctx_encoder_model_name_or_path is not None: + encoder_model = get_model( + self.ctx_encoder_args.ctx_encoder_model_name_or_path, + train=True, + requires_grad=False, + ) + else: + encoder_model = self.base_model self.ctx_encoder = EarlyExit( - get_base_model(self.base_model), self.ctx_encoder_args.layer_idx + get_base_model(encoder_model), self.ctx_encoder_args.layer_idx ) def to(self, *args, **kwargs): diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 287c32d..b6a67ba 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -108,17 +108,17 @@ def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs): def train_model( model, - tokenizer, + # tokenizer, training_args, train_dataset=None, val_dataset=None, test_dataset=None, train_collator=None, - generation_collator=None, + # generation_collator=None, compute_metrics=None, - preprocess_logits_for_metrics=None, - max_new_tokens=2**13, - gen_per_device_eval_batch_size=1, + # preprocess_logits_for_metrics=None, + # max_new_tokens=2**13, + # gen_per_device_eval_batch_size=1, ): checkpoint = None if training_args.resume_from_checkpoint is not None: @@ -132,7 +132,7 @@ def train_model( eval_dataset=val_dataset, data_collator=train_collator, compute_metrics=compute_metrics, - preprocess_logits_for_metrics=preprocess_logits_for_metrics, + # preprocess_logits_for_metrics=preprocess_logits_for_metrics, ) # Trainer loads the best model after training