diff --git a/configs/small_exp/qa_short_ctx_self_gen_lv1_closed_qa_1_small.yaml b/configs/small_exp/qa_short_ctx_self_gen_lv1_closed_qa_1_small.yaml index ba1630b..8f9f671 100644 --- a/configs/small_exp/qa_short_ctx_self_gen_lv1_closed_qa_1_small.yaml +++ b/configs/small_exp/qa_short_ctx_self_gen_lv1_closed_qa_1_small.yaml @@ -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 diff --git a/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_0.5_small.yaml b/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_0.5_small.yaml index 38f4560..6c965fb 100644 --- a/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_0.5_small.yaml +++ b/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_0.5_small.yaml @@ -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 diff --git a/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_1_small.yaml b/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_1_small.yaml index df050e9..0c5e25f 100644 --- a/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_1_small.yaml +++ b/configs/small_exp/qa_short_ctx_self_gen_lv3_closed_qa_1_small.yaml @@ -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 diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py index b0bbd9e..5c3d85e 100644 --- a/src/ctx_to_lora/data/processing.py +++ b/src/ctx_to_lora/data/processing.py @@ -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)