add validation set when not exist

This commit is contained in:
51616 2025-01-09 10:39:05 +00:00
parent b55b9a5643
commit eb522f678d
4 changed files with 28 additions and 26 deletions

View file

@ -38,5 +38,8 @@ target_modules:
train_ds_names:
- sggetao/PwC
val_ds_names:
- sggetao/PwC
test_ds_names:
- sggetao/PwC

View file

@ -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

View file

@ -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:

View file

@ -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(