self_gen parquet pattern + configs

This commit is contained in:
51616 2025-07-10 09:09:28 +00:00
parent c39e44cc90
commit f60d98aa8b
4 changed files with 30 additions and 22 deletions

View file

@ -39,7 +39,7 @@ target_modules:
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_1.0/fw_qa_v2_2k_len_level_1_small
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_1.0/fw_qa_v2/min_0_to_2000/*level_0.parquet
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_0.0/pwc_compact
# these provide exact tokens needed, no need to use self-gen data
- squad_compact

View file

@ -39,7 +39,7 @@ target_modules:
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_0.5/fw_qa_v2_2k_len_level_3_small
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_0.5/fw_qa_v2/min_0_to_2000/*level_3.parquet
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_0.0/pwc_compact
# these provide exact tokens needed, no need to use self-gen data
- squad_compact

View file

@ -39,7 +39,7 @@ target_modules:
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_1.0/fw_qa_v2_2k_len_level_3_small
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_1.0/fw_qa_v2/min_0_to_2000/*level_3.parquet
- self_gen/google/gemma-2-2b-it_temp_0.0_closed_qa_prob_0.0/pwc_compact
# these provide exact tokens needed, no need to use self-gen data
- squad_compact

View file

@ -321,26 +321,34 @@ def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]:
take = slice.split(":")[1]
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]
if ("[" in split) and split.endswith("]"):
kwargs["split"], slice = split.split("[")
slice = slice.strip("]")
skip = slice.split(":")[0]
if skip:
kwargs["skip"] = int(skip)
take = slice.split(":")[1]
if take:
kwargs["take"] = int(take)
files = glob(
f"{SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/{split}/*.parquet"
)
if not files:
raise FileNotFoundError(
f"No self-gen files found for base model {base_model_name} "
f"in {SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/"
if ds_name.endswith(".parquet"):
# ds_name is a glob pattern
files = glob(ds_name)
if not files:
raise FileNotFoundError(
f"The provided pattern does not match any files: {ds_name}"
)
else:
# e.g., "self_gen/google/gemma-2-2b-it/pwc"
base_model_name = "/".join(ds_name.split("/")[1:3])
base_ds = "/".join(ds_name.split("/")[3:])
if ("[" in split) and split.endswith("]"):
kwargs["split"], slice = split.split("[")
slice = slice.strip("]")
skip = slice.split(":")[0]
if skip:
kwargs["skip"] = int(skip)
take = slice.split(":")[1]
if take:
kwargs["take"] = int(take)
files = glob(
f"{SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/{split}/*.parquet"
)
if not files:
raise FileNotFoundError(
f"No self-gen files found for base model {base_model_name} "
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)