"""Pruned-vocab id remap for ALF-emb-micro.""" from __future__ import annotations from pathlib import Path class RemapTokenizer: def __init__(self, tok_dir: str | Path, keep_ids: list[int], unk_new: int | None = None): from transformers import AutoTokenizer import numpy as np import torch self.base = AutoTokenizer.from_pretrained(str(tok_dir)) self.map = {old: new for new, old in enumerate(keep_ids)} unk_old = self.base.unk_token_id self.unk_new = unk_new if unk_new is not None else self.map.get(unk_old, 0) vocab = int(getattr(self.base, "vocab_size", 0) or 0) size = max(vocab, max(keep_ids) + 1 if keep_ids else 1) lut = np.full(size, self.unk_new, dtype=np.int64) for old, new in self.map.items(): if 0 <= old < size: lut[old] = new self._lut_np = lut self._lut = torch.from_numpy(lut) def __call__(self, texts, **kwargs): enc = self.base(texts, **kwargs) ids = enc["input_ids"] if hasattr(ids, "clamp"): lut = self._lut.to(device=ids.device) safe = ids.clamp(0, lut.numel() - 1).long() enc["input_ids"] = lut[safe].to(dtype=ids.dtype) else: import numpy as np arr = np.asarray(ids, dtype=np.int64) mapped = self._lut_np[np.clip(arr, 0, len(self._lut_np) - 1)] tensors = kwargs.get("return_tensors") enc["input_ids"] = mapped if tensors in ("np", "pt") or hasattr(ids, "shape") else mapped.tolist() return enc