doc-to-lora/hyperlora/data_utils.py
2025-01-13 16:45:26 +00:00

566 lines
18 KiB
Python

import logging
import numpy as np
from typing import Any, Callable, Iterator, Optional
from datasets import load_dataset, IterableDataset
from training_utils import TRAINING_TASK
from transformers import PreTrainedTokenizerBase
IGNORE_INDEX = -100
logger = logging.getLogger()
DS_KWARGS = {
"hotpot_qa": dict(
train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[1000:]"),
validation=dict(
path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]"
),
test=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"),
),
"hotpot_qa_tiny": dict(
train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[1000:2000]"),
validation=dict(
path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]"
),
test=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"),
),
"pwc": dict(
train=dict(path="sggetao/PwC", split="train[1000:]"),
validation=dict(path="sggetao/PwC", split="train[:1000]"),
test=dict(path="sggetao/PwC", split="test"),
),
"pwc_tiny": dict(
train=dict(path="sggetao/PwC", split="train[1000:2000]"),
validation=dict(path="sggetao/PwC", split="train[:1000]"),
test=dict(path="sggetao/PwC", split="test"),
),
"fw_qa_tiny": dict(
train=dict(
path="parquet",
data_files="data/raw_datasets/fw_qa/00000.parquet",
split="train",
),
validation=dict(
path="parquet",
data_files="data/raw_datasets/fw_qa/00000_val.parquet",
split="train",
),
),
}
def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]:
if (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}"
)
return kwargs
return DS_KWARGS[ds_name][split]
def validate_columns(tokenized_ds):
cols = ["input_ids", "attention_mask", "labels"]
if "ctx_ids" in tokenized_ds.column_names:
cols += ["ctx_ids", "ctx_attn_mask"]
if "chat_ids" in tokenized_ds.column_names:
cols += ["chat_ids", "chat_attn_mask", "chat_labels"]
ref_cols = set(cols)
assert (
set(tokenized_ds.column_names) == ref_cols
), f"Columns mismatch: {set(tokenized_ds.column_names)} != {ref_cols}"
def filter_long_samples(samples):
return [len(ctx) < 10000 for ctx in samples["context"]]
def add_repeat_prompt_fn(samples):
unique_contexts = set()
ctxs, prompts, responses = [], [], []
for ctx, prompt, response in zip(
samples["context"], samples["prompt"], samples["response"]
):
# Only process if the context is not already in the set
if ctx in unique_contexts:
continue
unique_contexts.add(ctx)
ctxs.append(ctx)
responses.append(ctx)
prompts.append("Repeat the text above.")
logger.debug(f"Adding repeat prompt...")
logger.debug(f"# unique contexts: {len(unique_contexts)}")
return dict(
context=ctxs + samples["context"],
response=responses + samples["response"],
prompt=prompts + samples["prompt"],
)
def add_negative_prompt_fn(samples):
unique_contexts = set()
ctxs, prompts, responses = [], [], []
keywords = [
"repeat",
"rephrase",
"summarize",
"rewrite",
"title",
"keyword",
"continuation",
]
for ctx, prompt, response in zip(
samples["context"], samples["prompt"], samples["response"]
):
if ctx in unique_contexts:
continue
if any(keyword in prompt for keyword in keywords):
# Skip samples where the prompt contains any of the specified keywords
continue
unique_contexts.add(ctx)
ctxs.append(ctx)
prompts.append(prompt)
responses.append(response)
logger.debug(f"Adding negative prompt...")
logger.debug(f"# unique contexts: {len(unique_contexts)}")
# remove one last sample if the number of samples is odd
if len(ctxs) % 2 != 0:
ctxs.pop()
prompts.pop()
responses.pop()
# to make sure that the negative prompt/response is not the same as the original
indices = list(np.random.permutation(len(ctxs))) + list(
np.random.permutation(len(ctxs))
)
neg_ctxs, neg_prompts, neg_responses = [], [], []
for idx in range(0, len(indices), 2):
i = indices[idx]
j = indices[idx + 1]
neg_ctxs.append(ctxs[i])
neg_prompts.append(ctxs[j] + "\n\n" + prompts[j])
neg_responses.append(responses[j])
return dict(
context=neg_ctxs + samples["context"],
prompt=neg_prompts + samples["prompt"],
response=neg_responses + samples["response"],
)
def filter_none(samples):
out = [True] * len(samples["context"])
for k in samples:
for i, sample in enumerate(samples[k]):
if not bool(sample):
out[i] = False
return out
def get_tokenized_dataset(
ds_name: str,
split: str,
tokenizer: PreTrainedTokenizerBase,
tokenizer_kwargs: dict[str, Any],
ctx_tokenizer: PreTrainedTokenizerBase,
ctx_tokenizer_kwargs: dict[str, Any],
add_ctx_to_chat: bool,
add_repeat_prompt: bool,
add_negative_prompt: bool,
use_kl_loss: bool,
) -> dict[str, Any]:
logger.debug(f"Loading dataset {ds_name} with split {split}...")
need_ctx_ids = not add_ctx_to_chat
try:
ds = load_dataset(**get_ds_kwargs(ds_name, split))
except ValueError as e:
logger.info(
f"Failed to load dataset {ds_name} with split {split}. Error: {e}\nSkipping..."
)
return None
cols_to_remove = [
col for col in ds.column_names if col not in ["context", "prompt", "response"]
]
ds = ds.map(get_preprocessing_fn(ds_name))
ds = ds.remove_columns(cols_to_remove)
ds = ds.filter(filter_none, batched=True)
ds = ds.filter(filter_long_samples, batched=True)
if split == "train":
if add_negative_prompt:
ds = ds.map(add_negative_prompt_fn, batched=True, batch_size=None)
if add_repeat_prompt and "context_numbers" not in ds_name:
ds = ds.map(add_repeat_prompt_fn, batched=True, batch_size=None)
tokenized_ds = construct_and_tokenize_ctx_qa(
tokenizer,
tokenizer_kwargs,
ctx_tokenizer,
ctx_tokenizer_kwargs,
add_ctx_to_chat,
use_kl_loss,
need_ctx_ids,
ds,
)
return tokenized_ds
def construct_and_tokenize_ctx_qa(
tokenizer,
tokenizer_kwargs,
ctx_tokenizer,
ctx_tokenizer_kwargs,
add_ctx_to_chat,
use_kl_loss,
need_ctx_ids,
ds,
):
# for sft + chat_model, we need to convert the dataset to chat format
# add "messages" field
ds = ds.map(
convert_ctx_prompt_response_to_messages,
fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat},
)
# add "chat" field
ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer))
# tokenize the chat + mask the assistant inputs
# add "input_ids", "attention_mask", "labels"
tokenized_ds = ds.map(
tokenize_chat_messages,
fn_kwargs={
"tokenizer": tokenizer,
"mask_assistant_inputs": True,
"tokenizer_kwargs": tokenizer_kwargs,
},
)
# for use_kl_loss, we need "chat_ids" and "chat_attn_mask"
if use_kl_loss:
raise NotImplementedError("KL loss deprecated")
tokenized_ds = tokenized_ds.map(
convert_ctx_prompt_response_to_messages,
fn_kwargs={"add_ctx_to_chat": True},
remove_columns=["messages"],
)
tokenized_ds = tokenized_ds.map(
get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer),
remove_columns=["chat"],
)
tokenized_ds = tokenized_ds.map(
tokenize_chat_messages,
fn_kwargs={
"tokenizer": tokenizer,
"mask_assistant_inputs": True,
"tokenizer_kwargs": tokenizer_kwargs,
"for_kl_loss": True,
},
num_proc=16,
)
if need_ctx_ids:
# TODO: can we batch this?
# TODO: can we cache this?
# tokenize the ctx_text to get ctx_ids and ctx_attn_mask
tokenized_ds = tokenized_ds.map(
tokenize_ctx_text,
fn_kwargs={"tokenizer": ctx_tokenizer},
batched=True,
num_proc=16,
)
tokenized_ds = tokenized_ds.remove_columns(
["messages", "chat", "context", "prompt", "response"]
)
tokenized_ds.set_format(type="pt")
validate_columns(tokenized_ds)
return tokenized_ds
def get_sft_prompt_formatting_fn(
sft_mode: TRAINING_TASK,
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 sft_mode != TRAINING_TASK.COMPLETION:
raise NotImplementedError(
f"Training task {sft_mode} not supported. "
"Only completion is supported for now."
)
if tokenizer.chat_template is None:
raise NotImplementedError(
"Only chat models + SFT are supported. "
"Training with pre-training data is not supported yet."
)
# TODO: add support for causal_lm training (for pre-training data)
# TODO: add support for recon training (for pre-training data)
# TODO: add support for non-chat models
# def f(example):
# output_texts = (
# dict(text=[]) if sft_mode == "causal_lm" else dict(prompt=[], response=[])
# )
# df = pd.DataFrame(dict(example))
# for i, inp_txt in df.iterrows():
# if sft_mode == "causal_lm":
# text = metadata["text_template"].format(**inp_txt)
# output_texts["text"].append(text)
# elif sft_mode == "completion":
# prompt = metadata["user_prompt_template"].format(**inp_txt)
# output_texts["prompt"].append(prompt)
# output_texts["response"].append(str(inp_txt[metadata["response_field"]]))
# return output_texts
def f_intx(example):
chat_text = tokenizer.apply_chat_template(
example["messages"], tokenize=False, add_generation_prompt=False
)
return dict(chat=chat_text)
# return f if not apply_chat_template_fn is not None else f_intx
return f_intx
def convert_ctx_prompt_response_to_messages(
example: dict[str, Any],
add_ctx_to_chat: bool = True,
) -> 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
"""
if "prompt" not in example or "response" not in example:
raise ValueError(f"'prompt' and 'response' are required. Got: {example}")
system_msg = ""
if "system_message" in example:
system_msg = example["system_message"].strip()
user_msg = example["prompt"].strip()
if add_ctx_to_chat:
user_msg = example["context"].strip() + "\n\n" + user_msg.strip()
messages = [
{"role": "system", "content": system_msg.strip()},
{"role": "user", "content": user_msg.strip()},
{"role": "assistant", "content": example["response"].strip()},
]
return dict(messages=messages)
def get_preprocessing_fn(ds_name: str) -> Callable[[dict[str, Any]], dict[str, Any]]:
"""
Get preprocessing function for a specific dataset.
Args:
ds_name: Name of the dataset
Returns:
A preprocessing function that takes and returns a dictionary
"""
f = lambda x: x
if ds_name.startswith("pwc"):
def f(sample):
return {
"context": sample["input"],
"prompt": sample["prompt"],
"response": sample["answer"],
}
elif ds_name.startswith("hotpot_qa"):
def f(sample):
txt = ""
for p in sample["context"]["sentences"]:
txt += " " + "".join(p)
return {
"context": txt.strip(),
"prompt": sample["question"],
"response": sample["answer"],
}
return f
# taken from https://github.com/huggingface/trl/issues/632#issuecomment-1972630547
def get_assistant_start_end_indices(
messages: list[dict[str, str]],
conversation_text: str,
) -> list[tuple[int, int]]:
"""
Get the start and end indices of assistant messages in conversation text.
Args:
messages: List of message dictionaries with 'role' and 'content' keys
conversation_text: Full conversation text
Returns:
List of (start, end) index tuples for assistant messages
"""
start_indices = []
# end_indices = []
current_index = 0
for message in messages:
message_text = message["content"]
match_index = conversation_text[current_index:].find(message_text)
start_indices.append(current_index + match_index)
# end_indices.append(current_index + match_index + len(message_text))
current_index += match_index + len(message_text)
end_indices = [
len(conversation_text) if i == len(start_indices) - 1 else start_indices[i + 1]
for i, x in enumerate(start_indices)
]
roles = [message["role"] for message in messages]
return [
(s, e) for s, e, r in zip(start_indices, end_indices, roles) if r == "assistant"
]
def get_masked_labels(
conversation_ids: dict[str, list[Any]],
assistant_ranges: list[tuple[int, int]],
) -> Iterator[int]:
"""
Generate masked labels for conversation, masking non-assistant tokens.
NOTE: This will also includes extra tokens between assistant and user messages
in multi-turn conversations.
E.g., {assistant_msg} <|eot_id|><|start_header_id|>user<|end_header_id|> {user_msg}
for Llama 3 models.
Args:
conversation_ids: Dictionary with tokenized conversation info
assistant_ranges: List of (start, end) indices for assistant messages
Yields:
Token ID or IGNORE_INDEX for each position
"""
for id_, (id_s, id_e) in list(
zip(conversation_ids["input_ids"], conversation_ids["offset_mapping"])
):
if any(id_s >= s and id_e <= e for s, e in assistant_ranges):
yield id_
else:
yield IGNORE_INDEX
def tokenize_chat_messages(
example: dict[str, Any],
tokenizer: PreTrainedTokenizerBase,
mask_assistant_inputs: bool = True,
for_kl_loss: bool = False,
tokenizer_kwargs: Optional[dict[str, Any]] = None,
) -> dict[str, list[int]]:
"""
Tokenize chat messages and optionally mask non-assistant tokens.
Args:
example: Dictionary containing 'chat' and 'messages' keys
tokenizer: Tokenizer to use
mask_assistant_inputs: Whether to mask non-assistant tokens
for_kl_loss: change the column names to "chat_ids" and "chat_attn_mask"
tokenizer_kwargs: Additional arguments to pass to tokenizer
Returns:
Dictionary containing tokenized inputs and labels
"""
# should be used only with chat models
text = example["chat"]
messages = example["messages"]
n_response = len([m for m in messages if m["role"] == "assistant"])
if n_response != 1:
raise ValueError(f"Expected 1 assistant response. Got {n_response}.")
conversation_ids = tokenizer(
text,
return_offsets_mapping=mask_assistant_inputs,
add_special_tokens=False,
truncation=False,
**(tokenizer_kwargs or {}),
)
if len(conversation_ids["input_ids"]) >= tokenizer_kwargs["max_length"]:
raise ValueError(
f"Conversation length {len(conversation_ids['input_ids'])} exceeds max length {tokenizer_kwargs['max_length']}"
)
if mask_assistant_inputs:
assistant_ranges = get_assistant_start_end_indices(messages, text)
labels = get_masked_labels(conversation_ids, assistant_ranges)
conversation_ids["labels"] = list(labels)
del conversation_ids["offset_mapping"]
else:
conversation_ids["labels"] = conversation_ids["input_ids"]
if for_kl_loss:
conversation_ids["chat_ids"] = conversation_ids.pop("input_ids")
conversation_ids["chat_attn_mask"] = conversation_ids.pop("attention_mask")
conversation_ids["chat_labels"] = conversation_ids.pop("labels")
return conversation_ids
def tokenize_ctx_text(
example: dict[str, Any],
tokenizer: PreTrainedTokenizerBase,
) -> dict[str, Any]:
text = example["context"]
tokenized_text = tokenizer(text)
ctx_ids = tokenized_text["input_ids"]
ctx_attn_mask = tokenized_text["attention_mask"]
return dict(ctx_ids=ctx_ids, ctx_attn_mask=ctx_attn_mask)
if __name__ == "__main__":
from transformers import AutoTokenizer
model_name = "meta-llama/Llama-3.2-1B-Instruct"
messages = [
{"role": "user", "content": "Hello!"},
{"role": "assistant", "content": "Hello!"},
# {"role": "user", "content": "Not too bad"},
# {"role": "assistant", "content": "Cooooooool"},
]
tokenizer = AutoTokenizer.from_pretrained(model_name)
chat = tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=False, add_special_tokens=False
)
print(chat)
model_inputs = tokenize_chat_messages(
{"chat": chat, "messages": messages},
tokenizer,
for_kl_loss=True,
)
print(tokenizer(chat, add_special_tokens=False))
print(model_inputs)