Download data/document_grounding.py from gaaaaaaaaaaa/multimodal-reasoning: direct link, hf CLI and curl.
- Browser
- Download file 1.88 kB
-
https://huggingface.co/gaaaaaaaaaaa/multimodal-reasoning/resolve/main/data/document_grounding.py
- Command line
-
hf download hf://gaaaaaaaaaaa/multimodal-reasoning/data/document_grounding.py
-
curl -L -o document_grounding.py https://huggingface.co/gaaaaaaaaaaa/multimodal-reasoning/resolve/main/data/document_grounding.py
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 |