From 4d50ad84fad3e42150af1eb7a6ff8b34444b2cd7 Mon Sep 17 00:00:00 2001 From: 51616 Date: Fri, 10 Jan 2025 17:37:13 +0000 Subject: [PATCH] hotpot_qa fix!!! (strip) + ds_kwargs + easier ds names --- configs/hotpot_qa.yaml | 6 +- configs/pwc.yaml | 6 +- configs/pwc_and_ctx_numbers_256.yaml | 6 +- hyperlora/data_utils.py | 84 ++++++++++++++++++++-------- hyperlora/intx_sft.py | 3 +- 5 files changed, 72 insertions(+), 33 deletions(-) diff --git a/configs/hotpot_qa.yaml b/configs/hotpot_qa.yaml index cffffd4..a37b52e 100644 --- a/configs/hotpot_qa.yaml +++ b/configs/hotpot_qa.yaml @@ -36,10 +36,10 @@ target_modules: # data train_ds_names: -- hotpotqa/hotpot_qa +- hotpot_qa val_ds_names: -- hotpotqa/hotpot_qa +- hotpot_qa test_ds_names: -- hotpotqa/hotpot_qa +- hotpot_qa diff --git a/configs/pwc.yaml b/configs/pwc.yaml index 4f5c0ed..b508f26 100644 --- a/configs/pwc.yaml +++ b/configs/pwc.yaml @@ -36,10 +36,10 @@ target_modules: # data train_ds_names: -- sggetao/PwC +- pwc val_ds_names: -- sggetao/PwC +- pwc test_ds_names: -- sggetao/PwC +- pwc diff --git a/configs/pwc_and_ctx_numbers_256.yaml b/configs/pwc_and_ctx_numbers_256.yaml index ea4896b..f1d8087 100644 --- a/configs/pwc_and_ctx_numbers_256.yaml +++ b/configs/pwc_and_ctx_numbers_256.yaml @@ -36,7 +36,7 @@ target_modules: # data train_ds_names: -- sggetao/PwC +- pwc - data/raw_datasets/context_numbers_4 - data/raw_datasets/context_numbers_8 - data/raw_datasets/context_numbers_16 @@ -57,7 +57,7 @@ train_ds_names: val_ds_names: -- sggetao/PwC +- pwc - data/raw_datasets/context_numbers_16 - data/raw_datasets/context_numbers_32 - data/raw_datasets/context_numbers_64 @@ -70,4 +70,4 @@ test_ds_names: - data/raw_datasets/context_numbers_64 - data/raw_datasets/context_numbers_128 - data/raw_datasets/context_numbers_256 -- sggetao/PwC +- pwc diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index 68f30c8..42dd5e8 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -1,10 +1,7 @@ import logging -import random import numpy as np -from copy import copy -from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple, Union +from typing import Any, Callable, Iterator, Optional -import pandas as pd from datasets import load_dataset from training_utils import TRAINING_TASK from transformers import PreTrainedTokenizerBase @@ -14,15 +11,47 @@ IGNORE_INDEX = -100 logger = logging.getLogger() DS_KWARGS = { - "hotpotqa/hotpot_qa": dict( - train=dict(name="fullwiki", split="train[1000:2000]"), - validation=dict(name="fullwiki", split="train[:1000]"), - test=dict(name="fullwiki", split="validation"), + "hotpot_qa": dict( + train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[1000:]"), + validation=dict( + path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]" + ), + test=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"), ), - "sggetao/PwC": dict( - train=dict(split="train[1000:]"), - validation=dict(split="train[:1000]"), - test=dict(split="test"), + "hotpot_qa_tiny": dict( + train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[1000:2000]"), + validation=dict( + path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]" + ), + test=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"), + ), + "pwc": dict( + train=dict(path="sggetao/PwC", split="train[1000:]"), + validation=dict(path="sggetao/PwC", split="train[:1000]"), + test=dict(path="sggetao/PwC", split="test"), + ), + "pwc_tiny": dict( + train=dict(path="sggetao/PwC", split="train[1000:2000]"), + validation=dict(path="sggetao/PwC", split="train[:1000]"), + test=dict(path="sggetao/PwC", split="test"), + ), + "fineweb_tiny": dict( + train=dict( + path="parquet", + data_files="data/raw_datasets/fineweb_sharded/00000.parquet", + split="train[2000:]", + streaming=True, + ), + validation=dict( + path="parquet", + data_files="data/raw_datasets/fineweb_sharded/00000.parquet", + split="train[1000:2000]", + ), + test=dict( + path="parquet", + data_files="data/raw_datasets/fineweb_sharded/00000.parquet", + split="train[:1000]", + ), ), } @@ -156,7 +185,7 @@ def get_tokenized_dataset( logger.debug(f"Loading dataset {ds_name} with split {split}...") need_ctx_ids = not add_ctx_to_chat try: - ds = load_dataset(ds_name, **get_ds_kwargs(ds_name, split)) + ds = load_dataset(**get_ds_kwargs(ds_name, split)) except ValueError as e: logger.info( f"Failed to load dataset {ds_name} with split {split}. Error: {e}\nSkipping..." @@ -202,9 +231,11 @@ def get_tokenized_dataset( tokenized_ds = tokenized_ds.map( convert_ctx_prompt_response_to_messages, fn_kwargs={"add_ctx_to_chat": True}, + remove_columns=["messages"], ) tokenized_ds = tokenized_ds.map( - get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer) + get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer), + remove_columns=["chat"], ) tokenized_ds = tokenized_ds.map( tokenize_chat_messages, @@ -337,16 +368,16 @@ def get_preprocessing_fn(ds_name: str) -> Callable[[dict[str, Any]], dict[str, A """ f = lambda x: x - if ds_name == "sggetao/PwC": + if ds_name.startswith("pwc"): def f(sample): return { - "context": sample["context"], + "context": sample["input"], "prompt": sample["prompt"], "response": sample["answer"], } - elif ds_name == "hotpotqa/hotpot_qa": + elif ds_name.startswith("hotpot_qa"): def f(sample): txt = "" @@ -354,7 +385,7 @@ def get_preprocessing_fn(ds_name: str) -> Callable[[dict[str, Any]], dict[str, A txt += " " + "".join(p) return { - "context": txt, + "context": txt.strip(), "prompt": sample["question"], "response": sample["answer"], } @@ -453,9 +484,14 @@ def tokenize_chat_messages( text, return_offsets_mapping=mask_assistant_inputs, add_special_tokens=False, - truncation=True, + truncation=False, **(tokenizer_kwargs or {}), ) + if len(conversation_ids["input_ids"]) >= tokenizer_kwargs["max_length"]: + raise ValueError( + f"Conversation length {len(conversation_ids['input_ids'])} exceeds max length {tokenizer_kwargs['max_length']}" + ) + if mask_assistant_inputs: assistant_ranges = get_assistant_start_end_indices(messages, text) labels = get_masked_labels(conversation_ids, assistant_ranges) @@ -487,9 +523,9 @@ if __name__ == "__main__": model_name = "meta-llama/Llama-3.2-1B-Instruct" messages = [ {"role": "user", "content": "Hello!"}, - {"role": "assistant", "content": "Hey, how are you?"}, - {"role": "user", "content": "Not too bad"}, - {"role": "assistant", "content": "Cooooooool"}, + {"role": "assistant", "content": "Hello!"}, + # {"role": "user", "content": "Not too bad"}, + # {"role": "assistant", "content": "Cooooooool"}, ] tokenizer = AutoTokenizer.from_pretrained(model_name) chat = tokenizer.apply_chat_template( @@ -498,7 +534,9 @@ if __name__ == "__main__": print(chat) model_inputs = tokenize_chat_messages( - {"chat": chat, "messages": messages}, tokenizer + {"chat": chat, "messages": messages}, + tokenizer, + for_kl_loss=True, ) print(tokenizer(chat, add_special_tokens=False)) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 5df4173..72e9e5f 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -489,5 +489,6 @@ if __name__ == "__main__": os.environ["WANDB_WATCH"] = "" # "all" os.environ["WANDB_CONSOLE"] = "off" os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" - disable_caching() + if os.getenv("DEBUG", False): + disable_caching() main()