| """Calibration data loaders for Fisher factor estimation. | |
| We provide: | |
| - ``load_flores_texts(language, n_texts)`` — pulls FLORES+ dev/devtest text. | |
| - ``build_calibration_loader(...)`` — packs tokens into causal-LM blocks and | |
| returns a ``DataLoader`` over `dict(input_ids=..., labels=...)` batches. | |
| Notes: | |
| - Causal LM loss is ``loss = model(input_ids=ids, labels=ids).loss``; the same | |
| shifted-CE used in pretraining. Suitable for Fisher estimation. | |
| - For BERT/MLM calibration, prefer ``mlm`` collator outside this module. | |
| """ | |
| from __future__ import annotations | |
| from collections.abc import Iterable | |
| from typing import Any | |
| import torch | |
| from datasets import load_dataset | |
| from torch.utils.data import DataLoader, Dataset | |
| FLORES_LANG_TO_CODE = { | |
| "en": "eng_Latn", | |
| "de": "deu_Latn", | |
| "fr": "fra_Latn", | |
| "es": "spa_Latn", | |
| "ru": "rus_Cyrl", | |
| "zh": "cmn_Hans", # FLORES+ uses cmn (Mandarin) not zho | |
| "ar": "arb_Arab", | |
| "hi": "hin_Deva", | |
| "tr": "tur_Latn", | |
| "sw": "swh_Latn", | |
| "ur": "urd_Arab", | |
| "vi": "vie_Latn", | |
| } | |
| def load_flores_texts( | |
| language: str, | |
| *, | |
| n_texts: int | None = None, | |
| splits: tuple[str, ...] = ("dev", "devtest"), | |
| cache_dir: str | None = None, | |
| ) -> list[str]: | |
| """Return a list of texts in the given language from FLORES+ ``splits``.""" | |
| code = FLORES_LANG_TO_CODE.get(language, language) | |
| texts: list[str] = [] | |
| for split in splits: | |
| ds = load_dataset("openlanguagedata/flores_plus", split=split, cache_dir=cache_dir) | |
| # FLORES+ rows have an "iso_639_3" + "iso_15924" pair plus "text". | |
| ds = ds.filter(lambda r: f"{r['iso_639_3']}_{r['iso_15924']}" == code) | |
| texts.extend(ds["text"]) | |
| if n_texts is not None and len(texts) >= n_texts: | |
| break | |
| if n_texts is not None: | |
| texts = texts[:n_texts] | |
| return texts | |
| class _CausalLMBlocks(Dataset): | |
| def __init__(self, input_ids: torch.Tensor): | |
| self.input_ids = input_ids # [N, L] | |
| def __len__(self) -> int: | |
| return self.input_ids.shape[0] | |
| def __getitem__(self, idx: int): | |
| ids = self.input_ids[idx] | |
| return {"input_ids": ids, "labels": ids.clone()} | |
| def _collate(batch: list[dict[str, torch.Tensor]]) -> dict[str, torch.Tensor]: | |
| return { | |
| "input_ids": torch.stack([b["input_ids"] for b in batch]), | |
| "labels": torch.stack([b["labels"] for b in batch]), | |
| } | |
| def build_calibration_loader( | |
| texts: Iterable[str], | |
| tokenizer: Any, | |
| *, | |
| sequence_length: int = 512, | |
| batch_size: int = 4, | |
| drop_last: bool = True, | |
| ) -> DataLoader: | |
| """Concatenate texts, chunk into ``sequence_length`` blocks, return loader.""" | |
| eos = tokenizer.eos_token_id | |
| if eos is None: | |
| eos = tokenizer.pad_token_id | |
| if eos is None: | |
| raise ValueError("Tokenizer must have eos_token or pad_token") | |
| ids: list[int] = [] | |
| for text in texts: | |
| ids.extend(tokenizer.encode(text, add_special_tokens=False)) | |
| ids.append(eos) | |
| n_blocks = len(ids) // sequence_length | |
| if n_blocks == 0: | |
| raise RuntimeError( | |
| f"Not enough tokens for one block of {sequence_length}; got {len(ids)}" | |
| ) | |
| block_ids = torch.tensor(ids[: n_blocks * sequence_length], dtype=torch.long).view( | |
| n_blocks, sequence_length | |
| ) | |
| return DataLoader( | |
| _CausalLMBlocks(block_ids), | |
| batch_size=batch_size, | |
| shuffle=False, | |
| drop_last=drop_last, | |
| collate_fn=_collate, | |
| ) | |
Xet Storage Details
- Size:
- 3.5 kB
- Xet hash:
- 3aba72c146b4dc20b1826be658df595678680b9246440d598976075688727112
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.