diff --git a/src/ctx_to_lora/data/definitions.py b/src/ctx_to_lora/data/definitions.py index d1e09dc..ac35b62 100644 --- a/src/ctx_to_lora/data/definitions.py +++ b/src/ctx_to_lora/data/definitions.py @@ -150,6 +150,27 @@ DS_KWARGS = { split="train", ), ), + "fw_qa_v2_2k_len_level_3_tiny": dict( + train=dict( + path="parquet", + data_files="data/raw_datasets/fw_qa_v2/min_0_to_2000/000_00000_level_3.parquet", + split="train", + ), + ), + "fw_qa_v2_2k_len_level_3": dict( + train=dict( + path="parquet", + data_files="data/raw_datasets/fw_qa_v2/min_0_to_2000/*level_3.parquet", + split="train", + ), + # validation=dict( + # path="parquet", + # data_files=glob( + # "data/raw_datasets/fw_qa_v2/min_0_to_2000/*level_3_val.parquet" + # ), + # split="train", + # ), + ), # "fw_qa_3_mini": dict( # train=dict( # path="parquet", diff --git a/src/ctx_to_lora/data/packing.py b/src/ctx_to_lora/data/packing.py index 97e81dd..8080b28 100644 --- a/src/ctx_to_lora/data/packing.py +++ b/src/ctx_to_lora/data/packing.py @@ -250,7 +250,7 @@ if __name__ == "__main__": tokenizer_kwargs = {"max_length": base_model_max_len} # not used ctx_tokenizer_kwargs = {"max_length": base_model_max_len} # not used for now ds = get_tokenized_dataset( - ds_name="squad", + ds_name="fw_qa_v2_2k_len_level_0_tiny", split="train", base_model_max_len=model.base_model.config.max_position_embeddings, tokenizer=tokenizer, diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py index c723382..d53fe7e 100644 --- a/src/ctx_to_lora/data/processing.py +++ b/src/ctx_to_lora/data/processing.py @@ -1,6 +1,7 @@ import hashlib import json import logging +import os import random from collections.abc import Callable from glob import glob @@ -70,27 +71,27 @@ def get_preprocessing_fn( if "fw_qa_v2" in ds_name and "level" in ds_name: - def f_train(sample): + def f(sample): # get questions/answers from all levels in the ds q_cols = [col for col in sample.keys() if col.startswith("prompts_level")] - a_cols = [col for col in sample.keys() if col.startswith("answers_level")] + r_cols = [col for col in sample.keys() if col.startswith("responses_level")] questions = concat_list([sample[col] for col in q_cols]) - answers = concat_list([sample[col] for col in a_cols]) + responses = concat_list([sample[col] for col in r_cols]) + min_len = min(len(questions), len(responses)) + + if min_len == 0: + return { + "context": None, + "prompts": None, + "responses": None, + } + return { "context": sample["context"], - "prompts": questions, - "responses": answers, + "prompts": questions[:min_len], + "responses": responses[:min_len], } - def f_eval(sample): - return { - "context": sample["context"], - "prompt": sample["prompt"], - "response": sample["response"], - } - - f = f_eval if is_eval else f_train - elif ds_name.startswith("longbench"): def f(sample): @@ -262,7 +263,10 @@ def get_preprocessing_fn( def eval_intx_decorator(f): def g(sample): sample = f(sample) - sample["prompt"] = prompt_template.format(input=sample["prompt"]) + prompt_field = "prompt" if "prompt" in sample else "prompts" + sample[prompt_field] = prompt_template.format( + input=sample[prompt_field] + ) return sample return g @@ -693,6 +697,7 @@ def construct_and_tokenize_ctx_qa( remove_columns=[col for col in ds.column_names if col != "context"], ) # add input_ids, attention_mask, labels to "data" + os.environ["TOKENIZERS_PARALLELISM"] = "true" tokenized_ds = ds.map( get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer), batched=True, @@ -746,6 +751,7 @@ def construct_and_tokenize_ctx_qa( if need_ctx_ids: # tokenize the ctx_text to get ctx_ids and ctx_attn_mask + os.environ["TOKENIZERS_PARALLELISM"] = "true" tokenized_ds = tokenized_ds.map( tokenize_ctx_text, fn_kwargs={"tokenizer": ctx_tokenizer}, @@ -767,13 +773,16 @@ def construct_and_tokenize_ctx_qa( tokenized_ds = tokenized_ds.map( unpack_data, num_proc=num_proc, - remove_columns=["data"], + remove_columns=["data", "context"], ) else: + cols_to_remove = ["data", "context"] + if "ctx_ids" in tokenized_ds.column_names: + cols_to_remove += ["ctx_ids", "ctx_attn_mask"] tokenized_ds = tokenized_ds.map( unpack_data_eval, num_proc=num_proc, - remove_columns=["data"], + remove_columns=cols_to_remove, batched=True, batch_size=100_000, ) @@ -808,8 +817,6 @@ def construct_and_tokenize_ctx_qa( tokenized_ds = tokenized_ds.remove_columns(["context", "text"]) if not is_paraphrased: tokenized_ds = tokenized_ds.remove_columns(["qas"]) - else: - tokenized_ds = tokenized_ds.remove_columns(["context"]) if set_format: tokenized_ds.set_format(type=set_format) @@ -955,12 +962,18 @@ def unpack_data_eval(samples): # n_queries always == 1 for eval data = samples["data"] out = dict(input_ids=[], attention_mask=[], labels=[]) - for d in data: + if "ctx_ids" in samples: + out["ctx_ids"] = [] + out["ctx_attn_mask"] = [] + for i, d in enumerate(data): for tokens in zip( d["input_ids"], d["attention_mask"], d["labels"], ): + if "ctx_ids" in samples: + out["ctx_ids"].append(samples["ctx_ids"][i]) + out["ctx_attn_mask"].append(samples["ctx_attn_mask"][i]) out["input_ids"].append(tokens[0]) out["attention_mask"].append(tokens[1]) out["labels"].append(tokens[2])