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