doc-to-lora/src/ctx_to_lora/data/processing.py
Rujikorn Charakorn 7ac74e6d58
SFT data (#13)
* sft data from smoltalk + multi-turn ctx data

* add ctx-q-gsm (stateful-gsm from sleep-time compute)
2025-08-26 18:14:40 +09:00

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