From 56f393dfd509564052ba81126981fbe2a38840a9 Mon Sep 17 00:00:00 2001 From: 51616 Date: Sun, 22 Dec 2024 17:42:29 +0000 Subject: [PATCH] map fns return only new keys --- hyperlora/data_utils.py | 11 +++-------- 1 file changed, 3 insertions(+), 8 deletions(-) diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index 55ca005..68e3442 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -55,15 +55,10 @@ def get_sft_prompt_formatting_fn( # return output_texts def f_intx(example): - out = dict() - out["context"] = example["context"] - chat_text = tokenizer.apply_chat_template( example["messages"], tokenize=False, add_generation_prompt=False ) - out["messages"] = example["messages"] - out["chat"] = chat_text - return out + return dict(chat=chat_text) # return f if not apply_chat_template_fn is not None else f_intx return f_intx @@ -97,13 +92,13 @@ def convert_ctx_prompt_response_to_messages( if add_ctx_to_chat: user_msg = example["context"] + "\n" + user_msg - example["messages"] = [ + messages = [ {"role": "system", "content": system_msg}, {"role": "user", "content": user_msg}, {"role": "assistant", "content": example["response"]}, ] - return example + return dict(messages=messages) def get_preprocessing_fn(ds_name: str) -> Callable[[dict[str, Any]], dict[str, Any]]: