diff --git a/configs/pwc.yaml b/configs/pwc.yaml index ef7b2af..4f5c0ed 100644 --- a/configs/pwc.yaml +++ b/configs/pwc.yaml @@ -38,5 +38,8 @@ target_modules: train_ds_names: - sggetao/PwC +val_ds_names: +- sggetao/PwC + test_ds_names: - sggetao/PwC diff --git a/configs/pwc_and_ctx_numbers_256.yaml b/configs/pwc_and_ctx_numbers_256.yaml index 4b1609b..f38788b 100644 --- a/configs/pwc_and_ctx_numbers_256.yaml +++ b/configs/pwc_and_ctx_numbers_256.yaml @@ -2,7 +2,7 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] -add_repeat_prompt: false +add_repeat_prompt: true # eval_on_start: True # eval_strategy: "steps" # eval_steps: 500 @@ -58,6 +58,7 @@ train_ds_names: val_ds_names: +- sggetao/PwC - data/raw_datasets/context_numbers_16 - data/raw_datasets/context_numbers_32 - data/raw_datasets/context_numbers_64 diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index b41f033..c1f338c 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -1,3 +1,4 @@ +import logging from copy import copy from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple, Union @@ -8,6 +9,8 @@ from transformers import PreTrainedTokenizerBase IGNORE_INDEX = -100 +logger = logging.getLogger() + def validate_columns(tokenized_ds): cols = ["input_ids", "attention_mask", "labels"] @@ -42,7 +45,13 @@ def get_tokenized_dataset( ) -> dict[str, Any]: need_ctx_ids = not add_ctx_to_chat - ds = load_dataset(ds_name, split=split) + try: + ds = load_dataset(ds_name, split=split) + except ValueError as e: + logger.info( + f"Failed to load dataset {ds_name} with split {split}. Error: {e}\nSkipping..." + ) + return None ds = ds.map(get_preprocessing_fn(ds_name)) ds = ds.filter(filter_long_samples, batched=True) if add_repeat_prompt and "context_numbers" not in ds_name: diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 99eb191..ff0e6fb 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -255,19 +255,13 @@ def main(output_dir): logger.info("Loading dataset...") add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel) - need_ctx_features = isinstance(model, ModulatedPretrainedModel) - prompt_formatting_fn = get_sft_prompt_formatting_fn( - TRAINING_TASK.COMPLETION, tokenizer - ) tokenizer_kwargs = {"max_length": ctx_args.max_base_len} - # get_ctx_features_fn = model.get_ctx_features if need_ctx_features else None _get_tokenized_dataset = partial( get_tokenized_dataset, tokenizer=tokenizer, tokenizer_kwargs=tokenizer_kwargs, add_ctx_to_chat=add_ctx_to_chat, add_repeat_prompt=ctx_args.add_repeat_prompt, - # get_ctx_features_fn=get_ctx_features_fn, ) tokenized_ds = {} for split, ds_names in zip( @@ -276,32 +270,27 @@ def main(output_dir): ): if not ds_names: continue - if split == "train": - tokenized_ds[split] = concatenate_datasets( - [_get_tokenized_dataset(ds_name, split) for ds_name in ds_names] - ) - else: - tokenized_ds[split] = { - os.path.basename(ds_name): _get_tokenized_dataset(ds_name, split) - for ds_name in ds_names - } + tokenized_ds[split] = { + os.path.basename(ds_name): _get_tokenized_dataset(ds_name, split) + for ds_name in ds_names + } train_ds = tokenized_ds["train"] val_ds = dict() if "validation" in tokenized_ds: + n_val_samples = data_args.max_val_samples_per_ds for ds_name, ds in tokenized_ds["validation"].items(): + if ds is None: + # take some samples from train_ds + ds = train_ds[ds_name].take(n_val_samples) + train_ds[ds_name] = train_ds[ds_name].skip(n_val_samples) + val_ds[ds_name] = ds val_ds_size = len(val_ds[ds_name]) - val_indices = np.random.permutation(val_ds_size)[ - : data_args.max_val_samples_per_ds - ] + val_indices = np.random.permutation(val_ds_size)[:n_val_samples] val_ds[ds_name] = val_ds[ds_name].select(val_indices) - else: - # take some samples from train_ds - n_val_samples = data_args.max_val_samples_per_ds - val_ds["train_unseen"] = train_ds.take(n_val_samples) - train_ds = train_ds.skip(n_val_samples) + train_ds = concatenate_datasets(list(train_ds.values())) val_train_indices = np.random.permutation(len(train_ds))[:500] val_ds["train"] = train_ds.select(val_train_indices) @@ -318,10 +307,10 @@ def main(output_dir): logger.info(f"train_ds: {train_ds}") logger.info(f"val_ds: {val_ds}") logger.info(f"test_ds: {test_ds}") + # TODO: change to a faster collator? e.g., # https://huggingface.co/blog/packing-with-FA2 # data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8) - def train_collator(inp_list, tokenizer): # input is a list of tokenized sequences padding_kwargs = dict(