mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
* sft data from smoltalk + multi-turn ctx data * add ctx-q-gsm (stateful-gsm from sleep-time compute)
956 lines
31 KiB
Python
956 lines
31 KiB
Python
import json
|
|
import logging
|
|
import os
|
|
import random
|
|
from collections.abc import Callable
|
|
from glob import glob
|
|
from hashlib import sha256
|
|
from math import ceil, isclose
|
|
from os import path
|
|
from typing import Any
|
|
|
|
import datasets
|
|
import numpy as np
|
|
import torch
|
|
from datasets import Dataset, interleave_datasets, is_caching_enabled, load_dataset
|
|
from transformers import PreTrainedTokenizerBase
|
|
|
|
from ctx_to_lora.data.definitions import (
|
|
CTX_AFFIXES,
|
|
DS_KWARGS,
|
|
IGNORE_INDEX,
|
|
RAW_DATA_DIR,
|
|
REPEAT_PROMPTS,
|
|
SELF_GEN_DATA_DIR,
|
|
TRANSFORMED_DATA_DIR,
|
|
)
|
|
from ctx_to_lora.data.packing import pack_batch
|
|
from ctx_to_lora.data.preprocessing_fn import get_preprocessing_fn
|
|
from ctx_to_lora.utils import check_is_iterable, concat_list
|
|
|
|
logger = logging.getLogger()
|
|
|
|
COLS_TO_KEEP_PREPROCESSING = [
|
|
"context",
|
|
"prompts",
|
|
"responses",
|
|
"qas",
|
|
"variation",
|
|
"logprobs_vals",
|
|
"logprobs_indices",
|
|
"input_ids",
|
|
"ctx_ids",
|
|
"response_start_end",
|
|
]
|
|
|
|
COLS_TO_KEEP_TOKENIZED = [
|
|
"input_ids",
|
|
"labels",
|
|
"context",
|
|
"ctx_ids",
|
|
"logprobs_vals",
|
|
"logprobs_indices",
|
|
]
|
|
|
|
|
|
def get_ds_prob(train_ds_len: list[int], total_len: int):
|
|
# if a dataset is smaller than 1%, make it 1%
|
|
probs = [0 for _ in train_ds_len]
|
|
for i, ds_len in enumerate(train_ds_len):
|
|
if ds_len / total_len <= 0.01:
|
|
probs[i] = 0.01
|
|
res_probs = 1 - sum(probs)
|
|
res_total_len = sum([l for l in train_ds_len if (l / total_len) > 0.01])
|
|
for i, ds_len in enumerate(train_ds_len):
|
|
if (ds_len / total_len) > 0.01:
|
|
probs[i] = ds_len / res_total_len * res_probs
|
|
logger.debug(f"Dataset probabilities: {probs}")
|
|
assert isclose(sum(probs), 1.0), (
|
|
f"Probs sum to {sum(probs)} ({probs}), expected 1.0"
|
|
)
|
|
return probs
|
|
|
|
|
|
def load_answers(ds_name, split):
|
|
if ds_name.startswith("longbench"):
|
|
|
|
def extract_ans(sample):
|
|
return {"answers": sample["answers"]}
|
|
|
|
elif ds_name == "squad":
|
|
|
|
def extract_ans(sample):
|
|
return {"answers": sample["answers"]["text"]}
|
|
|
|
elif ds_name == "drop":
|
|
|
|
def extract_ans(sample):
|
|
return {"answers": sample["answers_spans"]["spans"]}
|
|
|
|
ds_kwargs = get_ds_kwargs(ds_name, split)
|
|
ds = load_dataset(**ds_kwargs, trust_remote_code=True)
|
|
ds = ds.map(extract_ans, num_proc=8, remove_columns=ds.column_names)
|
|
return ds
|
|
|
|
|
|
def get_repeat_prompt():
|
|
return random.choice(REPEAT_PROMPTS)
|
|
|
|
|
|
def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]:
|
|
# custom logic for slicing iterable datasets
|
|
take, skip = None, None
|
|
if ("[" in split) and split.endswith("]"):
|
|
split, slice = split.split("[")
|
|
slice = slice.strip("]")
|
|
skip = slice.split(":")[0]
|
|
take = slice.split(":")[1]
|
|
|
|
if ds_name.startswith("self_gen/"):
|
|
if ds_name.endswith(".parquet"):
|
|
# ds_name is a glob pattern
|
|
files = glob(f"{RAW_DATA_DIR}/{ds_name}")
|
|
if not files:
|
|
raise FileNotFoundError(
|
|
f"The provided pattern does not match any files: {RAW_DATA_DIR}/{ds_name}"
|
|
)
|
|
else:
|
|
# e.g., "self_gen/google/gemma-2-2b-it/pwc"
|
|
base_model_name = "/".join(ds_name.split("/")[1:3])
|
|
base_ds = "/".join(ds_name.split("/")[3:])
|
|
if ("[" in split) and split.endswith("]"):
|
|
kwargs["split"], slice = split.split("[")
|
|
slice = slice.strip("]")
|
|
skip = slice.split(":")[0]
|
|
if skip:
|
|
kwargs["skip"] = int(skip)
|
|
take = slice.split(":")[1]
|
|
if take:
|
|
kwargs["take"] = int(take)
|
|
files = glob(
|
|
f"{SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/{split}/*.parquet"
|
|
)
|
|
if not files:
|
|
raise FileNotFoundError(
|
|
f"No self-gen files found for base model {base_model_name} "
|
|
f"in {SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/"
|
|
)
|
|
kwargs = dict(path="parquet", data_files=files, split="train")
|
|
elif (ds_name not in DS_KWARGS) or (split not in DS_KWARGS[ds_name]):
|
|
kwargs = dict(path=ds_name, split=split)
|
|
logger.warning(
|
|
f"No dataset kwargs found for '{ds_name}' with split '{split}'.\n"
|
|
f"Using default kwargs: {kwargs}"
|
|
)
|
|
else:
|
|
kwargs = DS_KWARGS[ds_name][split]
|
|
|
|
if skip:
|
|
kwargs["skip"] = int(skip)
|
|
if take:
|
|
kwargs["take"] = int(take)
|
|
|
|
return kwargs
|
|
|
|
|
|
def len_filter(sample, max_length: int, keys: list[str]):
|
|
m = [len(sample[k]) <= max_length for k in keys]
|
|
return sum(m) == len(keys)
|
|
|
|
|
|
def filter_none(sample):
|
|
for v in sample.values():
|
|
if v is None:
|
|
return False
|
|
return True
|
|
|
|
|
|
def load_and_process_dataset(
|
|
ds_name: str,
|
|
split: str,
|
|
num_proc: int,
|
|
):
|
|
logger.info(f"Loading dataset {ds_name} with split {split}...")
|
|
try:
|
|
ds_kwargs = get_ds_kwargs(ds_name, split)
|
|
skip = ds_kwargs.pop("skip", None)
|
|
take = ds_kwargs.pop("take", None)
|
|
ds = load_dataset(**ds_kwargs, trust_remote_code=True)
|
|
if skip is not None:
|
|
ds = ds.skip(skip)
|
|
if take is not None:
|
|
ds = ds.take(take)
|
|
except ValueError as e:
|
|
raise ValueError(
|
|
f"Failed to load dataset {ds_name} with split {split}. Error: {e}"
|
|
)
|
|
cols_to_remove = [
|
|
col for col in ds.column_names if col not in COLS_TO_KEEP_PREPROCESSING
|
|
]
|
|
is_eval = split != "train"
|
|
ds = ds.map(
|
|
get_preprocessing_fn(ds_name, is_eval),
|
|
remove_columns=cols_to_remove,
|
|
num_proc=16,
|
|
)
|
|
ds = ds.filter(
|
|
filter_none,
|
|
batched=False,
|
|
num_proc=16,
|
|
)
|
|
return ds
|
|
|
|
|
|
def get_tokenized_dataset(
|
|
ds_name: str,
|
|
split: str,
|
|
max_qas_len: int,
|
|
max_qas_per_sample: int,
|
|
base_model_max_len: int,
|
|
tokenizer: PreTrainedTokenizerBase,
|
|
ctx_model_max_len: int,
|
|
ctx_tokenizer: PreTrainedTokenizerBase,
|
|
max_ctx_chunk_len: int,
|
|
min_ctx_chunk_len: int,
|
|
random_chunking: bool,
|
|
max_ctx_chunk_num: int,
|
|
add_ctx_to_chat: bool,
|
|
use_kl_loss: bool,
|
|
max_new_tokens: int = 256,
|
|
set_format: str | None = None,
|
|
) -> dict[str, Any]:
|
|
if max_qas_len > 0:
|
|
assert max_qas_len <= base_model_max_len, (
|
|
f"`max_qas_len` should be <= {base_model_max_len=}, got {max_qas_len=}"
|
|
)
|
|
logger.info(f"Loading dataset {ds_name} with split {split}...")
|
|
need_ctx_ids = not add_ctx_to_chat and bool(ctx_model_max_len)
|
|
|
|
load_and_process_kwargs = dict(
|
|
ds_name=ds_name,
|
|
split=split,
|
|
)
|
|
tokenize_kwargs = dict(
|
|
max_qas_len=max_qas_len,
|
|
max_qas_per_sample=max_qas_per_sample,
|
|
base_model_max_len=base_model_max_len,
|
|
ctx_model_max_len=ctx_model_max_len,
|
|
add_ctx_to_chat=add_ctx_to_chat,
|
|
max_ctx_chunk_len=max_ctx_chunk_len,
|
|
min_ctx_chunk_len=min_ctx_chunk_len,
|
|
random_chunking=random_chunking,
|
|
max_ctx_chunk_num=max_ctx_chunk_num,
|
|
need_ctx_ids=need_ctx_ids,
|
|
split=split,
|
|
max_new_tokens=max_new_tokens,
|
|
set_format=set_format,
|
|
)
|
|
|
|
all_kwargs = {**load_and_process_kwargs, **tokenize_kwargs}
|
|
kwargs_str = json.dumps(all_kwargs)
|
|
kwargs_str += tokenizer.name_or_path + ctx_tokenizer.name_or_path
|
|
logger.debug(f"Tokenizing dataset with kwargs: {kwargs_str}")
|
|
kwargs_str += repr(tokenizer) + repr(ctx_tokenizer)
|
|
ds_hash = sha256(kwargs_str.encode()).hexdigest()
|
|
logger.debug(f"Dataset hash: {ds_hash}")
|
|
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
|
|
|
|
if path.exists(ds_path) and ("train" in split) and is_caching_enabled():
|
|
# load the cached ds
|
|
logger.info(f"Loaded tokenized dataset from {ds_path}")
|
|
tokenized_ds = datasets.load_from_disk(ds_path)
|
|
if (not use_kl_loss) and ("logprobs_vals" in tokenized_ds.column_names):
|
|
tokenized_ds = tokenized_ds.remove_columns(
|
|
["logprobs_vals", "logprobs_indices"]
|
|
)
|
|
return tokenized_ds
|
|
|
|
num_proc = 4
|
|
ds = load_and_process_dataset(
|
|
**load_and_process_kwargs,
|
|
num_proc=num_proc,
|
|
)
|
|
if use_kl_loss:
|
|
if "train" in split and "logprobs_vals" not in ds.column_names:
|
|
raise ValueError(
|
|
"`use_kl_loss` is set to True but 'logprobs_vals' column "
|
|
"is not present in the dataset."
|
|
)
|
|
logger.info(f"Constructing and tokenizing {ds_name} with {split} split...")
|
|
|
|
tokenized_ds = construct_and_tokenize_ctx_qa(
|
|
ds=ds,
|
|
tokenizer=tokenizer,
|
|
ctx_tokenizer=ctx_tokenizer,
|
|
num_proc=num_proc,
|
|
**tokenize_kwargs,
|
|
)
|
|
|
|
if ("train" in split) and is_caching_enabled():
|
|
tokenized_ds = tokenized_ds.shuffle()
|
|
tokenized_ds.save_to_disk(ds_path, num_proc=16)
|
|
# force reload from disk for fingerprint consistency
|
|
tokenized_ds = datasets.load_from_disk(ds_path)
|
|
|
|
if (not use_kl_loss) and ("logprobs_vals" in tokenized_ds.column_names):
|
|
tokenized_ds = tokenized_ds.remove_columns(
|
|
["logprobs_vals", "logprobs_indices"]
|
|
)
|
|
return tokenized_ds
|
|
|
|
|
|
def construct_and_tokenize_ctx_qa(
|
|
max_qas_len,
|
|
max_qas_per_sample,
|
|
base_model_max_len,
|
|
tokenizer,
|
|
ctx_model_max_len,
|
|
ctx_tokenizer,
|
|
add_ctx_to_chat,
|
|
need_ctx_ids,
|
|
max_ctx_chunk_len,
|
|
min_ctx_chunk_len,
|
|
random_chunking,
|
|
max_ctx_chunk_num,
|
|
ds,
|
|
split,
|
|
max_new_tokens,
|
|
set_format=None,
|
|
num_proc=None,
|
|
):
|
|
is_train = "train" in split
|
|
# for sft + chat_model, we need to convert the dataset to chat format
|
|
if "input_ids" in ds.column_names and "response_start_end" in ds.column_names:
|
|
# already tokenized dataset (e.g., self-gen qa data)
|
|
tokenized_ds = ds.map(get_labels_from_input_ids, num_proc=16)
|
|
else:
|
|
# construct messages from prompts and responses
|
|
# add "messages_list" field
|
|
ds = ds.map(
|
|
convert_ctx_prompt_response_to_messages,
|
|
fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat},
|
|
num_proc=16,
|
|
)
|
|
# add `input_ids`, `attention_mask`, `labels`
|
|
os.environ["TOKENIZERS_PARALLELISM"] = "true"
|
|
logging.debug("Tokenizing inputs")
|
|
tokenized_ds = ds.map(
|
|
get_sft_prompt_formatting_fn(tokenizer),
|
|
batched=True,
|
|
batch_size=100_000,
|
|
)
|
|
|
|
tokenized_ds = tokenized_ds.remove_columns(
|
|
[col for col in tokenized_ds.column_names if col not in COLS_TO_KEEP_TOKENIZED],
|
|
)
|
|
tokenized_ds = tokenized_ds.filter(
|
|
lambda x: bool(x["input_ids"]), # remove empty "input_ids"
|
|
num_proc=16,
|
|
)
|
|
|
|
if need_ctx_ids:
|
|
if "ctx_ids" not in tokenized_ds.column_names:
|
|
# tokenize the ctx_text to get ctx_ids and ctx_attn_mask
|
|
os.environ["TOKENIZERS_PARALLELISM"] = "true"
|
|
logging.debug("Tokenizing context")
|
|
tokenized_ds = tokenized_ds.map(
|
|
tokenize_ctx_text,
|
|
fn_kwargs={"tokenizer": ctx_tokenizer},
|
|
batched=True,
|
|
batch_size=100_000,
|
|
remove_columns=["context"],
|
|
)
|
|
|
|
# # TODO: this can be removed once we implement ctx chunking for training
|
|
# if is_train:
|
|
# # drop
|
|
# tokenized_ds = tokenized_ds.filter(
|
|
# len_filter,
|
|
# fn_kwargs={"max_length": ctx_model_max_len, "keys": ["ctx_ids"]},
|
|
# num_proc=16,
|
|
# )
|
|
|
|
# ctx chunking
|
|
# Ideally we want to chunk the raw text directly
|
|
# however, since the contexts are tokenized during self-gen
|
|
# we can only chunk the tokenized context which requires some workaround
|
|
# e.g., removing/applying template to each chunk
|
|
# with some big caveats, e.g., losing order info
|
|
split_ctx_kwargs = {
|
|
"max_chunk_len": max_ctx_chunk_len,
|
|
"min_chunk_len": min_ctx_chunk_len
|
|
if min_ctx_chunk_len > 0
|
|
else max_ctx_chunk_len,
|
|
"random_chunking": random_chunking,
|
|
"max_num_split": max_ctx_chunk_num,
|
|
"model_name_or_path": tokenizer.name_or_path,
|
|
"is_train": is_train,
|
|
}
|
|
logging.info(f"Chunking context with {split_ctx_kwargs=}")
|
|
tokenized_ds = tokenized_ds.map(split_too_long_ctx, fn_kwargs=split_ctx_kwargs)
|
|
logging.info(
|
|
f"Avg. num chunks per ctx: {np.mean(list(map(len, tokenized_ds['ctx_ids'])))}"
|
|
)
|
|
|
|
split_qa_kwargs = {
|
|
"max_qas_len": max_qas_len,
|
|
"max_qas_per_sample": max_qas_per_sample,
|
|
}
|
|
logging.info(f"Split too long QAs with {split_qa_kwargs=}")
|
|
tokenized_ds = tokenized_ds.map(
|
|
split_too_long_qas,
|
|
fn_kwargs=split_qa_kwargs,
|
|
batched=True,
|
|
batch_size=12_500,
|
|
num_proc=16,
|
|
)
|
|
if "train" not in split:
|
|
# squeeze since we always have one query per sample in eval
|
|
tokenized_ds = tokenized_ds.map(squeeze_tokens, num_proc=num_proc)
|
|
tokenized_ds = tokenized_ds.map(
|
|
truncate_middle_if_too_long,
|
|
fn_kwargs={
|
|
"max_length": base_model_max_len,
|
|
"columns": ["input_ids", "labels"],
|
|
"max_new_tokens": max_new_tokens,
|
|
},
|
|
)
|
|
tokenized_ds = tokenized_ds.map(
|
|
add_length_info,
|
|
fn_kwargs={"columns": ["input_ids"]},
|
|
)
|
|
if "ctx_ids" in tokenized_ds.column_names:
|
|
# TODO: remove since we already have ctx chunking for eval
|
|
# tokenized_ds = tokenized_ds.map(
|
|
# truncate_middle_if_too_long,
|
|
# fn_kwargs={
|
|
# "max_length": ctx_model_max_len,
|
|
# "columns": ["ctx_ids"],
|
|
# # cxt encoder doesnt need to add new_tokens
|
|
# "max_new_tokens": 0,
|
|
# },
|
|
# )
|
|
tokenized_ds = tokenized_ds.map(
|
|
add_length_info,
|
|
fn_kwargs={"columns": ["ctx_ids"]},
|
|
)
|
|
|
|
if set_format:
|
|
tokenized_ds.set_format(type=set_format)
|
|
|
|
return tokenized_ds
|
|
|
|
|
|
def get_labels_from_input_ids(sample: dict[str, Any]) -> dict[str, Any]:
|
|
"""
|
|
Extract labels from input_ids and response_start.
|
|
|
|
Args:
|
|
sample: A dictionary containing 'input_ids' and 'response_start'
|
|
|
|
Returns:
|
|
A dictionary with 'labels' field added
|
|
"""
|
|
labels = []
|
|
for input_ids_i, (start_i, end_i) in zip(
|
|
sample["input_ids"], sample["response_start_end"]
|
|
):
|
|
len_input_ids = len(input_ids_i)
|
|
# pad labels with -100
|
|
pad_len_left = start_i
|
|
pad_len_right = len_input_ids - end_i
|
|
labels.append(
|
|
[IGNORE_INDEX] * pad_len_left
|
|
+ input_ids_i[start_i:]
|
|
+ [IGNORE_INDEX] * pad_len_right
|
|
)
|
|
|
|
sample["labels"] = labels
|
|
return sample
|
|
|
|
|
|
def get_sft_prompt_formatting_fn(
|
|
tokenizer: PreTrainedTokenizerBase,
|
|
) -> Callable[[dict[str, Any]], dict[str, Any]]:
|
|
"""
|
|
Get a function that formats examples for supervised fine-tuning.
|
|
|
|
Args:
|
|
sft_mode: The training task type
|
|
tokenizer: The tokenizer to use for chat template application
|
|
|
|
Returns:
|
|
A function that takes a training example and returns formatted data
|
|
|
|
Raises:
|
|
NotImplementedError: If sft_mode is not COMPLETION or tokenizer has no chat template
|
|
"""
|
|
|
|
if tokenizer.chat_template is None:
|
|
raise NotImplementedError("Only chat models are supported")
|
|
|
|
@torch.inference_mode()
|
|
def f_intx(samples):
|
|
# flatten all the messages into a list
|
|
# tokenize, the pack back correctly
|
|
messages_list = [x for x in samples["messages_list"]]
|
|
|
|
n_queries = [len(x) for x in messages_list]
|
|
messages = concat_list(messages_list)
|
|
logger.info(f"Tokenizing {len(messages)} messages...")
|
|
tokens = tokenizer.apply_chat_template(
|
|
messages,
|
|
tokenize=True,
|
|
add_special_tokens=False,
|
|
padding=False,
|
|
truncation=False,
|
|
return_attention_mask=False,
|
|
add_generation_prompt=False,
|
|
return_assistant_tokens_mask=True,
|
|
return_dict=True,
|
|
)
|
|
labels = []
|
|
for tok_ids, masks in zip(tokens["input_ids"], tokens["assistant_masks"]):
|
|
o = [id_ if mask else IGNORE_INDEX for id_, mask in zip(tok_ids, masks)]
|
|
labels.append(o)
|
|
|
|
del tokens["assistant_masks"]
|
|
tokens["labels"] = labels
|
|
per_ctx_tokens = {"input_ids": [], "labels": []}
|
|
i = 0
|
|
for n in n_queries:
|
|
per_ctx_tokens["input_ids"].append(tokens["input_ids"][i : i + n])
|
|
per_ctx_tokens["labels"].append(tokens["labels"][i : i + n])
|
|
i += n
|
|
|
|
return per_ctx_tokens
|
|
|
|
return f_intx
|
|
|
|
|
|
def convert_ctx_prompt_response_to_messages(
|
|
example: dict[str, Any],
|
|
add_ctx_to_chat: bool,
|
|
) -> dict[str, Any]:
|
|
"""
|
|
Convert context/prompt/response format to chat messages format.
|
|
|
|
Args:
|
|
example: Dictionary containing 'prompt' and 'response' keys
|
|
add_ctx_to_chat: Whether to prepend context to the user message
|
|
|
|
Returns:
|
|
Dictionary with added 'messages' key containing chat format
|
|
|
|
Raises:
|
|
ValueError: If 'prompt' or 'response' keys are missing
|
|
"""
|
|
prompt_field = "prompts"
|
|
res_field = "responses"
|
|
|
|
if prompt_field not in example or res_field not in example:
|
|
raise ValueError(
|
|
f"'{prompt_field}' and '{res_field}' are required. Got: {example}"
|
|
)
|
|
|
|
system_msg = ""
|
|
if "system_message" in example:
|
|
system_msg = example["system_message"].strip()
|
|
|
|
messages_list = []
|
|
for prompt, response in zip(example[prompt_field], example[res_field]):
|
|
user_msg = prompt.strip()
|
|
if add_ctx_to_chat:
|
|
user_msg = example["context"].strip() + "\n\n" + user_msg
|
|
|
|
messages_list.append(
|
|
[
|
|
{"role": "system", "content": system_msg.strip()},
|
|
{"role": "user", "content": user_msg.strip()},
|
|
{"role": "assistant", "content": response},
|
|
]
|
|
)
|
|
|
|
return {"messages_list": messages_list}
|
|
|
|
|
|
def split_too_long_ctx(
|
|
sample: dict[str, Any],
|
|
model_name_or_path: str,
|
|
max_chunk_len: int,
|
|
min_chunk_len: int,
|
|
max_num_split: int,
|
|
is_train: bool,
|
|
random_chunking: bool,
|
|
) -> dict[str, Any]:
|
|
"""
|
|
Split context into smaller chunks if it exceeds the maximum length.
|
|
|
|
Args:
|
|
samples: Dictionary containing 'ctx_ids' and 'ctx_attn_mask'
|
|
max_chunk_len: Maximum length for each context chunk
|
|
max_num_split: Maximum number of splits allowed
|
|
|
|
# Training
|
|
random_chunking: Wheter to use stochastic chunking with
|
|
`min_chunk_len` to `max_chunk_len` sizes
|
|
min_chunk_len: Minimum length for each context chunk (used when `random_chunking=True`)
|
|
|
|
Returns:
|
|
Dictionary with split context data
|
|
"""
|
|
|
|
chunk_len = max_chunk_len
|
|
if is_train:
|
|
# TODO: for training, we might wanna sort the context by num chunks
|
|
# since merging a batch of chunked loras need padded in the rank axis
|
|
# e.g., ctx1 has 5 chunks (rank-48), ctx2 has 10 chunks (rank-88)
|
|
|
|
# e.g., even if the ctx is not too long, we still split it randomly?
|
|
# say 0-4k is max len for one split,
|
|
# for some ctx shorter than 4k, we leave it as is
|
|
# for some we split to smaller chunks
|
|
if random_chunking:
|
|
chunk_len = random.randint(min_chunk_len, max_chunk_len)
|
|
|
|
ctx_affixes = CTX_AFFIXES[model_name_or_path]
|
|
prefix = ctx_affixes["prefix"]
|
|
suffix = ctx_affixes["suffix"]
|
|
ctx_ids = sample["ctx_ids"]
|
|
|
|
if chunk_len <= 0 and max_num_split <= 0:
|
|
return {"ctx_ids": [ctx_ids]}
|
|
|
|
if len(ctx_ids) <= chunk_len:
|
|
return {"ctx_ids": [ctx_ids]}
|
|
|
|
# uniform chunking
|
|
n_chunks = ceil(len(ctx_ids) / chunk_len)
|
|
avg_len = ceil(len(ctx_ids) / n_chunks)
|
|
|
|
# Split the context into smaller chunks
|
|
chunks = [ctx_ids[i : i + avg_len] for i in range(0, len(ctx_ids), avg_len)]
|
|
|
|
# this would exceed the avg_len a bit
|
|
chunks[0] = chunks[0] + suffix
|
|
for i in range(1, len(chunks) - 1):
|
|
chunks[i] = prefix + chunks[i] + suffix
|
|
chunks[-1] = prefix + chunks[-1]
|
|
|
|
return {"ctx_ids": chunks}
|
|
|
|
|
|
def split_too_long_qas(
|
|
samples: dict[str, any], max_qas_len: int, max_qas_per_sample: int
|
|
):
|
|
# samples keys: "input_ids", "attention_mask", "labels", "ctx_ids", "ctx_attn_mask"
|
|
# split the qas into multiple samples if they are too long
|
|
# e.g., if max_qas_len = 512, and qas is 1024 tokens long,
|
|
# we split it such that each sample has at most 512 tokens
|
|
# and the ctx_ids and ctx_attn_mask are the same for all samples
|
|
if max_qas_len < 0 and max_qas_per_sample < 0:
|
|
return samples
|
|
input_ids = samples["input_ids"]
|
|
labels = samples["labels"]
|
|
ctx_ids = samples["ctx_ids"]
|
|
target_logprobs_vals = samples.get("logprobs_vals", None)
|
|
target_logprobs_indices = samples.get("logprobs_indices", None)
|
|
|
|
# Pre-calculate total lengths to check if any splitting is needed
|
|
total_lengths = [sum(len(x) for x in seq) for seq in input_ids]
|
|
longest_old_qas_len = max(total_lengths) if total_lengths else 0
|
|
|
|
# Early exit if no splitting needed
|
|
if (max_qas_len < 0 or all(length <= max_qas_len for length in total_lengths)) and (
|
|
max_qas_per_sample < 0
|
|
or all(len(seq) <= max_qas_per_sample for seq in input_ids)
|
|
):
|
|
logger.debug(f"Longest old qas len: {longest_old_qas_len}")
|
|
logger.debug(f"Longest new qas len: {longest_old_qas_len}")
|
|
return samples
|
|
|
|
out = {k: list() for k in samples}
|
|
longest_new_qas_len = 0
|
|
n_skip = 0
|
|
has_target_logprobs = (
|
|
target_logprobs_vals is not None and target_logprobs_indices is not None
|
|
)
|
|
|
|
# Helper function to add a batch efficiently
|
|
def add_batch(
|
|
inp_ids_batch,
|
|
labels_batch,
|
|
ctx_id,
|
|
target_vals_batch=None,
|
|
target_indices_batch=None,
|
|
):
|
|
out["input_ids"].append(inp_ids_batch)
|
|
out["labels"].append(labels_batch)
|
|
out["ctx_ids"].append(ctx_id)
|
|
if has_target_logprobs:
|
|
out["logprobs_vals"].append(target_vals_batch)
|
|
out["logprobs_indices"].append(target_indices_batch)
|
|
|
|
for i, tot_inp_len in enumerate(total_lengths):
|
|
if (max_qas_len < 0 or tot_inp_len <= max_qas_len) and (
|
|
max_qas_per_sample < 0 or len(input_ids[i]) <= max_qas_per_sample
|
|
):
|
|
# No need to split - add entire sample
|
|
# logger.debug(f"Sample {i} is within limits, adding as is.")
|
|
for k in samples:
|
|
out[k].append(samples[k][i])
|
|
continue
|
|
|
|
# Need to split this sample
|
|
current_ctx_id = ctx_ids[i]
|
|
new_qas_len = 0
|
|
new_input_ids = []
|
|
new_labels = []
|
|
new_target_vals = [] if has_target_logprobs else None
|
|
new_target_indices = [] if has_target_logprobs else None
|
|
|
|
sequences = zip(input_ids[i], labels[i])
|
|
if has_target_logprobs:
|
|
sequences = zip(
|
|
input_ids[i],
|
|
labels[i],
|
|
target_logprobs_vals[i],
|
|
target_logprobs_indices[i],
|
|
)
|
|
|
|
for seq_data in sequences:
|
|
if has_target_logprobs:
|
|
inp_ids, label, target_vals, target_indices = seq_data
|
|
else:
|
|
inp_ids, label = seq_data
|
|
target_vals, target_indices = None, None
|
|
|
|
inp_len = len(inp_ids)
|
|
if (max_qas_len > 0) and (inp_len > max_qas_len):
|
|
# Skip individual sequences that are too long
|
|
n_skip += 1
|
|
continue
|
|
|
|
# Check if we can add to current batch (both length and sample count limits)
|
|
can_add_to_current = (
|
|
max_qas_len < 0 or new_qas_len + inp_len <= max_qas_len
|
|
) and (max_qas_per_sample < 0 or len(new_input_ids) < max_qas_per_sample)
|
|
|
|
if can_add_to_current:
|
|
# Add to current batch
|
|
new_qas_len += inp_len
|
|
new_input_ids.append(inp_ids)
|
|
new_labels.append(label)
|
|
if has_target_logprobs:
|
|
new_target_vals.append(target_vals)
|
|
new_target_indices.append(target_indices)
|
|
else:
|
|
# Current batch is full, save it and start new batch
|
|
# logger.debug(
|
|
# f"sample {i}: adding batch with {len(new_input_ids)} sequences"
|
|
# )
|
|
if new_input_ids: # Only add non-empty batches
|
|
add_batch(
|
|
new_input_ids,
|
|
new_labels,
|
|
current_ctx_id,
|
|
new_target_vals,
|
|
new_target_indices,
|
|
)
|
|
longest_new_qas_len = max(longest_new_qas_len, new_qas_len)
|
|
|
|
# Start new batch with current sequence
|
|
new_qas_len = inp_len
|
|
new_input_ids = [inp_ids]
|
|
new_labels = [label]
|
|
if has_target_logprobs:
|
|
new_target_vals = [target_vals]
|
|
new_target_indices = [target_indices]
|
|
|
|
# Add final batch if not empty
|
|
if new_input_ids:
|
|
add_batch(
|
|
new_input_ids,
|
|
new_labels,
|
|
current_ctx_id,
|
|
new_target_vals,
|
|
new_target_indices,
|
|
)
|
|
longest_new_qas_len = max(longest_new_qas_len, new_qas_len)
|
|
|
|
logger.debug(f"Longest old qas len: {longest_old_qas_len}")
|
|
logger.debug(f"Longest new qas len: {longest_new_qas_len}")
|
|
if n_skip:
|
|
logger.warning(
|
|
f"Skipped {n_skip} QA pairs because they were too long (> {max_qas_len=} tokens)"
|
|
)
|
|
|
|
return out
|
|
|
|
|
|
def unpack_data_eval(samples):
|
|
# n_queries always == 1 for eval
|
|
data = samples["data"]
|
|
out = dict(input_ids=[], labels=[])
|
|
if "ctx_ids" in samples:
|
|
out["ctx_ids"] = []
|
|
for i, d in enumerate(data):
|
|
for tokens in zip(
|
|
d["input_ids"],
|
|
d["labels"],
|
|
):
|
|
if "ctx_ids" in samples:
|
|
out["ctx_ids"].append(samples["ctx_ids"][i])
|
|
out["input_ids"].append(tokens[0])
|
|
out["labels"].append(tokens[2])
|
|
return out
|
|
|
|
|
|
def squeeze_tokens(sample: dict[str, Any]) -> dict[str, Any]:
|
|
"""
|
|
Squeeze the input_ids and labels to remove any extra dimensions.
|
|
|
|
Args:
|
|
sample: A dictionary containing 'input_ids' and 'labels'
|
|
Returns:
|
|
A dictionary with squeezed 'input_ids' and 'labels'
|
|
"""
|
|
first_id = sample["input_ids"][0]
|
|
if check_is_iterable(first_id):
|
|
sample["input_ids"] = first_id
|
|
first_label = sample["labels"][0]
|
|
if check_is_iterable(first_label):
|
|
sample["labels"] = first_label
|
|
return sample
|
|
|
|
|
|
def add_length_info(sample: dict[str, any], columns: list[str]) -> dict[str, any]:
|
|
out = {}
|
|
for k in columns:
|
|
if check_is_iterable(sample[k][0]):
|
|
# ctx_ids
|
|
out[f"{k}_len"] = sum([len(x) for x in sample[k]])
|
|
else:
|
|
# input_ids
|
|
label_idx = None
|
|
if k == "input_ids" and "labels" in sample:
|
|
label_idx = np.argmax(np.array(sample["labels"]) != -100)
|
|
out[f"{k}_len"] = len(sample[k][:label_idx])
|
|
return out
|
|
|
|
|
|
def truncate_middle_if_too_long(
|
|
sample: dict[str, any],
|
|
max_length: int,
|
|
columns: list[str],
|
|
max_new_tokens: int = 256,
|
|
) -> dict[str, any]:
|
|
"""
|
|
Truncate the middle of a list of tokens to fit within a maximum length.
|
|
|
|
Args:
|
|
tokens: List of token IDs
|
|
max_length: Maximum length for the truncated tokens
|
|
|
|
Returns:
|
|
List of truncated token IDs
|
|
"""
|
|
max_new_tokens_half = max_new_tokens // 2
|
|
# leave max_new_tokens for generation
|
|
half = max_length // 2 - max_new_tokens_half
|
|
for col in columns:
|
|
t = sample[col]
|
|
sample[col] = t[:half] + t[-half:] if len(t) > max_length else t
|
|
return sample
|
|
|
|
|
|
def tokenize_ctx_text(
|
|
samples: dict[str, Any],
|
|
tokenizer: PreTrainedTokenizerBase,
|
|
) -> dict[str, Any]:
|
|
if tokenizer.chat_template:
|
|
tokenized_text = tokenizer.apply_chat_template(
|
|
[
|
|
[
|
|
{"role": "system", "content": ""},
|
|
{"role": "user", "content": ctx.strip()},
|
|
]
|
|
if isinstance(ctx, str)
|
|
else ctx
|
|
for ctx in samples["context"]
|
|
],
|
|
tokenize=True,
|
|
add_generation_prompt=True,
|
|
return_attention_mask=False,
|
|
padding=False,
|
|
truncation=False,
|
|
add_special_tokens=False, # special tokens are already added by the chat template
|
|
return_dict=True,
|
|
)
|
|
else:
|
|
raise NotImplementedError("Only support chat models.")
|
|
|
|
ctx_ids = tokenized_text["input_ids"]
|
|
return dict(ctx_ids=ctx_ids)
|
|
|
|
|
|
def pack(
|
|
ds_dict: dict[str, Dataset],
|
|
max_packed_inp_len: int,
|
|
max_packed_ctx_len: int,
|
|
max_packed_size: int,
|
|
seed: int,
|
|
num_proc: int = 0,
|
|
):
|
|
kwargs = dict(
|
|
max_packed_inp_len=max_packed_inp_len,
|
|
max_packed_ctx_len=max_packed_ctx_len,
|
|
max_packed_size=max_packed_size,
|
|
)
|
|
train_ds_lens = [len(ds) for ds in ds_dict.values()]
|
|
total_samples = sum(train_ds_lens)
|
|
logging.info(f"Total samples before packing: {total_samples}")
|
|
logging.info("Packing dataset")
|
|
|
|
sorted_keys = sorted(ds_dict)
|
|
ds_fingerprint = "|".join([ds_dict[k]._fingerprint for k in sorted_keys])
|
|
ds_hash = sha256((ds_fingerprint + json.dumps(kwargs)).encode()).hexdigest()
|
|
ds_path = f"{TRANSFORMED_DATA_DIR}/packed_{ds_hash}"
|
|
logger.info(
|
|
f"Packing ds {ds_hash} with {max_packed_inp_len=} and {max_packed_ctx_len=}"
|
|
)
|
|
if path.exists(ds_path) and is_caching_enabled():
|
|
logger.info(f"Loading a cached packed dataset for {ds_path}")
|
|
packed_ds = datasets.load_from_disk(ds_path)
|
|
else:
|
|
train_ds = interleave_datasets(
|
|
list(ds_dict.values()),
|
|
probabilities=get_ds_prob(train_ds_lens, total_samples),
|
|
seed=seed,
|
|
stopping_strategy="all_exhausted",
|
|
)
|
|
logger.info(f"Train dataset length: {len(train_ds)}")
|
|
packed_ds = train_ds.map(
|
|
pack_batch,
|
|
fn_kwargs={
|
|
"max_packed_inp_len": max_packed_inp_len,
|
|
"max_packed_ctx_len": max_packed_ctx_len,
|
|
"max_packed_size": max_packed_size,
|
|
"metadata_path": f"{ds_path}/packing_metadata.json",
|
|
},
|
|
batched=True,
|
|
batch_size=125_000,
|
|
num_proc=num_proc,
|
|
remove_columns=train_ds.column_names,
|
|
)
|
|
# this would generate another cache file for the already concat'd + packed ds
|
|
# TODO: saving here is not space efficient at all...
|
|
# the contexts are being duplicated for each datapoint when splitting QAs
|
|
packed_ds.save_to_disk(ds_path, num_proc=num_proc)
|
|
|
|
logger.info(f"Packed dataset length: {len(packed_ds)}")
|
|
logger.info(
|
|
f"Avg. # of samples per packed sequence: {total_samples / len(packed_ds)}"
|
|
)
|
|
|
|
return packed_ds
|