def __init__(
self,
name: str,
*,
pat_str: str | None = None,
mergeable_ranks: Mapping[bytes, int] | None = None,
special_tokens: Mapping[str, int] | None = None,
explicit_n_vocab: int | None = None,
unit: Literal["byte", "unicode"] = "byte",
format: Literal["gpt2", "unitoken"] | None = None,
_encoder: BpeEncoder | None = None,
_ordinary_encoder: BpeEncoder | None = None,
_token_bytes: Mapping[int, bytes] | None = None,
_merges: Sequence[tuple[bytes, bytes]] | None = None,
) -> None:
self.name = name
self._pat_str = pat_str
self._special_tokens = dict(special_tokens or {})
self._special_tokens_set = set(self._special_tokens)
self._mergeable_ranks = dict(mergeable_ranks or {})
self._explicit_n_vocab = explicit_n_vocab
self._unit = unit
self._format = format or ("unitoken" if unit == "unicode" else "gpt2")
token_bytes = dict(_token_bytes or {})
for token, idx in self._mergeable_ranks.items():
token_bytes[idx] = token
for token, idx in self._special_tokens.items():
token_bytes[idx] = token.encode("utf-8")
self._token_bytes = token_bytes
if _encoder is not None:
self._encoder = _encoder
else:
vocab = {token: idx for idx, token in self._token_bytes.items()}
merges = list(_merges) if _merges is not None else _infer_merges_from_ranks(self._mergeable_ranks)
self._encoder = BpeEncoder(
unit=unit,
special_tokens=list(self._special_tokens),
merges=merges,
vocab=vocab,
pat_str=pat_str,
)
if _ordinary_encoder is not None:
self._ordinary_encoder = _ordinary_encoder
else:
vocab = {token: idx for idx, token in self._token_bytes.items()}
merges = list(_merges) if _merges is not None else _infer_merges_from_ranks(self._mergeable_ranks)
self._ordinary_encoder = BpeEncoder(
unit=unit,
special_tokens=[],
merges=merges,
vocab=vocab,
pat_str=pat_str,
)
if explicit_n_vocab is not None and self.n_vocab != explicit_n_vocab:
raise ValueError(f"explicit_n_vocab={explicit_n_vocab} does not match n_vocab={self.n_vocab}")