etomoscow/mff_lora / code /src /mfflora /data /calibration.py
etomoscow's picture
download
raw
3.5 kB
"""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.