multimodal-reasoning / data /document_grounding.py
gaaaaaaaaaaa's picture
Upload 52 files
a797f9a verified
Raw History Blame Contribute Delete
1.88 kB
from __future__ import annotations
from typing import Any
from datasets import load_dataset
import utils.configs_loader as ProjectConfigs
def load_document_grounding_dpo(tokenizer, dataset_key: str = "DOCUMENT_GROUNDING_DPO"):
cfg = ProjectConfigs.get(("DATASET", dataset_key))
ds = load_dataset(
"json",
data_files=cfg.Path,
split=cfg.Split,
)
test_size = getattr(cfg, "TestSize", 0.0)
seed = getattr(cfg, "Seed", 42)
if test_size and test_size > 0:
split = ds.train_test_split(
test_size=test_size,
seed=seed,
)
train_ds = split["train"]
eval_ds = split["test"]
else:
train_ds = ds
eval_ds = None
def convert(example: dict[str, Any]) -> dict[str, str]:
prompt_messages = example["prompt"]
chosen_messages = example["chosen"]
rejected_messages = example["rejected"]
prompt_text = tokenizer.apply_chat_template(
prompt_messages,
tokenize=False,
add_generation_prompt=True,
)
chosen_text = chosen_messages[0]["content"]
rejected_text = rejected_messages[0]["content"]
if tokenizer.eos_token:
chosen_text += tokenizer.eos_token
rejected_text += tokenizer.eos_token
return {
"prompt": prompt_text,
"chosen": chosen_text,
"rejected": rejected_text,
}
remove_cols = train_ds.column_names
train_ds = train_ds.map(
convert,
remove_columns=remove_cols,
desc="Converting document grounding DPO train set",
)
if eval_ds is not None:
eval_ds = eval_ds.map(
convert,
remove_columns=remove_cols,
desc="Converting document grounding DPO eval set",
)
return train_ds, eval_ds