From 2381b6a93ec72cac6f20142b01f3a6e435c839ce Mon Sep 17 00:00:00 2001 From: 51616 Date: Wed, 8 Jan 2025 16:17:07 +0000 Subject: [PATCH] add repeat prompt --- configs/context_numbers_10.yaml | 1 + configs/context_numbers_128.yaml | 1 + configs/context_numbers_128_only.yaml | 1 + configs/context_numbers_256.yaml | 1 + configs/context_numbers_256_only.yaml | 1 + configs/context_numbers_32.yaml | 1 + configs/context_numbers_32_only.yaml | 1 + configs/context_numbers_64.yaml | 1 + configs/context_numbers_64_only.yaml | 1 + configs/context_numbers_debug.yaml | 1 + configs/context_numbers_easy.yaml | 1 + configs/default.yaml | 1 + configs/pwc.yaml | 2 +- configs/pwc_and_ctx_numbers_256.yaml | 3 ++- hyperlora/configs.py | 4 ++++ hyperlora/data_utils.py | 17 +++++++++++++++++ hyperlora/intx_sft.py | 1 + 17 files changed, 37 insertions(+), 2 deletions(-) diff --git a/configs/context_numbers_10.yaml b/configs/context_numbers_10.yaml index 4106235..157e51f 100644 --- a/configs/context_numbers_10.yaml +++ b/configs/context_numbers_10.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_128.yaml b/configs/context_numbers_128.yaml index ba88b13..b57fb2f 100644 --- a/configs/context_numbers_128.yaml +++ b/configs/context_numbers_128.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_128_only.yaml b/configs/context_numbers_128_only.yaml index 258d709..2cd5f1d 100644 --- a/configs/context_numbers_128_only.yaml +++ b/configs/context_numbers_128_only.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_256.yaml b/configs/context_numbers_256.yaml index 0d124f6..f36f4f4 100644 --- a/configs/context_numbers_256.yaml +++ b/configs/context_numbers_256.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_256_only.yaml b/configs/context_numbers_256_only.yaml index 8e95d3d..ce3f84b 100644 --- a/configs/context_numbers_256_only.yaml +++ b/configs/context_numbers_256_only.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_32.yaml b/configs/context_numbers_32.yaml index a15dc20..81a329e 100644 --- a/configs/context_numbers_32.yaml +++ b/configs/context_numbers_32.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_32_only.yaml b/configs/context_numbers_32_only.yaml index a542023..7bf48d8 100644 --- a/configs/context_numbers_32_only.yaml +++ b/configs/context_numbers_32_only.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_64.yaml b/configs/context_numbers_64.yaml index 482c9cb..ed6f48e 100644 --- a/configs/context_numbers_64.yaml +++ b/configs/context_numbers_64.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_64_only.yaml b/configs/context_numbers_64_only.yaml index d84211a..082421d 100644 --- a/configs/context_numbers_64_only.yaml +++ b/configs/context_numbers_64_only.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 diff --git a/configs/context_numbers_debug.yaml b/configs/context_numbers_debug.yaml index 70f315a..c9d3102 100644 --- a/configs/context_numbers_debug.yaml +++ b/configs/context_numbers_debug.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false eval_on_start: True eval_strategy: "steps" eval_steps: 500 diff --git a/configs/context_numbers_easy.yaml b/configs/context_numbers_easy.yaml index 6d8c60c..fe392b9 100644 --- a/configs/context_numbers_easy.yaml +++ b/configs/context_numbers_easy.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false eval_on_start: True eval_strategy: "steps" eval_steps: 500 diff --git a/configs/default.yaml b/configs/default.yaml index 15f3035..8bf7fa8 100644 --- a/configs/default.yaml +++ b/configs/default.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false eval_on_start: True eval_strategy: "steps" eval_steps: 500 diff --git a/configs/pwc.yaml b/configs/pwc.yaml index 535d527..ef7b2af 100644 --- a/configs/pwc.yaml +++ b/configs/pwc.yaml @@ -18,7 +18,7 @@ label_names: ["labels"] per_device_train_batch_size: 32 per_device_eval_batch_size: 32 -max_val_samples_per_ds: 20 +max_val_samples_per_ds: 1000 # optim: schedule_free_adamw learning_rate: 0.00001 # lr_scheduler_type: "constant_with_warmup" diff --git a/configs/pwc_and_ctx_numbers_256.yaml b/configs/pwc_and_ctx_numbers_256.yaml index cf14a24..4b1609b 100644 --- a/configs/pwc_and_ctx_numbers_256.yaml +++ b/configs/pwc_and_ctx_numbers_256.yaml @@ -2,6 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] +add_repeat_prompt: false # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 @@ -18,7 +19,7 @@ label_names: ["labels"] per_device_train_batch_size: 32 per_device_eval_batch_size: 1 -max_val_samples_per_ds: 20 +max_val_samples_per_ds: 1000 # optim: schedule_free_adamw learning_rate: 0.00001 # lr_scheduler_type: "constant_with_warmup" diff --git a/hyperlora/configs.py b/hyperlora/configs.py index ead6daa..57280c4 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -235,6 +235,10 @@ class CtxTrainingArguments: default=None, metadata={"help": "Wandb notes for the experiment."}, ) + add_repeat_prompt: bool = field( + default=True, + metadata={"help": "Whether to add repeat prompt to the dataset."}, + ) @dataclass diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index 4dc57a4..a2e235c 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -19,17 +19,34 @@ def validate_columns(tokenized_ds): ), f"Columns mismatch: {set(tokenized_ds.column_names)} != {ref_cols}" +def filter_long_samples(samples): + return [len(ctx) < 10000 for ctx in samples["context"]] + + +def add_repeat_prompt_fn(samples): + ctxs = samples["context"] * 2 + responses = samples["answer"] * 2 + prompts = samples["prompt"] + prompts += ["Repeat the text above."] * len(prompts) + + return {"context": ctxs, "response": responses, "prompt": prompts} + + def get_tokenized_dataset( ds_name: str, split: str, tokenizer: PreTrainedTokenizerBase, tokenizer_kwargs: dict[str, Any], add_ctx_to_chat: bool, + add_repeat_prompt: bool, ) -> dict[str, Any]: need_ctx_ids = not add_ctx_to_chat ds = load_dataset(ds_name, split=split) ds = ds.map(get_preprocessing_fn(ds_name)) + ds = ds.filter(filter_long_samples, batched=True) + if add_repeat_prompt and "context_numbers" not in ds_name: + ds = ds.map(add_repeat_prompt_fn, batched=True, remove_columns=ds.column_names) # for sft + chat_model, we need to convert the dataset to chat format # add "messages" field ds = ds.map( diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 05e8ea3..a52d16b 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -266,6 +266,7 @@ def main(output_dir): tokenizer=tokenizer, tokenizer_kwargs=tokenizer_kwargs, add_ctx_to_chat=add_ctx_to_chat, + add_repeat_prompt=ctx_args.add_repeat_prompt, # get_ctx_features_fn=get_ctx_features_fn, ) tokenized_ds = {}