diff --git a/src/ctx_to_lora/data/definitions.py b/src/ctx_to_lora/data/definitions.py index 02b6f1f..7f7effa 100644 --- a/src/ctx_to_lora/data/definitions.py +++ b/src/ctx_to_lora/data/definitions.py @@ -477,6 +477,13 @@ DS_KWARGS = { split="train", ), ), + "tatqa": dict( + train=dict( + path="nvidia/ChatQA-Training-Data", + name="tatqa", + split="train", + ) + ), "booksum": dict( train=dict( path="kmfoda/booksum", diff --git a/src/ctx_to_lora/data/preprocessing_fn.py b/src/ctx_to_lora/data/preprocessing_fn.py index 66c2564..e7839b6 100644 --- a/src/ctx_to_lora/data/preprocessing_fn.py +++ b/src/ctx_to_lora/data/preprocessing_fn.py @@ -173,17 +173,16 @@ def get_preprocessing_fn( q = closed_qa_prompting(q) if not is_eval else q return {"context": ctx, "prompt": q, "response": response} - elif ds_name in ["narrativeqa", "quoref"]: # , "ropes"]: + elif ds_name in ["narrativeqa", "quoref", "tatqa"]: # , "ropes"]: def f(sample): response = sample["answers"][0] if isinstance(response, list): response = response[0] q = sample["messages"][-1]["content"] - prompt = closed_qa_prompting(q) if not is_eval else q return { "context": sample["document"], - "prompt": prompt, + "prompt": q, "response": response, } @@ -261,6 +260,7 @@ def get_preprocessing_fn( "prompt": sample["instruction"], "response": "```python\n" + sample["code"].strip() + "\n```", } + elif "openhermes" == ds_name: # system prompt is from a set of predefined prompts → query # user prompt is the content → ctx