From 8b15901fb93cd5956be4f762d4244c594fab1c84 Mon Sep 17 00:00:00 2001 From: 51616 Date: Tue, 3 Jun 2025 15:32:38 +0000 Subject: [PATCH] self-gen data naming + num_proc to 2 --- configs/self_gen_3_mini.yaml | 34 ++++++++++++++--------------- configs/self_gen_pwc_hotpot_qa.yaml | 8 +++---- src/ctx_to_lora/data/processing.py | 18 +++++++-------- 3 files changed, 30 insertions(+), 30 deletions(-) diff --git a/configs/self_gen_3_mini.yaml b/configs/self_gen_3_mini.yaml index 234fae1..be673f0 100644 --- a/configs/self_gen_3_mini.yaml +++ b/configs/self_gen_3_mini.yaml @@ -1,6 +1,6 @@ output_dir: "" # just a placeholder bf16: true -model_name_or_path: meta-llama/Llama-3.2-1B-Instruct +model_name_or_path: google/gemma-2-2b-it label_names: ["labels"] # eval_on_start: True # eval_strategy: "steps" @@ -37,24 +37,24 @@ target_modules: - down_proj # data train_ds_names: - - self_gen/fw_qa_3_mini # 100k - - self_gen/ctx_qa # 300k - - self_gen/pwc # 240k - - self_gen/hotpot_qa # 90k - - self_gen/squad # 90k - - self_gen/drop # 77k - - self_gen/narrativeqa # 40k - - self_gen/quoref # 11k - - self_gen/ropes # 11k - - self_gen/synthetic_convqa # 40k + - self_gen/google/gemma-2-2b-it/fw_qa_3_mini # 100k + - self_gen/google/gemma-2-2b-it/ctx_qa # 300k + - self_gen/google/gemma-2-2b-it/pwc # 240k + - self_gen/google/gemma-2-2b-it/hotpot_qa # 90k + - self_gen/google/gemma-2-2b-it/squad # 90k + - self_gen/google/gemma-2-2b-it/drop # 77k + - self_gen/google/gemma-2-2b-it/narrativeqa # 40k + - self_gen/google/gemma-2-2b-it/quoref # 11k + - self_gen/google/gemma-2-2b-it/ropes # 11k + - self_gen/google/gemma-2-2b-it/synthetic_convqa # 40k val_ds_names: - - self_gen/fw_qa_3 - - self_gen/fw_qa_xl - - self_gen/ctx_qa - - self_gen/pwc - - self_gen/hotpot_qa - - self_gen/squad + - self_gen/google/gemma-2-2b-it/fw_qa_3 + - self_gen/google/gemma-2-2b-it/fw_qa_xl + - self_gen/google/gemma-2-2b-it/ctx_qa + - self_gen/google/gemma-2-2b-it/pwc + - self_gen/google/gemma-2-2b-it/hotpot_qa + - self_gen/google/gemma-2-2b-it/squad - fw_qa_3 - fw_qa_xl - ctx_qa diff --git a/configs/self_gen_pwc_hotpot_qa.yaml b/configs/self_gen_pwc_hotpot_qa.yaml index 97bd271..6da0f94 100644 --- a/configs/self_gen_pwc_hotpot_qa.yaml +++ b/configs/self_gen_pwc_hotpot_qa.yaml @@ -39,9 +39,9 @@ target_modules: # data train_ds_names: - - self_gen/pwc - - self_gen/hotpot_qa + - self_gen/google/gemma-2-2b-it/pwc + - self_gen/google/gemma-2-2b-it/hotpot_qa val_ds_names: - - self_gen/pwc - - self_gen/hotpot_qa + - self_gen/google/gemma-2-2b-it/pwc + - self_gen/google/gemma-2-2b-it/hotpot_qa diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py index 44ce9eb..add15bc 100644 --- a/src/ctx_to_lora/data/processing.py +++ b/src/ctx_to_lora/data/processing.py @@ -245,14 +245,8 @@ def get_preprocessing_fn( def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]: - if (ds_name not in DS_KWARGS) or (split not in DS_KWARGS[ds_name]): - kwargs = dict(path=ds_name, split=split) - logger.warning( - f"No dataset kwargs found for '{ds_name}' with split '{split}'.\n" - f"Using default kwargs: {kwargs}" - ) - elif ds_name.startswith("self_gen/"): - # e.g., "self_gen/google/gemma-2-2b-it/pwc/val" + if ds_name.startswith("self_gen/"): + # e.g., "self_gen/google/gemma-2-2b-it/pwc" base_model_name = "/".join(ds_name.split("/")[1:-1]) base_ds = ds_name.split("/")[-1] files = glob( @@ -264,6 +258,12 @@ def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]: f"in {SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/" ) kwargs = dict(path="parquet", data_files=files, split="train") + elif (ds_name not in DS_KWARGS) or (split not in DS_KWARGS[ds_name]): + kwargs = dict(path=ds_name, split=split) + logger.warning( + f"No dataset kwargs found for '{ds_name}' with split '{split}'.\n" + f"Using default kwargs: {kwargs}" + ) else: kwargs = DS_KWARGS[ds_name][split] @@ -522,7 +522,7 @@ def get_tokenized_dataset( repeat_prob=repeat_prob, streaming=streaming, ) - num_proc = None if streaming and split == "train" else 8 + num_proc = None if streaming and split == "train" else 2 ds = load_and_process_dataset( **load_and_process_kwargs,