diff --git a/src/memos/configs/embedder.py b/src/memos/configs/embedder.py index 050043ab0..0b4da0d2f 100644 --- a/src/memos/configs/embedder.py +++ b/src/memos/configs/embedder.py @@ -13,8 +13,8 @@ class BaseEmbedderConfig(BaseConfig): default=None, description="Number of dimensions for the embedding" ) max_tokens: int | None = Field( - default=8192, - description="Maximum number of tokens per text. Texts exceeding this limit will be automatically truncated. Set to None to disable truncation.", + default=None, + description="Maximum number of tokens per text. When None (default) the embedder falls back to a provider-aware safe limit (currently 3072 tokens, matching text-embedding-3). Set to 0 or a concrete positive integer to explicitly disable or override truncation.", ) headers_extra: dict[str, Any] | None = Field( default=None, diff --git a/src/memos/embedders/base.py b/src/memos/embedders/base.py index 4cc8e6c85..69bcc3988 100644 --- a/src/memos/embedders/base.py +++ b/src/memos/embedders/base.py @@ -13,6 +13,8 @@ logger = get_logger(__name__) EmbeddingCallable = TypeVar("EmbeddingCallable", bound=Callable[..., Any]) +_SAFE_EMBEDDING_MAX_TOKENS: int = 3072 + def log_embedding_call(func: EmbeddingCallable) -> EmbeddingCallable: """Log embedding request dimensions and timing without text or vectors.""" @@ -136,26 +138,58 @@ def __init__(self, config: BaseEmbedderConfig): def _truncate_texts(self, texts: list[str], approx_char_per_token=1.0) -> (list)[str]: """ - Truncate texts to fit within max_tokens limit if configured. + Truncate texts to fit within a per-text token limit. + + When the configuration does not specify an explicit limit, the + per-provider default supplied by :meth:`_effective_max_tokens` is + used. This avoids hard-coded 8192 limits that exceed the actual + token budget (e.g. text-embedding-3's 3072-token input limit). + + Truncation is token-aware: each text is measured with + :func:`_count_tokens_for_embedding` and, if necessary, truncated via + :func:`_truncate_text_to_tokens` (binary search over prefix length). Args: texts: List of texts to truncate. + approx_char_per_token: Ignored; retained for backwards + compatibility with older call sites. Returns: List of truncated texts. """ - if not hasattr(self, "config") or self.config.max_tokens is None: + if not hasattr(self, "config"): + return texts + max_tokens = self._effective_max_tokens() + if max_tokens is None or max_tokens <= 0: return texts - max_tokens = self.config.max_tokens - truncated = [] + truncated: list[str] = [] for t in texts: - if len(t) < max_tokens * approx_char_per_token: + # Cheap fast-path: if the text is obviously too short to exceed + # the token budget, skip the (potentially expensive) token count. + # ~1 char/token for CJK text is the worst case. + if len(t) <= max_tokens: truncated.append(t) - else: - truncated.append(t[:max_tokens]) + continue + truncated.append(_truncate_text_to_tokens(t, max_tokens)) return truncated + def _effective_max_tokens(self) -> int | None: + """Return the per-text token limit used for truncation. + + Callers can override ``self.config.max_tokens`` explicitly. When it + is not set (``None`` or ``0`` after coercion through the config) + this falls back to a safe provider-aware default so that APIs like + text-embedding-3 (3072 input tokens) are never sent inputs that + would otherwise trigger an out-of-budget error. + """ + config = getattr(self, "config", None) + if config is not None: + configured = getattr(config, "max_tokens", None) + if configured is not None and configured > 0: + return configured + return _SAFE_EMBEDDING_MAX_TOKENS + @abstractmethod def embed(self, texts: list[str]) -> list[list[float]]: """Generate embeddings for the given texts.""" diff --git a/tests/embedders/test_base.py b/tests/embedders/test_base.py index 47a05864d..e59496631 100644 --- a/tests/embedders/test_base.py +++ b/tests/embedders/test_base.py @@ -3,7 +3,14 @@ import pytest -from memos.embedders.base import BaseEmbedder, log_embedding_call +from memos.configs.embedder import BaseEmbedderConfig +from memos.embedders.base import ( + _SAFE_EMBEDDING_MAX_TOKENS, + BaseEmbedder, + _count_tokens_for_embedding, + _truncate_text_to_tokens, + log_embedding_call, +) from tests.utils import check_module_base_class @@ -11,6 +18,89 @@ def test_base_embedder_class(): check_module_base_class(BaseEmbedder) +class _ConcreteEmbedder(BaseEmbedder): + def embed(self, texts: list[str]) -> list[list[float]]: + return [[0.0] for _ in texts] + + +class TestEffectiveMaxTokensDefaults: + def test_none_max_tokens_falls_back_to_safe_default(self): + config = BaseEmbedderConfig( + model_name_or_path="text-embedding-3-large", + max_tokens=None, + ) + emb = _ConcreteEmbedder(config) + assert emb._effective_max_tokens() == _SAFE_EMBEDDING_MAX_TOKENS + assert _SAFE_EMBEDDING_MAX_TOKENS == 3072 + + def test_zero_max_tokens_falls_back_to_safe_default(self): + config = BaseEmbedderConfig( + model_name_or_path="text-embedding-3-large", + max_tokens=0, + ) + emb = _ConcreteEmbedder(config) + assert emb._effective_max_tokens() == _SAFE_EMBEDDING_MAX_TOKENS + + def test_explicit_max_tokens_honoured(self): + config = BaseEmbedderConfig( + model_name_or_path="text-embedding-3-large", + max_tokens=1024, + ) + emb = _ConcreteEmbedder(config) + assert emb._effective_max_tokens() == 1024 + + +class TestTruncateTextToTokens: + def test_short_text_returned_unchanged(self): + text = "short text" + assert _truncate_text_to_tokens(text, 1000) == text + + def test_empty_and_none_inputs(self): + assert _truncate_text_to_tokens("", 10) == "" + assert _truncate_text_to_tokens("any", 0) == "any" + assert _truncate_text_to_tokens("any", None) == "any" + + def test_long_cjk_text_truncates_within_budget(self): + text = "记忆内容测试" * 1000 # ~6000 chars -> ~6000 tokens with simple heuristic + limit = 500 + truncated = _truncate_text_to_tokens(text, limit) + assert _count_tokens_for_embedding(truncated) <= limit + assert len(truncated) >= limit + + +class TestTruncateTextsDefault3072: + def test_short_texts_untouched(self): + config = BaseEmbedderConfig( + model_name_or_path="text-embedding-3-large", + max_tokens=None, + ) + emb = _ConcreteEmbedder(config) + texts = ["a", "bb", "ccc"] + assert emb._truncate_texts(texts) == texts + + def test_very_long_text_is_truncated_to_3072_tokens(self): + config = BaseEmbedderConfig( + model_name_or_path="text-embedding-3-large", + max_tokens=None, + ) + emb = _ConcreteEmbedder(config) + long_text = "记忆内容长文本token" * 2000 # should clearly exceed 3072 tokens + result = emb._truncate_texts([long_text, "short"]) + assert len(result) == 2 + assert result[1] == "short" + assert _count_tokens_for_embedding(result[0]) <= 3072 + assert result[0] != long_text + + def test_explicit_override_disables_default(self): + config = BaseEmbedderConfig( + model_name_or_path="text-embedding-3-large", + max_tokens=10, + ) + emb = _ConcreteEmbedder(config) + result = emb._truncate_texts(["a" * 1000]) + assert _count_tokens_for_embedding(result[0]) <= 10 + + def test_log_embedding_call_records_safe_structured_summary(): class StubEmbedder: config = SimpleNamespace(model_name_or_path="embedding-model")