self-gen data naming + num_proc to 2

This commit is contained in:
51616 2025-06-03 15:32:38 +00:00
parent 5aa83b96e2
commit 8b15901fb9
3 changed files with 30 additions and 30 deletions

View file

@ -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

View file

@ -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

View file

@ -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,