diff --git a/sdk/python/feast/context/__init__.py b/sdk/python/feast/context/__init__.py new file mode 100644 index 00000000000..62de13ebead --- /dev/null +++ b/sdk/python/feast/context/__init__.py @@ -0,0 +1,39 @@ +"""Token-budget-aware helpers for assembling LLM context in feature views. + +Plain functions over immutable values, so an OnDemandFeatureView calling them +gets the same prompt offline and online. +""" + +from feast.context.errors import ( + ContextError, + TokenBudgetExceededError, + TokenizerNotFoundError, + TokenizerUnavailableError, +) +from feast.context.token_budget import TokenBudget +from feast.context.tokenizer import ( + ApproximateTokenizer, + Cl100kBaseTokenizer, + O200kBaseTokenizer, + TiktokenTokenizer, + Tokenizer, + TokenizerName, + TokenizerSpec, + get_tokenizer, +) + +__all__ = [ + "ApproximateTokenizer", + "Cl100kBaseTokenizer", + "ContextError", + "O200kBaseTokenizer", + "TiktokenTokenizer", + "TokenBudget", + "TokenBudgetExceededError", + "Tokenizer", + "TokenizerName", + "TokenizerNotFoundError", + "TokenizerSpec", + "TokenizerUnavailableError", + "get_tokenizer", +] diff --git a/sdk/python/feast/context/errors.py b/sdk/python/feast/context/errors.py new file mode 100644 index 00000000000..2504eabebd9 --- /dev/null +++ b/sdk/python/feast/context/errors.py @@ -0,0 +1,53 @@ +from typing import Iterable + +from fastapi import status as HttpStatusCode + +from feast.errors import FeastError + + +class ContextError(FeastError): + """Base class for all feast.context failures.""" + + +class TokenizerNotFoundError(ContextError): + def __init__(self, name: str, available: Iterable[str]): + super().__init__( + f"Unknown tokenizer '{name}'. Built-in tokenizers: " + f"{', '.join(sorted(available))}. For your own, pass the class " + f"path of a Tokenizer subclass, e.g. 'my_pkg.MyTokenizer'." + ) + + def http_status_code(self) -> int: + return HttpStatusCode.HTTP_400_BAD_REQUEST + + +class TokenizerUnavailableError(ContextError): + """A known tokenizer cannot be loaded. + + Covers a missing dependency, and one present but unable to load its + encoding — tiktoken fetching a BPE file on a host with no network, say. + """ + + def __init__(self, name: str, reason: str, install_hint: str): + super().__init__( + f"Tokenizer '{name}' is unavailable: {reason}. Install its " + f"dependencies with: {install_hint}" + ) + self.tokenizer_name = name + + def http_status_code(self) -> int: + return HttpStatusCode.HTTP_400_BAD_REQUEST + + +class TokenBudgetExceededError(ContextError): + def __init__(self, requested: int, remaining: int, max_tokens: int): + super().__init__( + f"Cannot consume {requested} tokens: only {remaining} of " + f"{max_tokens} tokens remain in the budget." + ) + self.requested = requested + self.remaining = remaining + self.max_tokens = max_tokens + + def http_status_code(self) -> int: + return HttpStatusCode.HTTP_400_BAD_REQUEST diff --git a/sdk/python/feast/context/token_budget.py b/sdk/python/feast/context/token_budget.py new file mode 100644 index 00000000000..2b2ccc8657d --- /dev/null +++ b/sdk/python/feast/context/token_budget.py @@ -0,0 +1,163 @@ +"""Token accounting for prompt assembly.""" + +from dataclasses import dataclass, field +from typing import Optional + +from feast.context.errors import TokenBudgetExceededError +from feast.context.tokenizer import ( + ApproximateTokenizer, + Tokenizer, + TokenizerName, + TokenizerSpec, + get_tokenizer, +) + + +@dataclass(frozen=True) +class TokenBudget: + """An immutable token allowance measured with a specific tokenizer. + + ``consume()`` returns a new budget instead of mutating this one, so a + budget is safe to share across threads and to reuse between assemblies:: + + budget = TokenBudget.of(4096, "cl100k_base") + budget.count_tokens("Hello world") # -> 2 + budget = budget.consume("Hello world") + budget.remaining # -> 4094 + + Attributes: + max_tokens: Size of the allowance. Zero admits only empty content. + tokenizer: The tokenizer doing the counting. A name is accepted too + and resolved at construction; :meth:`of` types that properly and + can forbid the fallback to the estimate. + consumed: Tokens already spent. + """ + + max_tokens: int + tokenizer: Tokenizer = field(default_factory=get_tokenizer) + consumed: int = 0 + + def __post_init__(self) -> None: + if self.max_tokens < 0: + raise ValueError(f"max_tokens must not be negative, got {self.max_tokens}.") + if self.consumed < 0: + raise ValueError(f"consumed must not be negative, got {self.consumed}.") + if self.consumed > self.max_tokens: + raise ValueError( + f"consumed ({self.consumed}) must not exceed max_tokens " + f"({self.max_tokens})." + ) + if not isinstance(self.tokenizer, Tokenizer): + # A name from config lands here; resolve it once, so counting + # never re-enters resolution and equality stays value-based. + object.__setattr__(self, "tokenizer", get_tokenizer(self.tokenizer)) + + @classmethod + def of( + cls, + max_tokens: int, + tokenizer: TokenizerSpec = TokenizerName.CL100K_BASE, + *, + consumed: int = 0, + fallback: bool = True, + ) -> "TokenBudget": + """Build a budget from a tokenizer name, instance, or class path. + + Args: + fallback: Set False to raise instead of degrading to the character + estimate when the tokenizer cannot be loaded. + """ + return cls( + max_tokens=max_tokens, + tokenizer=get_tokenizer(tokenizer, fallback=fallback), + consumed=consumed, + ) + + @property + def remaining(self) -> int: + """Tokens still available.""" + return self.max_tokens - self.consumed + + @property + def is_exhausted(self) -> bool: + """Whether the allowance is fully spent.""" + return self.remaining == 0 + + @property + def tokenizer_name(self) -> str: + """Name of the tokenizer actually counting.""" + return self.tokenizer.name + + @property + def is_approximate(self) -> bool: + """Whether counts are estimated rather than exact. + + True when the requested tokenizer could not be loaded and the budget + fell back to counting characters. + """ + return isinstance(self.tokenizer, ApproximateTokenizer) + + def count_tokens(self, text: str) -> int: + """Tokens ``text`` would occupy, regardless of what remains.""" + return self.tokenizer.count_tokens(text) + + def fits(self, text: str) -> bool: + """Whether ``text`` fits in the remaining allowance.""" + return self.count_tokens(text) <= self.remaining + + def fits_tokens(self, tokens: int) -> bool: + """Whether ``tokens`` more tokens fit in the remaining allowance.""" + _check_not_negative(tokens) + return tokens <= self.remaining + + def consume(self, text: str) -> "TokenBudget": + """Charge ``text`` against the budget, returning the new budget. + + Raises: + TokenBudgetExceededError: ``text`` does not fit; guard with + ``fits()`` or use ``try_consume()``. + """ + return self.consume_tokens(self.count_tokens(text)) + + def consume_tokens(self, tokens: int) -> "TokenBudget": + """Charge ``tokens`` against the budget, returning the new budget. + + Raises: + TokenBudgetExceededError: ``tokens`` exceeds what remains. + """ + if not self.fits_tokens(tokens): + raise TokenBudgetExceededError(tokens, self.remaining, self.max_tokens) + return self._replace_consumed(self.consumed + tokens) + + def try_consume(self, text: str) -> Optional["TokenBudget"]: + """Charge ``text`` if it fits, else return None. + + Lets a caller keep whichever candidate sections still fit, without + catching exceptions. + """ + tokens = self.count_tokens(text) + if not self.fits_tokens(tokens): + return None + return self._replace_consumed(self.consumed + tokens) + + def reset(self) -> "TokenBudget": + """A budget with the same limit and tokenizer, nothing consumed.""" + return self._replace_consumed(0) + + def _replace_consumed(self, consumed: int) -> "TokenBudget": + return TokenBudget( + max_tokens=self.max_tokens, + tokenizer=self.tokenizer, + consumed=consumed, + ) + + def __str__(self) -> str: + return ( + f"TokenBudget({self.consumed}/{self.max_tokens} tokens used, " + f"tokenizer={self.tokenizer_name})" + ) + + +def _check_not_negative(tokens: int) -> None: + if tokens < 0: + raise ValueError(f"tokens must not be negative, got {tokens}.") diff --git a/sdk/python/feast/context/tokenizer.py b/sdk/python/feast/context/tokenizer.py new file mode 100644 index 00000000000..09af77def68 --- /dev/null +++ b/sdk/python/feast/context/tokenizer.py @@ -0,0 +1,237 @@ +"""Tokenizers and how a name resolves to one. + +Resolution follows the same convention as online store types in +``repo_config``: a built-in name maps to a class path, and anything else is +itself the fully-qualified path of a :class:`Tokenizer` subclass with a +no-argument constructor:: + + TokenBudget(max_tokens=4096, tokenizer="cl100k_base") + TokenBudget(max_tokens=4096, tokenizer="my_pkg.tokenizers.LlamaTokenizer") +""" + +import logging +import math +from abc import ABC, abstractmethod +from enum import Enum +from functools import lru_cache +from typing import Any, Union + +from feast.context.errors import TokenizerNotFoundError, TokenizerUnavailableError +from feast.errors import FeastInvalidBaseClass +from feast.importer import import_class + +logger = logging.getLogger(__name__) + +#: Average characters per token for English prose. +DEFAULT_CHARS_PER_TOKEN = 4 + +#: Quoted in error messages when tiktoken is missing. +TIKTOKEN_INSTALL_HINT = "pip install tiktoken" + + +class TokenizerName(str, Enum): + """Built-in tokenizer names.""" + + #: GPT-4, GPT-3.5-turbo, text-embedding-3-*. + CL100K_BASE = "cl100k_base" + #: GPT-4o and the o-series models. + O200K_BASE = "o200k_base" + #: Character estimate; no third-party dependency. + APPROXIMATE = "approximate" + + +TOKENIZER_CLASS_FOR_TYPE = { + TokenizerName.CL100K_BASE.value: "feast.context.tokenizer.Cl100kBaseTokenizer", + TokenizerName.O200K_BASE.value: "feast.context.tokenizer.O200kBaseTokenizer", + TokenizerName.APPROXIMATE.value: "feast.context.tokenizer.ApproximateTokenizer", +} + + +class Tokenizer(ABC): + """Counts the tokens a model would consume for a piece of text. + + Implementations must be stateless and thread-safe: one instance per name is + shared across every :class:`TokenBudget`. A third-party subclass needs a + no-argument constructor and a class name ending in ``Tokenizer``, which is + what resolution by class path looks for. + + Two tokenizers that count identically compare equal, so budgets built from + the same name are interchangeable as dict keys and set members. + """ + + @property + @abstractmethod + def name(self) -> str: + """Identifier reported in budgets and logs.""" + + @abstractmethod + def count_tokens(self, text: str) -> int: + """Number of tokens in ``text``. Returns 0 for the empty string.""" + + def _identity(self) -> tuple[object, ...]: + """Whatever makes two instances count the same way.""" + return (self.name,) + + def __eq__(self, other: object) -> bool: + return isinstance(other, Tokenizer) and self._identity() == other._identity() + + def __hash__(self) -> int: + return hash(self._identity()) + + def __repr__(self) -> str: + return f"{type(self).__name__}(name={self.name!r})" + + +class ApproximateTokenizer(Tokenizer): + """Character-count estimate used when no real tokenizer is available. + + Counts ``ceil(len(text) / chars_per_token)``: close enough for budgeting, + not for exact context-window arithmetic. Code and non-Latin scripts drift + furthest from the ratio. + """ + + def __init__(self, chars_per_token: int = DEFAULT_CHARS_PER_TOKEN) -> None: + if chars_per_token <= 0: + raise ValueError( + f"chars_per_token must be positive, got {chars_per_token}." + ) + self._chars_per_token = chars_per_token + + @property + def name(self) -> str: + return TokenizerName.APPROXIMATE.value + + @property + def chars_per_token(self) -> int: + return self._chars_per_token + + def count_tokens(self, text: str) -> int: + return math.ceil(len(text) / self._chars_per_token) + + def _identity(self) -> tuple[object, ...]: + return (self.name, self._chars_per_token) + + +class TiktokenTokenizer(Tokenizer): + """Exact token counts from a tiktoken encoding. + + Raises: + TokenizerUnavailableError: tiktoken is missing or the encoding failed + to load. + TokenizerNotFoundError: tiktoken does not know ``encoding_name``. + """ + + def __init__(self, encoding_name: str) -> None: + self._encoding_name = encoding_name + self._encoding = _load_tiktoken_encoding(encoding_name) + + @property + def name(self) -> str: + return self._encoding_name + + def count_tokens(self, text: str) -> int: + if not text: + return 0 + # Feature values are arbitrary text: count a literal "<|endoftext|>" + # rather than raise, which is tiktoken's default for special tokens. + return len(self._encoding.encode(text, disallowed_special=())) + + +class Cl100kBaseTokenizer(TiktokenTokenizer): + """cl100k_base, in the no-argument form that resolution by name needs.""" + + def __init__(self) -> None: + super().__init__(TokenizerName.CL100K_BASE.value) + + +class O200kBaseTokenizer(TiktokenTokenizer): + """o200k_base, in the no-argument form that resolution by name needs.""" + + def __init__(self) -> None: + super().__init__(TokenizerName.O200K_BASE.value) + + +def _load_tiktoken_encoding(encoding_name: str) -> Any: + """Load a tiktoken encoding. tiktoken memoizes these internally.""" + try: + import tiktoken + except ImportError as e: + raise TokenizerUnavailableError( + encoding_name, f"tiktoken is not installed ({e})", TIKTOKEN_INSTALL_HINT + ) from e + + try: + return tiktoken.get_encoding(encoding_name) + except ValueError as e: + raise TokenizerNotFoundError( + encoding_name, tiktoken.list_encoding_names() + ) from e + except Exception as e: + # Encodings download on first use; an offline host fails here. + raise TokenizerUnavailableError( + encoding_name, + f"tiktoken failed to load the encoding ({e})", + TIKTOKEN_INSTALL_HINT, + ) from e + + +#: What callers may pass wherever a tokenizer is expected. +TokenizerSpec = Union[str, TokenizerName, Tokenizer] + + +def get_tokenizer( + spec: TokenizerSpec = TokenizerName.CL100K_BASE, + *, + fallback: bool = True, +) -> Tokenizer: + """Resolve ``spec`` to a tokenizer instance. + + A :class:`Tokenizer` passes through; a name is built once, then cached. + + Args: + spec: Tokenizer instance, built-in name, or the class path of a + ``Tokenizer`` subclass with a no-argument constructor. + fallback: On missing dependencies, return + :class:`ApproximateTokenizer` with a warning instead of raising. + Set False when exact counts are required. + + Raises: + TokenizerNotFoundError: ``spec`` is neither a built-in name nor a + class path ending in ``Tokenizer``. + TokenizerUnavailableError: the tokenizer cannot be built and + ``fallback`` is False. + """ + if isinstance(spec, Tokenizer): + return spec + name = spec.value if isinstance(spec, TokenizerName) else spec + if not isinstance(name, str) or not name.strip(): + raise TypeError(f"Tokenizer name must be a non-empty string, got {spec!r}.") + + try: + return _build_tokenizer(name.strip()) + except TokenizerUnavailableError as e: + if not fallback: + raise + logger.warning( + "%s Falling back to a character-based estimate; token counts will " + "be approximate.", + e, + ) + # Through the cache, so budgets that fell back still compare equal. + return _build_tokenizer(TokenizerName.APPROXIMATE.value) + + +@lru_cache(maxsize=None) +def _build_tokenizer(name: str) -> Tokenizer: + """Import and instantiate the tokenizer ``name`` refers to.""" + tokenizer_type = TOKENIZER_CLASS_FOR_TYPE.get(name, name) + if "." not in tokenizer_type or not tokenizer_type.endswith("Tokenizer"): + raise TokenizerNotFoundError(name, TOKENIZER_CLASS_FOR_TYPE.keys()) + + module_name, class_name = tokenizer_type.rsplit(".", 1) + tokenizer = import_class(module_name, class_name, "Tokenizer")() + if not isinstance(tokenizer, Tokenizer): + # import_class checks the base class by name, so an unrelated class + # called "Tokenizer" gets this far. + raise FeastInvalidBaseClass(tokenizer_type, "Tokenizer") + return tokenizer diff --git a/sdk/python/tests/unit/context/__init__.py b/sdk/python/tests/unit/context/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/sdk/python/tests/unit/context/conftest.py b/sdk/python/tests/unit/context/conftest.py new file mode 100644 index 00000000000..2bed5b15656 --- /dev/null +++ b/sdk/python/tests/unit/context/conftest.py @@ -0,0 +1,22 @@ +import builtins + +import pytest + +from feast.context import tokenizer as tokenizer_module + + +@pytest.fixture +def no_tiktoken(monkeypatch): + """Make ``import tiktoken`` fail, as it does when it is not installed.""" + real_import = builtins.__import__ + + def fake_import(name, *args, **kwargs): + if name == "tiktoken": + raise ImportError("No module named 'tiktoken'") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fake_import) + # Drop tokenizers an earlier test may have built and cached. + tokenizer_module._build_tokenizer.cache_clear() + yield + tokenizer_module._build_tokenizer.cache_clear() diff --git a/sdk/python/tests/unit/context/test_token_budget.py b/sdk/python/tests/unit/context/test_token_budget.py new file mode 100644 index 00000000000..9480c488b53 --- /dev/null +++ b/sdk/python/tests/unit/context/test_token_budget.py @@ -0,0 +1,200 @@ +import dataclasses +import inspect + +import pytest + +from feast.context import TokenBudget +from feast.context.errors import TokenBudgetExceededError, TokenizerNotFoundError +from feast.context.tokenizer import ( + ApproximateTokenizer, + Tokenizer, + TokenizerName, + get_tokenizer, +) + +# Four characters per token: "abcd" is one token under the estimate. +ESTIMATING = ApproximateTokenizer() + + +class HalvingTokenizer(Tokenizer): + """A third-party tokenizer, resolved by class path.""" + + @property + def name(self) -> str: + return "halving" + + def count_tokens(self, text: str) -> int: + return len(text) // 2 + + +def budget(max_tokens: int = 100, consumed: int = 0) -> TokenBudget: + """A budget on the deterministic estimate, so counts need no tiktoken.""" + return TokenBudget(max_tokens=max_tokens, tokenizer=ESTIMATING, consumed=consumed) + + +class TestConstruction: + def test_defaults_to_cl100k_base(self): + field = TokenBudget.__dataclass_fields__["tokenizer"] + assert field.default_factory is get_tokenizer + assert ( + inspect.signature(get_tokenizer).parameters["spec"].default + is TokenizerName.CL100K_BASE + ) + + def test_accepts_a_tokenizer_name(self): + assert TokenBudget.of(10, "approximate").tokenizer is get_tokenizer( + "approximate" + ) + + def test_accepts_a_tokenizer_instance(self): + assert TokenBudget(max_tokens=10, tokenizer=ESTIMATING).tokenizer is ESTIMATING + + def test_resolves_a_name_passed_to_the_constructor(self): + # Config supplies a string; the field still holds a Tokenizer. + resolved = TokenBudget(max_tokens=10, tokenizer="approximate").tokenizer # type: ignore[arg-type] + assert resolved is get_tokenizer("approximate") + + def test_resolves_a_custom_tokenizer_by_class_path(self): + b = TokenBudget.of(10, "tests.unit.context.test_token_budget.HalvingTokenizer") + assert b.tokenizer_name == "halving" + assert b.count_tokens("abcdefgh") == 4 + assert b.consume("abcdefgh").remaining == 6 + + def test_unknown_tokenizer_raises(self): + with pytest.raises(TokenizerNotFoundError): + TokenBudget.of(10, "gpt-42") + + def test_rejects_negative_max_tokens(self): + with pytest.raises(ValueError, match="max_tokens"): + budget(max_tokens=-1) + + def test_rejects_negative_consumed(self): + with pytest.raises(ValueError, match="consumed"): + budget(consumed=-1) + + def test_rejects_consumed_above_max_tokens(self): + with pytest.raises(ValueError, match="must not exceed"): + budget(max_tokens=10, consumed=11) + + def test_is_immutable(self): + with pytest.raises(dataclasses.FrozenInstanceError): + budget().max_tokens = 5 # type: ignore[misc] + + def test_budgets_with_the_same_state_are_equal_and_hash_alike(self): + # Default-constructed, so a tokenizer compared by identity would fail. + assert TokenBudget(max_tokens=10) == TokenBudget(max_tokens=10) + assert len({TokenBudget(max_tokens=10), TokenBudget(max_tokens=10)}) == 1 + + def test_replace_keeps_the_resolved_tokenizer(self): + replaced = dataclasses.replace(budget(max_tokens=10), consumed=4) + assert replaced.tokenizer is ESTIMATING + assert replaced.remaining == 6 + + +class TestCounting: + def test_counts_tokens(self): + assert budget().count_tokens("abcdefgh") == 2 + + def test_empty_string_costs_nothing(self): + assert budget().count_tokens("") == 0 + + def test_is_approximate_flags_the_estimate(self): + assert budget().is_approximate is True + + +class TestFits: + @pytest.mark.parametrize( + "max_tokens, text, expected", + [ + (10, "abcdefgh", True), + (2, "abcdefgh", True), # exactly filling the budget fits + (1, "a" * 40, False), + (0, "", True), # nothing always fits + ], + ) + def test_fits_compares_against_what_remains(self, max_tokens, text, expected): + assert budget(max_tokens=max_tokens).fits(text) is expected + + def test_fits_tokens_compares_against_what_remains(self): + b = budget(max_tokens=10, consumed=8) + assert b.fits_tokens(2) is True + assert b.fits_tokens(3) is False + + def test_fits_tokens_rejects_negative_counts(self): + with pytest.raises(ValueError, match="tokens"): + budget().fits_tokens(-1) + + +class TestConsume: + def test_returns_a_new_budget(self): + original = budget(max_tokens=10) + after = original.consume("abcdefgh") + assert after is not original + assert original.remaining == 10 + assert after.remaining == 8 + assert after.consumed == 2 + + def test_exhausts_the_budget_exactly(self): + after = budget(max_tokens=2).consume("abcdefgh") + assert after.remaining == 0 + assert after.is_exhausted is True + + def test_over_budget_raises_and_leaves_the_original_untouched(self): + b = budget(max_tokens=2) + with pytest.raises(TokenBudgetExceededError, match="only 2 of 2 tokens remain"): + b.consume("a" * 40) + assert b.remaining == 2 + + def test_consume_tokens_charges_a_count_directly(self): + assert budget(max_tokens=10).consume_tokens(4).remaining == 6 + + def test_consume_tokens_rejects_negative_counts(self): + with pytest.raises(ValueError, match="tokens"): + budget().consume_tokens(-1) + + +class TestTryConsume: + def test_returns_the_new_budget_when_content_fits(self): + after = budget(max_tokens=10).try_consume("abcdefgh") + assert after is not None + assert after.remaining == 8 + + def test_returns_none_when_content_does_not_fit(self): + assert budget(max_tokens=2).try_consume("a" * 40) is None + + def test_greedy_selection_keeps_what_fits(self): + b = budget(max_tokens=3) + kept = [] + for section in ["abcd", "a" * 40, "abcdefgh"]: + candidate = b.try_consume(section) + if candidate is not None: + b = candidate + kept.append(section) + assert kept == ["abcd", "abcdefgh"] + assert b.remaining == 0 + + +class TestReset: + def test_clears_consumption_and_keeps_limit_and_tokenizer(self): + after = budget(max_tokens=10).consume("abcdefgh").reset() + assert after.consumed == 0 + assert after == budget(max_tokens=10) + + +class TestTiktokenBudget: + @pytest.mark.parametrize( + "name", [TokenizerName.CL100K_BASE.value, TokenizerName.O200K_BASE.value] + ) + def test_counts_exact_tokens(self, name): + pytest.importorskip("tiktoken") + b = TokenBudget.of(4096, name) + assert b.tokenizer_name == name + assert b.is_approximate is False + assert b.count_tokens("Hello world") == 2 + assert b.consume("Hello world").remaining == 4094 + + def test_the_default_budget_counts_cl100k_base(self): + pytest.importorskip("tiktoken") + assert TokenBudget(max_tokens=4096).tokenizer_name == ( + TokenizerName.CL100K_BASE.value + ) diff --git a/sdk/python/tests/unit/context/test_tokenizer.py b/sdk/python/tests/unit/context/test_tokenizer.py new file mode 100644 index 00000000000..2b909cad65b --- /dev/null +++ b/sdk/python/tests/unit/context/test_tokenizer.py @@ -0,0 +1,141 @@ +import pytest + +from feast.context.errors import TokenizerNotFoundError, TokenizerUnavailableError +from feast.context.tokenizer import ( + ApproximateTokenizer, + Cl100kBaseTokenizer, + O200kBaseTokenizer, + TiktokenTokenizer, + TokenizerName, + get_tokenizer, +) +from feast.errors import FeastInvalidBaseClass + + +class FakeTokenizer: + """Named like a tokenizer, but not a Tokenizer subclass.""" + + +class TestApproximateTokenizer: + @pytest.mark.parametrize( + "text, expected", + [ + ("abcdefgh", 2), + ("abcde", 2), # partial tokens round up + ("", 0), + ("a" * 100_000, 25_000), + ], + ) + def test_counts_four_characters_per_token(self, text, expected): + assert ApproximateTokenizer().count_tokens(text) == expected + + def test_chars_per_token_is_configurable(self): + assert ApproximateTokenizer(chars_per_token=2).count_tokens("abcd") == 2 + + @pytest.mark.parametrize("chars_per_token", [0, -1]) + def test_rejects_non_positive_chars_per_token(self, chars_per_token): + with pytest.raises(ValueError): + ApproximateTokenizer(chars_per_token=chars_per_token) + + +class TestTiktokenTokenizer: + @pytest.mark.parametrize( + "encoding", [TokenizerName.CL100K_BASE.value, TokenizerName.O200K_BASE.value] + ) + def test_counts_tokens_for_supported_encodings(self, encoding): + pytest.importorskip("tiktoken") + tokenizer = TiktokenTokenizer(encoding) + assert tokenizer.name == encoding + assert tokenizer.count_tokens("Hello world") == 2 + assert tokenizer.count_tokens("") == 0 + # Unicode and long text: exact counts are the encoding's business, but + # neither may crash or come back empty. + assert tokenizer.count_tokens("北京の天気") > 0 + assert tokenizer.count_tokens("word " * 10_000) >= 10_000 + + def test_special_token_text_is_counted_not_rejected(self): + pytest.importorskip("tiktoken") + tokenizer = Cl100kBaseTokenizer() + assert tokenizer.count_tokens("<|endoftext|> appears in this review") > 0 + + def test_unknown_encoding_raises(self): + pytest.importorskip("tiktoken") + with pytest.raises(TokenizerNotFoundError): + TiktokenTokenizer("no_such_encoding") + + +class TestEquality: + def test_same_name_and_settings_are_equal(self): + assert ApproximateTokenizer() == ApproximateTokenizer() + assert len({ApproximateTokenizer(), ApproximateTokenizer()}) == 1 + + def test_different_settings_are_not_equal(self): + assert ApproximateTokenizer() != ApproximateTokenizer(chars_per_token=3) + + def test_same_encoding_from_different_classes_is_equal(self): + pytest.importorskip("tiktoken") + assert Cl100kBaseTokenizer() == TiktokenTokenizer( + TokenizerName.CL100K_BASE.value + ) + assert Cl100kBaseTokenizer() != O200kBaseTokenizer() + + +class TestGetTokenizer: + @pytest.mark.parametrize( + "name, expected_class", + [ + (TokenizerName.CL100K_BASE.value, Cl100kBaseTokenizer), + (TokenizerName.O200K_BASE.value, O200kBaseTokenizer), + (TokenizerName.APPROXIMATE.value, ApproximateTokenizer), + ], + ) + def test_resolves_builtin_names(self, name, expected_class): + if expected_class is not ApproximateTokenizer: + pytest.importorskip("tiktoken") + assert isinstance(get_tokenizer(name), expected_class) + + def test_resolves_a_class_path(self): + tokenizer = get_tokenizer("feast.context.tokenizer.ApproximateTokenizer") + assert isinstance(tokenizer, ApproximateTokenizer) + + def test_rejects_a_class_path_that_is_not_a_tokenizer(self): + with pytest.raises(FeastInvalidBaseClass): + get_tokenizer("tests.unit.context.test_tokenizer.FakeTokenizer") + + def test_rejects_a_path_that_is_not_a_tokenizer_class_name(self): + with pytest.raises(TokenizerNotFoundError): + get_tokenizer("feast.repo_config.RepoConfig") + + def test_unknown_name_raises_with_builtin_names(self): + with pytest.raises(TokenizerNotFoundError, match="cl100k_base"): + get_tokenizer("gpt-42") + + def test_builds_once_and_caches(self): + assert get_tokenizer("approximate") is get_tokenizer(" approximate ") + + def test_accepts_enum_member_and_equivalent_string(self): + assert get_tokenizer(TokenizerName.APPROXIMATE) is get_tokenizer("approximate") + + def test_tokenizer_instance_passes_through(self): + tokenizer = ApproximateTokenizer() + assert get_tokenizer(tokenizer) is tokenizer + + @pytest.mark.parametrize("name", [" ", None]) + def test_rejects_invalid_names(self, name): + with pytest.raises(TypeError): + get_tokenizer(name) + + +class TestFallback: + def test_falls_back_to_estimate_without_tiktoken(self, no_tiktoken, caplog): + tokenizer = get_tokenizer(TokenizerName.CL100K_BASE) + + assert isinstance(tokenizer, ApproximateTokenizer) + assert tokenizer.count_tokens("abcdefgh") == 2 + # The shared instance, so budgets that fell back stay comparable. + assert tokenizer is get_tokenizer(TokenizerName.APPROXIMATE) + assert any("tiktoken" in record.message for record in caplog.records) + + def test_fallback_disabled_raises(self, no_tiktoken): + with pytest.raises(TokenizerUnavailableError, match="pip install tiktoken"): + get_tokenizer(TokenizerName.CL100K_BASE, fallback=False)