Files
fml09 3950fda4e0 fix: bound MPS embedding memory on Apple Silicon (#244)
* fix: bound and recycle MPS embedding memory

* refactor: use CocoIndex GPU runner for MPS memory safety
2026-08-01 23:02:20 -07:00

216 lines
7.2 KiB
Python

"""Tests for embedder creation helpers."""
from __future__ import annotations
import os
from typing import Any
import numpy as np
import pytest
from cocoindex.ops.sentence_transformers import SentenceTransformerEmbedder
from cocoindex_code.litellm_embedder import PacedLiteLLMEmbedder
from cocoindex_code.settings import EmbeddingSettings
from cocoindex_code.shared import (
check_embedding,
configure_mps_environment,
create_embedder,
is_sentence_transformers_installed,
)
def test_create_embedder_uses_default_litellm_pacing() -> None:
embedder = create_embedder(
EmbeddingSettings(
provider="litellm",
model="text-embedding-3-small",
)
)
assert isinstance(embedder, PacedLiteLLMEmbedder)
assert embedder._min_request_interval_seconds == 0.005
def test_create_embedder_uses_paced_litellm_embedder() -> None:
embedder = create_embedder(
EmbeddingSettings(
provider="litellm",
model="text-embedding-3-small",
min_interval_ms=300,
)
)
assert isinstance(embedder, PacedLiteLLMEmbedder)
assert embedder._min_request_interval_seconds == 0.3
def test_create_embedder_litellm_passes_indexing_params_as_constructor_default() -> None:
"""Indexing params become default kwargs forwarded into every litellm call —
covering paths that don't go through INDEXING_EMBED_PARAMS (dim probe, etc.).
"""
embedder = create_embedder(
EmbeddingSettings(provider="litellm", model="cohere/embed-english-v3.0"),
indexing_params={"input_type": "search_document"},
)
assert isinstance(embedder, PacedLiteLLMEmbedder)
assert embedder._kwargs == {"input_type": "search_document"}
def test_create_embedder_sentence_transformers_ignores_indexing_params() -> None:
"""The SentenceTransformer constructor doesn't accept arbitrary kwargs;
indexing_params is silently ignored for that provider.
"""
embedder = create_embedder(
EmbeddingSettings(
provider="sentence-transformers", model="sentence-transformers/all-MiniLM-L6-v2"
),
indexing_params={"prompt_name": "passage"},
)
# No exception, and prompt_name is not stashed on the constructor —
# it's a per-call argument supplied via the embed() call site.
assert not isinstance(embedder, PacedLiteLLMEmbedder)
def test_create_embedder_uses_cocoindex_sentence_transformer_for_mps() -> None:
embedder = create_embedder(
EmbeddingSettings(
provider="sentence-transformers",
model="sentence-transformers/all-MiniLM-L6-v2",
device="mps",
)
)
assert isinstance(embedder, SentenceTransformerEmbedder)
def test_configure_mps_environment_enables_cocoindex_gpu_subprocess(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.delenv("COCOINDEX_RUN_GPU_IN_SUBPROCESS", raising=False)
monkeypatch.delenv("PYTORCH_MPS_HIGH_WATERMARK_RATIO", raising=False)
monkeypatch.delenv("PYTORCH_MPS_LOW_WATERMARK_RATIO", raising=False)
enabled = configure_mps_environment(
EmbeddingSettings(
provider="sentence-transformers",
model="sentence-transformers/all-MiniLM-L6-v2",
device="mps",
)
)
assert enabled is True
assert os.environ["COCOINDEX_RUN_GPU_IN_SUBPROCESS"] == "1"
assert os.environ["PYTORCH_MPS_LOW_WATERMARK_RATIO"] == "0.4"
assert os.environ["PYTORCH_MPS_HIGH_WATERMARK_RATIO"] == "0.5"
def test_explicit_mps_allocator_env_takes_precedence(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("COCOINDEX_RUN_GPU_IN_SUBPROCESS", "0")
monkeypatch.setenv("PYTORCH_MPS_LOW_WATERMARK_RATIO", "0.25")
monkeypatch.setenv("PYTORCH_MPS_HIGH_WATERMARK_RATIO", "0.3")
enabled = configure_mps_environment(
EmbeddingSettings(
provider="sentence-transformers",
model="sentence-transformers/all-MiniLM-L6-v2",
device="mps",
)
)
assert enabled is False
assert os.environ["COCOINDEX_RUN_GPU_IN_SUBPROCESS"] == "0"
assert os.environ["PYTORCH_MPS_LOW_WATERMARK_RATIO"] == "0.25"
assert os.environ["PYTORCH_MPS_HIGH_WATERMARK_RATIO"] == "0.3"
def test_configure_mps_environment_does_not_change_non_mps_provider(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.delenv("COCOINDEX_RUN_GPU_IN_SUBPROCESS", raising=False)
monkeypatch.delenv("PYTORCH_MPS_HIGH_WATERMARK_RATIO", raising=False)
monkeypatch.delenv("PYTORCH_MPS_LOW_WATERMARK_RATIO", raising=False)
enabled = configure_mps_environment(
EmbeddingSettings(provider="litellm", model="text-embedding-3-small")
)
assert enabled is False
assert "COCOINDEX_RUN_GPU_IN_SUBPROCESS" not in os.environ
assert "PYTORCH_MPS_LOW_WATERMARK_RATIO" not in os.environ
assert "PYTORCH_MPS_HIGH_WATERMARK_RATIO" not in os.environ
def test_configure_mps_environment_auto_detects_macos(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr("cocoindex_code.shared.sys.platform", "darwin")
monkeypatch.delenv("COCOINDEX_RUN_GPU_IN_SUBPROCESS", raising=False)
enabled = configure_mps_environment(
EmbeddingSettings(
provider="sentence-transformers",
model="sentence-transformers/all-MiniLM-L6-v2",
)
)
assert enabled is True
def test_is_sentence_transformers_installed_true_in_dev() -> None:
# Dev env pulls in sentence-transformers via the `dev` extras group.
assert is_sentence_transformers_installed() is True
def test_is_sentence_transformers_installed_false_when_find_spec_returns_none(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import importlib.util
monkeypatch.setattr(importlib.util, "find_spec", lambda name: None)
assert is_sentence_transformers_installed() is False
class _StubOkEmbedder:
def __init__(self) -> None:
self.last_kwargs: dict[str, Any] | None = None
async def embed(self, text: str, **kwargs: Any) -> Any:
self.last_kwargs = dict(kwargs)
return np.zeros(384, dtype=np.float32)
class _StubErrEmbedder:
async def embed(self, text: str, **kwargs: Any) -> Any:
raise RuntimeError("boom")
async def test_check_embedding_ok() -> None:
result = await check_embedding(_StubOkEmbedder())
assert result.error is None
assert result.dim == 384
assert result.traceback is None
async def test_check_embedding_error() -> None:
result = await check_embedding(_StubErrEmbedder())
assert result.dim is None
assert result.error is not None
assert result.error.startswith("RuntimeError:")
assert "boom" in result.error
# The full traceback is captured so `ccc doctor` can surface it for debugging.
assert result.traceback is not None
assert "Traceback (most recent call last):" in result.traceback
assert "boom" in result.traceback
async def test_check_embedding_forwards_params() -> None:
stub = _StubOkEmbedder()
await check_embedding(stub, {"prompt_name": "passage"})
assert stub.last_kwargs == {"prompt_name": "passage"}
async def test_check_embedding_no_params_forwards_empty() -> None:
stub = _StubOkEmbedder()
await check_embedding(stub)
assert stub.last_kwargs == {}