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