Merge pull request #531 from AIOSAI/work/system-s135-maintenance-registry-noise-fix-doctor-provide
feat(system): S135 maintenance: registry noise fix, doctor provider-hooks check, stop-hook CWD scoping, branch settings cleanup
This commit is contained in:
@@ -29,18 +29,41 @@ def _find_repo_root() -> Path | None:
|
||||
AIPASS_ROOT = _find_repo_root()
|
||||
|
||||
|
||||
def _get_cwd_branch() -> str | None:
|
||||
"""Detect which branch directory (src/aipass/<name>) the CWD is in."""
|
||||
cwd = Path.cwd().resolve()
|
||||
if AIPASS_ROOT is None:
|
||||
return None
|
||||
src = AIPASS_ROOT / "src" / "aipass"
|
||||
try:
|
||||
rel = cwd.relative_to(src)
|
||||
return rel.parts[0] if rel.parts else None
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def get_modified_py_files() -> list[str]:
|
||||
"""Get Python files modified in the working tree (unstaged + staged)."""
|
||||
"""Get Python files modified in the working tree, scoped to the CWD branch.
|
||||
|
||||
Only returns files inside the current branch's directory (or repo-root files).
|
||||
This prevents dispatched agents' changes from triggering violations on the
|
||||
orchestrator or other agents sharing the worktree.
|
||||
"""
|
||||
if AIPASS_ROOT is None:
|
||||
return []
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", "diff", "--name-only", "HEAD"], capture_output=True, text=True, timeout=5, cwd=str(AIPASS_ROOT)
|
||||
)
|
||||
cwd_branch = _get_cwd_branch()
|
||||
files = []
|
||||
for line in result.stdout.strip().split("\n"):
|
||||
line = line.strip()
|
||||
if line.endswith(".py") and not line.startswith(".claude/"):
|
||||
if cwd_branch and line.startswith("src/aipass/"):
|
||||
file_branch = line.split("/")[2] if len(line.split("/")) > 2 else None
|
||||
if file_branch and file_branch != cwd_branch:
|
||||
continue
|
||||
full = AIPASS_ROOT / line
|
||||
if full.exists():
|
||||
files.append(str(full))
|
||||
|
||||
+1
-1
@@ -48,7 +48,7 @@ trinity = [
|
||||
memory = [
|
||||
"numpy>=2.0",
|
||||
"chromadb>=1.0",
|
||||
"sentence-transformers>=2.0",
|
||||
"fastembed>=0.4",
|
||||
]
|
||||
seedgo = []
|
||||
dev = [
|
||||
|
||||
@@ -1,97 +1,4 @@
|
||||
{
|
||||
"hooks": {
|
||||
"PostToolUse": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/auto_fix_diagnostics.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"PreToolUse": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/pre_edit_gate.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"Stop": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/subagent_stop_gate.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"PreCompact": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/pre_compact.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"UserPromptSubmit": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "cat .aipass/aipass_global_prompt.md 2>/dev/null || true"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 -c \"from pathlib import Path; p=next((x/'.aipass'/'aipass_local_prompt.md' for x in [Path.cwd(),*Path.cwd().parents] if (x/'.aipass'/'aipass_local_prompt.md').exists()),None); p and print(p.read_text(encoding='utf-8'),end='')\""
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/branch_prompt_loader.py"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/email_notification.py"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/identity_injector.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"env": {
|
||||
"AIPASS_HOME": "/home/patrick/Projects/AIPass"
|
||||
}
|
||||
|
||||
@@ -271,17 +271,19 @@ def _check_services(verbose: bool = False) -> List[CheckResult]:
|
||||
logger.warning("[doctor] pytest collect timed out: %s", exc)
|
||||
results.append(CheckResult("pytest collect", GLYPH_WARN, "timed out", ""))
|
||||
|
||||
# hooks wired
|
||||
auto_fix = Path("~/.claude/hooks/auto_fix_diagnostics.py").expanduser()
|
||||
if auto_fix.exists():
|
||||
results.append(CheckResult("hooks", GLYPH_PASS, "auto_fix_diagnostics wired", ""))
|
||||
# hooks wired — check provider-level enforcement hooks
|
||||
provider_hooks_dir = Path("~/.claude/hooks").expanduser()
|
||||
provider_hooks = ["auto_fix_diagnostics.py", "pre_edit_gate.py", "subagent_stop_gate.py"]
|
||||
missing = [h for h in provider_hooks if not (provider_hooks_dir / h).exists()]
|
||||
if not missing:
|
||||
results.append(CheckResult("hooks", GLYPH_PASS, f"{len(provider_hooks)} provider hooks wired", ""))
|
||||
else:
|
||||
results.append(
|
||||
CheckResult(
|
||||
"hooks",
|
||||
GLYPH_WARN,
|
||||
"auto_fix_diagnostics.py not found",
|
||||
"Run setup to wire Claude Code hooks",
|
||||
f"{len(missing)} provider hook(s) missing: {', '.join(missing)}",
|
||||
"Copy from .claude/hooks/ to ~/.claude/hooks/ — see .claude/hooks/README.md",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -1,97 +1,4 @@
|
||||
{
|
||||
"hooks": {
|
||||
"PostToolUse": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/auto_fix_diagnostics.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"PreToolUse": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/pre_edit_gate.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"Stop": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/subagent_stop_gate.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"PreCompact": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/pre_compact.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"UserPromptSubmit": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "cat .aipass/aipass_global_prompt.md 2>/dev/null || true"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 -c \"from pathlib import Path; p=next((x/'.aipass'/'aipass_local_prompt.md' for x in [Path.cwd(),*Path.cwd().parents] if (x/'.aipass'/'aipass_local_prompt.md').exists()),None); p and print(p.read_text(encoding='utf-8'),end='')\""
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/branch_prompt_loader.py"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/email_notification.py"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/identity_injector.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"env": {
|
||||
"AIPASS_HOME": "/home/patrick/Projects/AIPass"
|
||||
}
|
||||
|
||||
@@ -1,97 +1,4 @@
|
||||
{
|
||||
"hooks": {
|
||||
"PostToolUse": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/auto_fix_diagnostics.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"PreToolUse": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/pre_edit_gate.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"Stop": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/subagent_stop_gate.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"PreCompact": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/pre_compact.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"UserPromptSubmit": [
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "cat .aipass/aipass_global_prompt.md 2>/dev/null || true"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 -c \"from pathlib import Path; p=next((x/'.aipass'/'aipass_local_prompt.md' for x in [Path.cwd(),*Path.cwd().parents] if (x/'.aipass'/'aipass_local_prompt.md').exists()),None); p and print(p.read_text(encoding='utf-8'),end='')\""
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/branch_prompt_loader.py"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/email_notification.py"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"matcher": "",
|
||||
"hooks": [
|
||||
{
|
||||
"type": "command",
|
||||
"command": "python3 .claude/hooks/identity_injector.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"env": {
|
||||
"AIPASS_HOME": "/home/patrick/Projects/AIPass"
|
||||
}
|
||||
|
||||
@@ -128,10 +128,7 @@ def find_registry() -> Path:
|
||||
if hit is not None:
|
||||
if _registry_matches_credential(hit):
|
||||
return hit
|
||||
logger.warning(
|
||||
"Skipping mismatched registry at %s — continuing walk-up",
|
||||
hit,
|
||||
)
|
||||
continue
|
||||
|
||||
# AIPASS_HOME fallback — for external projects where CWD walk finds nothing
|
||||
aipass_home = os.environ.get("AIPASS_HOME")
|
||||
|
||||
@@ -251,14 +251,14 @@ def process_file_to_vectors(
|
||||
# Chunk content
|
||||
chunks = chunk_content(content, chunk_size, chunk_overlap)
|
||||
|
||||
# Import ChromaDB and sentence transformers (late import for venv compatibility)
|
||||
# Import ChromaDB and fastembed (late import for venv compatibility)
|
||||
try:
|
||||
import chromadb
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
client = chromadb.PersistentClient(path=str(CHROMA_PATH))
|
||||
collection = client.get_or_create_collection(name=collection_name)
|
||||
model = SentenceTransformer("all-MiniLM-L6-v2")
|
||||
model = TextEmbedding("sentence-transformers/all-MiniLM-L6-v2")
|
||||
|
||||
# Generate embeddings and store
|
||||
documents = []
|
||||
@@ -280,7 +280,7 @@ def process_file_to_vectors(
|
||||
)
|
||||
|
||||
# Batch encode
|
||||
embeddings = model.encode(documents).tolist()
|
||||
embeddings = [e.tolist() for e in model.embed(documents)]
|
||||
|
||||
# Upsert (update if exists, insert if not)
|
||||
collection.upsert(documents=documents, embeddings=embeddings, ids=ids, metadatas=metadatas)
|
||||
|
||||
@@ -24,8 +24,7 @@ Design:
|
||||
|
||||
Dependencies (optional):
|
||||
- chromadb
|
||||
- sentence-transformers
|
||||
- torch
|
||||
- fastembed
|
||||
"""
|
||||
|
||||
from typing import List, Dict, Any
|
||||
@@ -56,7 +55,7 @@ class QueryEncoder:
|
||||
Uses all-MiniLM-L6-v2 model with same settings.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = "all-MiniLM-L6-v2"):
|
||||
def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2"):
|
||||
"""
|
||||
Initialize query encoder
|
||||
|
||||
@@ -64,28 +63,17 @@ class QueryEncoder:
|
||||
model_name: HuggingFace model identifier (must match embedder.py)
|
||||
|
||||
Raises:
|
||||
ImportError: If sentence-transformers or torch are not installed
|
||||
ImportError: If fastembed is not installed
|
||||
"""
|
||||
# Late imports (heavy optional dependencies)
|
||||
try:
|
||||
import torch
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError as e:
|
||||
logger.info(f"[vector_search] Optional ML dependencies not available: {e}")
|
||||
raise ImportError(
|
||||
f"Search requires sentence-transformers and torch. "
|
||||
f"Install with: pip install sentence-transformers torch. "
|
||||
f"Original error: {e}"
|
||||
)
|
||||
raise ImportError(f"Search requires fastembed. Install with: pip install fastembed. Original error: {e}")
|
||||
|
||||
self.model_name = model_name
|
||||
self.model = SentenceTransformer(model_name)
|
||||
|
||||
# GPU optimization if available
|
||||
self.use_gpu = torch.cuda.is_available()
|
||||
if self.use_gpu:
|
||||
self.model = self.model.to("cuda")
|
||||
|
||||
self.model = TextEmbedding(model_name)
|
||||
self.dimension = 384 # all-MiniLM-L6-v2 output dimension
|
||||
|
||||
def encode(self, query: str) -> List[float]:
|
||||
@@ -98,21 +86,8 @@ class QueryEncoder:
|
||||
Returns:
|
||||
384-dimensional embedding as list
|
||||
"""
|
||||
import torch
|
||||
|
||||
# Encode with same settings as embedder.py
|
||||
embedding = self.model.encode(
|
||||
query,
|
||||
convert_to_tensor=False, # Return numpy
|
||||
normalize_embeddings=True, # Critical for L2 distance
|
||||
show_progress_bar=False,
|
||||
)
|
||||
|
||||
# Cleanup GPU memory if used
|
||||
if self.use_gpu:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return embedding.tolist()
|
||||
embeddings = list(self.model.embed([query]))
|
||||
return embeddings[0].tolist()
|
||||
|
||||
|
||||
# Global encoder instance (singleton pattern)
|
||||
|
||||
@@ -1,16 +1,16 @@
|
||||
# =================== AIPass ====================
|
||||
# Name: embed_subprocess.py
|
||||
# Description: Embedding Subprocess Handler
|
||||
# Version: 1.0.0
|
||||
# Version: 2.0.0
|
||||
# Created: 2026-03-12
|
||||
# Modified: 2026-03-12
|
||||
# Modified: 2026-05-07
|
||||
# =============================================
|
||||
|
||||
"""
|
||||
Embedding Subprocess Handler
|
||||
|
||||
Called via subprocess from rollover orchestrator to ensure sentence-transformers
|
||||
and torch run in the memory-specific venv (AIPASS_MEMORY_PYTHON).
|
||||
Called via subprocess from rollover orchestrator to generate embeddings
|
||||
using fastembed (ONNX runtime, no torch dependency).
|
||||
|
||||
Input: JSON on stdin with texts to encode
|
||||
Output: JSON on stdout with embeddings
|
||||
@@ -21,7 +21,7 @@ import json
|
||||
|
||||
|
||||
def main():
|
||||
"""Process embedding request from stdin JSON"""
|
||||
"""Process embedding request from stdin JSON."""
|
||||
try:
|
||||
input_data = json.load(sys.stdin)
|
||||
texts = input_data.get("texts", [])
|
||||
@@ -30,40 +30,19 @@ def main():
|
||||
print(json.dumps({"success": True, "embeddings": [], "count": 0, "dimension": 384}))
|
||||
return
|
||||
|
||||
# Import here — runs in memory venv where these are installed
|
||||
from sentence_transformers import SentenceTransformer
|
||||
import torch
|
||||
from fastembed import TextEmbedding
|
||||
|
||||
model = SentenceTransformer("all-MiniLM-L6-v2")
|
||||
model = TextEmbedding("sentence-transformers/all-MiniLM-L6-v2")
|
||||
|
||||
use_gpu = torch.cuda.is_available()
|
||||
if use_gpu:
|
||||
model = model.to("cuda")
|
||||
batch_size = 64
|
||||
else:
|
||||
batch_size = 16
|
||||
|
||||
# Pre-sort by length (reduces padding waste)
|
||||
sorted_pairs = sorted(enumerate(texts), key=lambda x: len(x[1]))
|
||||
sorted_indices, sorted_texts = zip(*sorted_pairs)
|
||||
|
||||
# Encode
|
||||
embeddings = model.encode(
|
||||
list(sorted_texts),
|
||||
batch_size=batch_size,
|
||||
convert_to_tensor=False,
|
||||
normalize_embeddings=True,
|
||||
show_progress_bar=False,
|
||||
)
|
||||
embeddings = list(model.embed(list(sorted_texts)))
|
||||
|
||||
# Restore original order
|
||||
ordered = [None] * len(texts)
|
||||
for orig_idx, sorted_idx in enumerate(sorted_indices):
|
||||
ordered[sorted_idx] = embeddings[orig_idx].tolist()
|
||||
|
||||
if use_gpu:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
print(json.dumps({"success": True, "embeddings": ordered, "count": len(ordered), "dimension": 384}))
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -1,31 +1,22 @@
|
||||
# =================== AIPass ====================
|
||||
# Name: embedder.py
|
||||
# Description: Vector Embedding Handler
|
||||
# Version: 0.2.0
|
||||
# Version: 0.3.0
|
||||
# Created: 2025-11-16
|
||||
# Modified: 2026-03-06
|
||||
# Modified: 2026-05-07
|
||||
# =============================================
|
||||
|
||||
"""
|
||||
Vector Embedding Handler
|
||||
|
||||
Generates semantic embeddings using sentence-transformers/all-MiniLM-L6-v2.
|
||||
Implements production best practices from research.
|
||||
Generates semantic embeddings using fastembed/all-MiniLM-L6-v2 (ONNX).
|
||||
|
||||
Purpose:
|
||||
Convert text memories into 384-dimensional vectors for semantic search.
|
||||
Optimized for batch processing (100 lines during rollover).
|
||||
|
||||
Best Practices Applied:
|
||||
- Pre-sort by length (30% padding reduction)
|
||||
- Built-in normalization (L2 distance requirement)
|
||||
- GPU memory cleanup (prevent VRAM leaks)
|
||||
- Batch size optimization (64 GPU, 16 CPU)
|
||||
- Singleton pattern (model loaded once)
|
||||
|
||||
Dependencies (optional):
|
||||
- sentence-transformers
|
||||
- torch
|
||||
- fastembed
|
||||
"""
|
||||
|
||||
from typing import List, Dict, Any
|
||||
@@ -35,147 +26,58 @@ from aipass.memory.apps.handlers.json import json_handler
|
||||
|
||||
logger = get_system_logger()
|
||||
|
||||
# No service imports - handlers are pure workers (3-tier architecture)
|
||||
# No module imports (handler independence)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# EMBEDDING SERVICE (Singleton)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class EmbeddingService:
|
||||
"""
|
||||
Production-ready embedding service
|
||||
"""Embedding service using fastembed (ONNX runtime, no torch dependency)."""
|
||||
|
||||
Implements best practices:
|
||||
- Batch size optimization (64 GPU, 16 CPU)
|
||||
- Pre-sorting by length (reduces padding waste 30%)
|
||||
- Built-in normalization (critical for L2 distance)
|
||||
- GPU memory cleanup (prevents VRAM leaks)
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = "all-MiniLM-L6-v2"):
|
||||
"""
|
||||
Initialize embedding service
|
||||
|
||||
Args:
|
||||
model_name: HuggingFace model identifier
|
||||
|
||||
Raises:
|
||||
ImportError: If sentence-transformers or torch are not installed
|
||||
"""
|
||||
# Late imports (heavy optional dependencies)
|
||||
def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2"):
|
||||
try:
|
||||
import torch
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError as e:
|
||||
logger.info(f"[embedder] Optional ML dependencies not available: {e}")
|
||||
raise ImportError(
|
||||
f"Embedding requires sentence-transformers and torch. "
|
||||
f"Install with: pip install sentence-transformers torch. "
|
||||
f"Original error: {e}"
|
||||
)
|
||||
raise ImportError(f"Embedding requires fastembed. Install with: pip install fastembed. Original error: {e}")
|
||||
|
||||
self.model_name = model_name
|
||||
self.model = SentenceTransformer(model_name)
|
||||
|
||||
# GPU optimization if available
|
||||
self.use_gpu = torch.cuda.is_available()
|
||||
if self.use_gpu:
|
||||
self.model = self.model.to("cuda")
|
||||
self.batch_size = 64
|
||||
else:
|
||||
self.batch_size = 16
|
||||
|
||||
self.dimension = 384 # all-MiniLM-L6-v2 output dimension
|
||||
self.model = TextEmbedding(model_name)
|
||||
self.dimension = 384
|
||||
|
||||
def encode_batch(self, texts: List[str]) -> Dict[str, Any]:
|
||||
"""
|
||||
Encode batch of texts with all optimizations
|
||||
|
||||
Best practices applied:
|
||||
1. Pre-sort by length (reduces padding waste)
|
||||
2. Batch processing (optimal batch size)
|
||||
3. Built-in normalization (L2 distance requirement)
|
||||
4. GPU cleanup (prevent VRAM leaks)
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
Dict with embeddings and metadata
|
||||
"""
|
||||
import torch
|
||||
|
||||
"""Encode texts to embeddings with pre-sort optimization."""
|
||||
if not texts:
|
||||
return {"embeddings": [], "count": 0, "dimension": self.dimension}
|
||||
|
||||
# Pre-sort by length (reduces padding waste by 30%)
|
||||
sorted_pairs = sorted(enumerate(texts), key=lambda x: len(x[1]))
|
||||
sorted_indices: list[int] = [p[0] for p in sorted_pairs]
|
||||
sorted_text_list: list[str] = [p[1] for p in sorted_pairs]
|
||||
|
||||
# Encode with optimal settings
|
||||
embeddings = self.model.encode(
|
||||
sorted_text_list,
|
||||
batch_size=self.batch_size,
|
||||
convert_to_tensor=False, # Return numpy for Chroma
|
||||
normalize_embeddings=True, # Critical for L2 distance
|
||||
show_progress_bar=False,
|
||||
)
|
||||
embeddings = list(self.model.embed(sorted_text_list))
|
||||
|
||||
# Restore original order
|
||||
ordered_embeddings: List[Any] = [None] * len(texts)
|
||||
for original_idx, sorted_idx in enumerate(sorted_indices):
|
||||
ordered_embeddings[sorted_idx] = embeddings[original_idx]
|
||||
|
||||
# Cleanup GPU memory if used
|
||||
if self.use_gpu:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return {"embeddings": ordered_embeddings, "count": len(ordered_embeddings), "dimension": self.dimension}
|
||||
|
||||
|
||||
# Global service instance (singleton pattern)
|
||||
_embedding_service = None
|
||||
|
||||
|
||||
def _get_service() -> EmbeddingService:
|
||||
"""
|
||||
Get or create embedding service singleton
|
||||
|
||||
Lazy initialization - model loaded on first use
|
||||
"""
|
||||
global _embedding_service
|
||||
if _embedding_service is None:
|
||||
_embedding_service = EmbeddingService()
|
||||
return _embedding_service
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# PUBLIC API
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def encode_batch(texts: List[str]) -> Dict[str, Any]:
|
||||
"""
|
||||
Encode batch of texts to embeddings
|
||||
|
||||
This is the main public API. Delegates to singleton service
|
||||
to avoid reloading the model.
|
||||
Encode batch of texts to embeddings.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
Dict with embeddings and metadata
|
||||
|
||||
Example:
|
||||
result = encode_batch(["memory 1", "memory 2"])
|
||||
if result['success']:
|
||||
embeddings = result['embeddings']
|
||||
# Each embedding is 384-dim numpy array
|
||||
"""
|
||||
if not texts:
|
||||
return {"success": True, "embeddings": [], "count": 0, "message": "No texts provided"}
|
||||
@@ -197,45 +99,27 @@ def encode_batch(texts: List[str]) -> Dict[str, Any]:
|
||||
|
||||
def encode_memories(memories: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""
|
||||
Encode memory entries to embeddings
|
||||
|
||||
Extracts text from memory entries and encodes them.
|
||||
Preserves original memory structure for metadata.
|
||||
Encode memory entries to embeddings.
|
||||
|
||||
Args:
|
||||
memories: List of memory entry dicts (from extraction)
|
||||
|
||||
Returns:
|
||||
Dict with embeddings and original memories
|
||||
|
||||
Example:
|
||||
memories = [{"content": "...", "timestamp": "..."}]
|
||||
result = encode_memories(memories)
|
||||
embeddings = result['embeddings']
|
||||
original = result['memories']
|
||||
"""
|
||||
if not memories:
|
||||
return {"success": True, "embeddings": [], "memories": [], "count": 0, "message": "No memories provided"}
|
||||
|
||||
# Extract text content from memories
|
||||
texts = []
|
||||
for memory in memories:
|
||||
# Try common fields for text content
|
||||
text = (
|
||||
memory.get("content")
|
||||
or memory.get("text")
|
||||
or memory.get("message")
|
||||
or str(memory) # Fallback to string representation
|
||||
)
|
||||
text = memory.get("content") or memory.get("text") or memory.get("message") or str(memory)
|
||||
texts.append(text)
|
||||
|
||||
# Encode texts
|
||||
encode_result = encode_batch(texts)
|
||||
|
||||
if not encode_result["success"]:
|
||||
return encode_result
|
||||
|
||||
# Combine embeddings with original memories
|
||||
return {
|
||||
"success": True,
|
||||
"embeddings": encode_result["embeddings"],
|
||||
@@ -246,12 +130,7 @@ def encode_memories(memories: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
|
||||
|
||||
def get_model_info() -> Dict[str, Any]:
|
||||
"""
|
||||
Get embedding model information
|
||||
|
||||
Returns:
|
||||
Dict with model metadata
|
||||
"""
|
||||
"""Get embedding model information."""
|
||||
try:
|
||||
service = _get_service()
|
||||
|
||||
@@ -259,8 +138,6 @@ def get_model_info() -> Dict[str, Any]:
|
||||
"success": True,
|
||||
"model_name": service.model_name,
|
||||
"dimension": service.dimension,
|
||||
"batch_size": service.batch_size,
|
||||
"gpu_enabled": service.use_gpu,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"[embedder] Failed to get model info: {e}")
|
||||
|
||||
@@ -336,14 +336,14 @@ class TestProcessFileToVectors:
|
||||
mock_chromadb = MagicMock()
|
||||
mock_chromadb.PersistentClient.return_value = mock_client
|
||||
|
||||
# Mock sentence_transformers
|
||||
# Mock fastembed
|
||||
mock_model = MagicMock()
|
||||
mock_model.encode.return_value = MagicMock(tolist=MagicMock(return_value=[[0.1, 0.2]]))
|
||||
mock_st = MagicMock()
|
||||
mock_st.SentenceTransformer.return_value = mock_model
|
||||
mock_model.embed.return_value = iter([MagicMock(tolist=MagicMock(return_value=[0.1, 0.2]))])
|
||||
mock_fastembed = MagicMock()
|
||||
mock_fastembed.TextEmbedding.return_value = mock_model
|
||||
|
||||
monkeypatch.setitem(sys.modules, "chromadb", mock_chromadb)
|
||||
monkeypatch.setitem(sys.modules, "sentence_transformers", mock_st)
|
||||
monkeypatch.setitem(sys.modules, "fastembed", mock_fastembed)
|
||||
|
||||
result = mod.process_file_to_vectors(test_file, "test_collection")
|
||||
|
||||
@@ -368,7 +368,7 @@ class TestProcessFileToVectors:
|
||||
|
||||
# Remove chromadb from modules so the import inside the function fails
|
||||
monkeypatch.delitem(sys.modules, "chromadb", raising=False)
|
||||
monkeypatch.delitem(sys.modules, "sentence_transformers", raising=False)
|
||||
monkeypatch.delitem(sys.modules, "fastembed", raising=False)
|
||||
|
||||
# Patch the builtins __import__ to raise for chromadb
|
||||
original_import = __builtins__.__import__ if hasattr(__builtins__, "__import__") else __import__
|
||||
|
||||
@@ -87,26 +87,22 @@ def _prepare_vector_search_mocks(monkeypatch):
|
||||
MagicMock(),
|
||||
)
|
||||
|
||||
# Mock sentence_transformers and torch for QueryEncoder
|
||||
# Mock fastembed for QueryEncoder
|
||||
mock_model = MagicMock()
|
||||
mock_model.encode.return_value = MagicMock(tolist=MagicMock(return_value=[0.1] * 384))
|
||||
mock_model.to.return_value = mock_model
|
||||
mock_model.embed.side_effect = lambda texts: iter(
|
||||
[MagicMock(tolist=MagicMock(return_value=[0.1] * 384)) for _ in texts]
|
||||
)
|
||||
|
||||
mock_st_cls = MagicMock(return_value=mock_model)
|
||||
mock_te_cls = MagicMock(return_value=mock_model)
|
||||
|
||||
mock_sentence_transformers = MagicMock()
|
||||
mock_sentence_transformers.SentenceTransformer = mock_st_cls
|
||||
monkeypatch.setitem(sys.modules, "sentence_transformers", mock_sentence_transformers)
|
||||
|
||||
mock_torch = MagicMock()
|
||||
mock_torch.cuda.is_available.return_value = False
|
||||
monkeypatch.setitem(sys.modules, "torch", mock_torch)
|
||||
mock_fastembed = MagicMock()
|
||||
mock_fastembed.TextEmbedding = mock_te_cls
|
||||
monkeypatch.setitem(sys.modules, "fastembed", mock_fastembed)
|
||||
|
||||
return {
|
||||
"client": mock_client,
|
||||
"model": mock_model,
|
||||
"st_cls": mock_st_cls,
|
||||
"torch": mock_torch,
|
||||
"st_cls": mock_te_cls,
|
||||
}
|
||||
|
||||
|
||||
@@ -443,7 +439,7 @@ class TestEncodeQuery:
|
||||
assert result["success"] is True
|
||||
assert len(result["embedding"]) == 384
|
||||
assert result["dimension"] == 384
|
||||
assert result["model"] == "all-MiniLM-L6-v2"
|
||||
assert result["model"] == "sentence-transformers/all-MiniLM-L6-v2"
|
||||
|
||||
def test_encode_rejects_empty_query(self, monkeypatch):
|
||||
"""Empty query string returns error."""
|
||||
@@ -467,7 +463,7 @@ class TestEncodeQuery:
|
||||
"""If model.encode raises, error is caught and returned."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["model"].encode.side_effect = RuntimeError("CUDA out of memory")
|
||||
mocks["model"].embed.side_effect = RuntimeError("CUDA out of memory")
|
||||
|
||||
result = mod.encode_query("test query")
|
||||
|
||||
@@ -481,7 +477,7 @@ class TestEncodeQuery:
|
||||
mod.encode_query("first query")
|
||||
mod.encode_query("second query")
|
||||
|
||||
# SentenceTransformer should only be constructed once
|
||||
# TextEmbedding should only be constructed once
|
||||
mocks["st_cls"].assert_called_once()
|
||||
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
# META DATA HEADER
|
||||
# Name: tests/test_vector.py
|
||||
# Date: 2026-04-03
|
||||
# Version: 1.0.0
|
||||
# Version: 2.0.0
|
||||
# Category: memory/tests
|
||||
# =============================================
|
||||
|
||||
@@ -10,12 +10,12 @@
|
||||
|
||||
Covers:
|
||||
- vector/embedder.py EmbeddingService class (init, encode_batch with
|
||||
pre-sort by length and order restoration, GPU cleanup path)
|
||||
pre-sort by length and order restoration)
|
||||
- vector/embedder.py Public API functions (encode_batch, encode_memories,
|
||||
get_model_info)
|
||||
- vector/embedder.py Singleton management (_get_service, global reset)
|
||||
|
||||
All tests use mocks/tmp_path -- no live sentence-transformers, torch, or GPU access.
|
||||
All tests use mocks/tmp_path -- no live fastembed or ONNX access.
|
||||
"""
|
||||
|
||||
import sys
|
||||
@@ -29,7 +29,7 @@ pytest.importorskip("chromadb")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Import helper -- torch and sentence_transformers must be mocked
|
||||
# Import helper -- fastembed must be mocked
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@@ -39,14 +39,10 @@ def _import_embedder(monkeypatch):
|
||||
Returns:
|
||||
Tuple of (embedder module, dict of mock objects)
|
||||
"""
|
||||
mock_torch = MagicMock()
|
||||
mock_torch.cuda.is_available.return_value = False
|
||||
monkeypatch.setitem(sys.modules, "torch", mock_torch)
|
||||
|
||||
mock_st = MagicMock()
|
||||
mock_fastembed = MagicMock()
|
||||
mock_model = MagicMock()
|
||||
mock_st.SentenceTransformer.return_value = mock_model
|
||||
monkeypatch.setitem(sys.modules, "sentence_transformers", mock_st)
|
||||
mock_fastembed.TextEmbedding.return_value = mock_model
|
||||
monkeypatch.setitem(sys.modules, "fastembed", mock_fastembed)
|
||||
|
||||
# Clear cached module for fresh import
|
||||
sys.modules.pop("aipass.memory.apps.handlers.vector.embedder", None)
|
||||
@@ -57,8 +53,7 @@ def _import_embedder(monkeypatch):
|
||||
from aipass.memory.apps.handlers.vector import embedder
|
||||
|
||||
return embedder, {
|
||||
"torch": mock_torch,
|
||||
"st": mock_st,
|
||||
"fastembed": mock_fastembed,
|
||||
"model": mock_model,
|
||||
}
|
||||
|
||||
@@ -90,8 +85,8 @@ class TestPublicEncodeBatch:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
fake_embeddings = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]])
|
||||
mocks["model"].encode.return_value = fake_embeddings
|
||||
fake_embeddings = [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])]
|
||||
mocks["model"].embed.return_value = iter(fake_embeddings)
|
||||
|
||||
result = embedder.encode_batch(["hello world", "test text"])
|
||||
|
||||
@@ -103,7 +98,7 @@ class TestPublicEncodeBatch:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["model"].encode.side_effect = RuntimeError("CUDA out of memory")
|
||||
mocks["model"].embed.side_effect = RuntimeError("ONNX runtime error")
|
||||
|
||||
result = embedder.encode_batch(["some text"])
|
||||
|
||||
@@ -114,7 +109,7 @@ class TestPublicEncodeBatch:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["st"].SentenceTransformer.side_effect = RuntimeError("Model not found")
|
||||
mocks["fastembed"].TextEmbedding.side_effect = RuntimeError("Model not found")
|
||||
|
||||
result = embedder.encode_batch(["some text"])
|
||||
|
||||
@@ -145,8 +140,8 @@ class TestPublicEncodeMemories:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
fake_embeddings = np.array([[0.1, 0.2]])
|
||||
mocks["model"].encode.return_value = fake_embeddings
|
||||
fake_embeddings = [np.array([0.1, 0.2])]
|
||||
mocks["model"].embed.return_value = iter(fake_embeddings)
|
||||
|
||||
memories = [{"content": "Important observation", "timestamp": "2026-01-01"}]
|
||||
result = embedder.encode_memories(memories)
|
||||
@@ -154,8 +149,7 @@ class TestPublicEncodeMemories:
|
||||
assert result["success"] is True
|
||||
assert result["count"] == 1
|
||||
assert result["memories"] == memories
|
||||
# Verify the model was called with extracted text
|
||||
call_args = mocks["model"].encode.call_args
|
||||
call_args = mocks["model"].embed.call_args
|
||||
texts_passed = call_args[0][0]
|
||||
assert texts_passed == ["Important observation"]
|
||||
|
||||
@@ -163,15 +157,15 @@ class TestPublicEncodeMemories:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
fake_embeddings = np.array([[0.1, 0.2]])
|
||||
mocks["model"].encode.return_value = fake_embeddings
|
||||
fake_embeddings = [np.array([0.1, 0.2])]
|
||||
mocks["model"].embed.return_value = iter(fake_embeddings)
|
||||
|
||||
memories = [{"text": "Session summary", "date": "2026-02-01"}]
|
||||
result = embedder.encode_memories(memories)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["count"] == 1
|
||||
call_args = mocks["model"].encode.call_args
|
||||
call_args = mocks["model"].embed.call_args
|
||||
texts_passed = call_args[0][0]
|
||||
assert texts_passed == ["Session summary"]
|
||||
|
||||
@@ -179,16 +173,15 @@ class TestPublicEncodeMemories:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
fake_embeddings = np.array([[0.1, 0.2]])
|
||||
mocks["model"].encode.return_value = fake_embeddings
|
||||
fake_embeddings = [np.array([0.1, 0.2])]
|
||||
mocks["model"].embed.return_value = iter(fake_embeddings)
|
||||
|
||||
memories = [{"arbitrary_key": "value123", "number": 42}]
|
||||
result = embedder.encode_memories(memories)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["count"] == 1
|
||||
# The fallback is str(memory) which includes the full dict repr
|
||||
call_args = mocks["model"].encode.call_args
|
||||
call_args = mocks["model"].embed.call_args
|
||||
texts_passed = call_args[0][0]
|
||||
assert "arbitrary_key" in texts_passed[0]
|
||||
|
||||
@@ -196,7 +189,7 @@ class TestPublicEncodeMemories:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["model"].encode.side_effect = RuntimeError("Encoding crashed")
|
||||
mocks["model"].embed.side_effect = RuntimeError("Encoding crashed")
|
||||
|
||||
memories = [{"content": "test memory"}]
|
||||
result = embedder.encode_memories(memories)
|
||||
@@ -207,8 +200,8 @@ class TestPublicEncodeMemories:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
fake_embeddings = np.array([[0.1], [0.2], [0.3]])
|
||||
mocks["model"].encode.return_value = fake_embeddings
|
||||
fake_embeddings = [np.array([0.1]), np.array([0.2]), np.array([0.3])]
|
||||
mocks["model"].embed.return_value = iter(fake_embeddings)
|
||||
|
||||
memories = [
|
||||
{"content": "first"},
|
||||
@@ -221,11 +214,10 @@ class TestPublicEncodeMemories:
|
||||
assert result["count"] == 3
|
||||
assert result["memories"] is memories
|
||||
|
||||
call_args = mocks["model"].encode.call_args
|
||||
call_args = mocks["model"].embed.call_args
|
||||
texts_passed = call_args[0][0]
|
||||
assert texts_passed[0] == "first"
|
||||
assert texts_passed[1] == "second"
|
||||
# Third falls back to str()
|
||||
assert "third" in texts_passed[2]
|
||||
|
||||
|
||||
@@ -244,16 +236,14 @@ class TestPublicGetModelInfo:
|
||||
result = embedder.get_model_info()
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["model_name"] == "all-MiniLM-L6-v2"
|
||||
assert result["model_name"] == "sentence-transformers/all-MiniLM-L6-v2"
|
||||
assert result["dimension"] == 384
|
||||
assert result["batch_size"] == 16 # CPU batch size (GPU is mocked off)
|
||||
assert result["gpu_enabled"] is False
|
||||
|
||||
def test_service_init_failure_returns_error(self, monkeypatch):
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["st"].SentenceTransformer.side_effect = ImportError("no model")
|
||||
mocks["fastembed"].TextEmbedding.side_effect = ImportError("no model")
|
||||
|
||||
result = embedder.get_model_info()
|
||||
|
||||
@@ -273,28 +263,21 @@ class TestEmbeddingServiceEncodeBatch:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
# Track what texts the model receives (should be sorted by length)
|
||||
received_texts: list[Any] = []
|
||||
|
||||
def fake_encode(texts, **kwargs):
|
||||
def fake_embed(texts):
|
||||
received_texts.append(list(texts))
|
||||
# Return embeddings matching the sorted input length
|
||||
return np.array([[float(i)] * 3 for i in range(len(texts))])
|
||||
return iter([np.array([float(i)] * 3) for i in range(len(texts))])
|
||||
|
||||
mocks["model"].encode.side_effect = fake_encode
|
||||
mocks["model"].embed.side_effect = fake_embed
|
||||
|
||||
service = embedder.EmbeddingService()
|
||||
texts = ["long text here", "ab", "medium text"]
|
||||
result = service.encode_batch(texts)
|
||||
|
||||
# Model should receive texts sorted by length
|
||||
assert received_texts[0] == ["ab", "medium text", "long text here"]
|
||||
|
||||
# But returned embeddings should be in original order
|
||||
assert result["count"] == 3
|
||||
# Index 0 was "long text here" (sorted position 2) -> embedding [2,2,2]
|
||||
# Index 1 was "ab" (sorted position 0) -> embedding [0,0,0]
|
||||
# Index 2 was "medium text" (sorted position 1) -> embedding [1,1,1]
|
||||
embs = result["embeddings"]
|
||||
np.testing.assert_array_equal(embs[0], [2.0, 2.0, 2.0])
|
||||
np.testing.assert_array_equal(embs[1], [0.0, 0.0, 0.0])
|
||||
@@ -315,54 +298,10 @@ class TestEmbeddingServiceEncodeBatch:
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["model"].encode.return_value = np.array([[0.5, 0.6]])
|
||||
mocks["model"].embed.return_value = iter([np.array([0.5, 0.6])])
|
||||
|
||||
service = embedder.EmbeddingService()
|
||||
result = service.encode_batch(["only one"])
|
||||
|
||||
assert result["count"] == 1
|
||||
np.testing.assert_array_equal(result["embeddings"][0], [0.5, 0.6])
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Tests: EmbeddingService -- GPU path
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestEmbeddingServiceGPU:
|
||||
"""Test EmbeddingService GPU detection and cleanup."""
|
||||
|
||||
def test_gpu_enabled_sets_larger_batch_size(self, monkeypatch):
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["torch"].cuda.is_available.return_value = True
|
||||
|
||||
service = embedder.EmbeddingService()
|
||||
|
||||
assert service.use_gpu is True
|
||||
assert service.batch_size == 64
|
||||
|
||||
def test_gpu_cache_cleared_after_encode(self, monkeypatch):
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["torch"].cuda.is_available.return_value = True
|
||||
mocks["model"].encode.return_value = np.array([[0.1, 0.2]])
|
||||
|
||||
service = embedder.EmbeddingService()
|
||||
service.encode_batch(["test text"])
|
||||
|
||||
mocks["torch"].cuda.empty_cache.assert_called_once()
|
||||
|
||||
def test_cpu_does_not_clear_gpu_cache(self, monkeypatch):
|
||||
embedder, mocks = _import_embedder(monkeypatch)
|
||||
_reset_globals(embedder)
|
||||
|
||||
mocks["torch"].cuda.is_available.return_value = False
|
||||
mocks["model"].encode.return_value = np.array([[0.1, 0.2]])
|
||||
|
||||
service = embedder.EmbeddingService()
|
||||
service.encode_batch(["test text"])
|
||||
|
||||
mocks["torch"].cuda.empty_cache.assert_not_called()
|
||||
|
||||
Reference in New Issue
Block a user