mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-30 01:52:29 +00:00
up
This commit is contained in:
parent
6ad02f0a8f
commit
9f90c448f1
4 changed files with 146 additions and 128 deletions
|
|
@ -134,6 +134,11 @@ class BaseComponent(ABC):
|
|||
return Path.cwd() / "metadata"
|
||||
return self.vault_path / self.app_context.app_config.metadata_dir
|
||||
|
||||
@property
|
||||
def component_metadata_path(self) -> Path:
|
||||
"""Resolved component metadata directory: vault_metadata_path / component_type."""
|
||||
return self.vault_metadata_path / self.component_type.value
|
||||
|
||||
def to_vault_relative(self, path: str | Path) -> str:
|
||||
"""Return path relative to vault_path; absolute path string if outside."""
|
||||
abs_path = Path(path).absolute()
|
||||
|
|
|
|||
|
|
@ -13,9 +13,11 @@ from ..base_component import BaseComponent
|
|||
from ...enumeration import ComponentEnum
|
||||
from ...schema import EmbNode
|
||||
|
||||
Miss = tuple[int, str, str] # (result_index, text, cache_key)
|
||||
|
||||
|
||||
class BaseEmbeddingModel(BaseComponent):
|
||||
"""Embedding model with LRU cache and disk persistence."""
|
||||
"""Embedding model with LRU cache, disk persistence, and concurrent batching."""
|
||||
|
||||
component_type = ComponentEnum.EMBEDDING_MODEL
|
||||
|
||||
|
|
@ -29,6 +31,7 @@ class BaseEmbeddingModel(BaseComponent):
|
|||
max_batch_size: int = 10,
|
||||
max_input_length: int = 8192,
|
||||
max_cache_size: int = 10000,
|
||||
max_concurrency: int = 2,
|
||||
enable_cache: bool = True,
|
||||
cache_version: str = "v1",
|
||||
max_retries: int = 3,
|
||||
|
|
@ -43,21 +46,24 @@ class BaseEmbeddingModel(BaseComponent):
|
|||
self.max_batch_size = max_batch_size
|
||||
self.max_input_length = max_input_length
|
||||
self.max_cache_size = max_cache_size
|
||||
self.max_concurrency = max_concurrency
|
||||
self.enable_cache = enable_cache
|
||||
self.cache_version = cache_version
|
||||
self.max_retries = max_retries
|
||||
self._embedding_cache: OrderedDict[str, np.ndarray] = OrderedDict()
|
||||
self._cache: OrderedDict[str, np.ndarray] = OrderedDict()
|
||||
self._key_suffix = f"|{model_name}|{dimensions}".encode()
|
||||
self.is_healthy: bool = True
|
||||
|
||||
@property
|
||||
def cache_path(self) -> Path:
|
||||
"""Disk path for the embedding cache file."""
|
||||
return self.vault_metadata_path / "embedding_cache" / f"{self.name}_{self.cache_version}.npz"
|
||||
|
||||
async def _start(self) -> None:
|
||||
"""Load cache from disk on startup."""
|
||||
await self.load()
|
||||
|
||||
async def _close(self) -> None:
|
||||
await self.dump()
|
||||
|
||||
async def health_check(self, timeout: float = 2.0) -> bool:
|
||||
"""Probe the provider; sets and returns is_healthy."""
|
||||
tag = f"[EMBEDDING HEALTH CHECK] name={self.name} model={self.model_name}"
|
||||
|
|
@ -75,71 +81,21 @@ class BaseEmbeddingModel(BaseComponent):
|
|||
self.logger.error(f"{tag} -> FAIL {type(e).__name__}: {e}")
|
||||
return self.is_healthy
|
||||
|
||||
async def _close(self) -> None:
|
||||
"""Persist cache to disk on shutdown."""
|
||||
await self.dump()
|
||||
|
||||
# -- Public API --
|
||||
|
||||
async def get_embedding(self, input_text: str, **kwargs) -> np.ndarray | None:
|
||||
"""Get embedding for a single text."""
|
||||
results = await self.get_embeddings([input_text], **kwargs)
|
||||
return results[0] if results else None
|
||||
|
||||
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[np.ndarray | None]:
|
||||
"""Get embeddings for a list of texts, with caching and batching."""
|
||||
truncated = [t[: self.max_input_length] for t in input_text]
|
||||
results: list[np.ndarray | None] = [None] * len(truncated)
|
||||
to_compute: list[tuple[int, str]] = []
|
||||
|
||||
# Split into cache hits and misses
|
||||
for idx, text in enumerate(truncated):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
to_compute.append((idx, text))
|
||||
|
||||
# Batch-compute misses with retry
|
||||
if to_compute:
|
||||
for i in range(0, len(to_compute), self.max_batch_size):
|
||||
batch = to_compute[i : i + self.max_batch_size]
|
||||
indices = [idx for idx, _ in batch]
|
||||
texts = [text for _, text in batch]
|
||||
|
||||
embeddings = None
|
||||
for attempt in range(self.max_retries):
|
||||
try:
|
||||
embeddings = await self._get_embeddings(texts, **kwargs)
|
||||
if embeddings and len(embeddings) == len(texts):
|
||||
break
|
||||
except (TimeoutError, ConnectionError, OSError):
|
||||
if attempt < self.max_retries - 1:
|
||||
await asyncio.sleep(2**attempt)
|
||||
except Exception:
|
||||
self.logger.exception("Embedding request failed")
|
||||
break
|
||||
|
||||
if not embeddings or len(embeddings) != len(texts):
|
||||
continue
|
||||
|
||||
# Normalize dimensions and cache
|
||||
for orig_idx, text, emb in zip(indices, texts, embeddings):
|
||||
if emb is None:
|
||||
continue
|
||||
emb_array = np.asarray(emb, dtype=np.float16)
|
||||
if len(emb_array) != self.dimensions:
|
||||
if len(emb_array) < self.dimensions:
|
||||
emb_array = np.pad(emb_array, (0, self.dimensions - len(emb_array)))
|
||||
else:
|
||||
emb_array = emb_array[: self.dimensions]
|
||||
results[orig_idx] = emb_array
|
||||
self._put_to_cache(text, emb_array)
|
||||
|
||||
"""Get embeddings for texts. Cache hits return immediately; misses run concurrently."""
|
||||
texts = [self._truncate(t) for t in input_text]
|
||||
results, misses = self._partition_by_cache(texts)
|
||||
if misses:
|
||||
await self._fill_misses(misses, results, **kwargs)
|
||||
return results
|
||||
|
||||
async def get_node_embeddings(self, nodes: list[EmbNode], **kwargs) -> list[EmbNode]:
|
||||
"""Compute and assign embeddings for EmbNode objects."""
|
||||
embeddings = await self.get_embeddings([n.text for n in nodes], **kwargs)
|
||||
if len(embeddings) == len(nodes):
|
||||
for node, vec in zip(nodes, embeddings):
|
||||
|
|
@ -151,64 +107,134 @@ class BaseEmbeddingModel(BaseComponent):
|
|||
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]:
|
||||
"""Get raw embeddings from the underlying provider."""
|
||||
|
||||
# -- Cache Operations --
|
||||
# -- Batching --
|
||||
|
||||
def _get_from_cache(self, text: str) -> np.ndarray | None:
|
||||
"""Lookup text in LRU cache, promoting on hit."""
|
||||
if not self.enable_cache:
|
||||
return None
|
||||
key = self._get_cache_key(text)
|
||||
if key not in self._embedding_cache:
|
||||
return None
|
||||
self._embedding_cache.move_to_end(key)
|
||||
return self._embedding_cache[key]
|
||||
def _truncate(self, text: str) -> str:
|
||||
return text if len(text) <= self.max_input_length else text[: self.max_input_length]
|
||||
|
||||
def _put_to_cache(self, text: str, embedding: np.ndarray) -> None:
|
||||
"""Insert into LRU cache, evicting oldest if full."""
|
||||
def _partition_by_cache(self, texts: list[str]) -> tuple[list[np.ndarray | None], list[Miss]]:
|
||||
"""Split texts into pre-filled results (hits) and a miss list to compute."""
|
||||
results: list[np.ndarray | None] = [None] * len(texts)
|
||||
misses: list[Miss] = []
|
||||
for idx, text in enumerate(texts):
|
||||
key = self._cache_key(text)
|
||||
hit = self._cache_get(key)
|
||||
if hit is not None:
|
||||
results[idx] = hit
|
||||
else:
|
||||
misses.append((idx, text, key))
|
||||
return results, misses
|
||||
|
||||
async def _fill_misses(self, misses: list[Miss], results: list[np.ndarray | None], **kwargs) -> None:
|
||||
"""Compute miss embeddings in concurrent batches and write into results + cache."""
|
||||
size = self.max_batch_size
|
||||
batches = [misses[i : i + size] for i in range(0, len(misses), size)]
|
||||
sem = asyncio.Semaphore(self.max_concurrency)
|
||||
|
||||
async def run(batch: list[Miss]) -> list[tuple[int, str, np.ndarray]]:
|
||||
async with sem:
|
||||
return await self._compute_batch(batch, **kwargs)
|
||||
|
||||
for done in await asyncio.gather(*(run(b) for b in batches)):
|
||||
for idx, key, emb in done:
|
||||
results[idx] = emb
|
||||
self._cache_put(key, emb)
|
||||
|
||||
async def _compute_batch(self, batch: list[Miss], **kwargs) -> list[tuple[int, str, np.ndarray]]:
|
||||
"""Call provider for one batch with retry; returns [(idx, key, embedding)]."""
|
||||
texts = [text for _, text, _ in batch]
|
||||
embeddings = await self._call_with_retry(texts, **kwargs)
|
||||
if not embeddings or len(embeddings) != len(texts):
|
||||
return []
|
||||
out: list[tuple[int, str, np.ndarray]] = []
|
||||
for (idx, _text, key), raw in zip(batch, embeddings):
|
||||
if raw is None:
|
||||
continue
|
||||
emb = self._normalize_dim(np.asarray(raw, dtype=np.float16))
|
||||
out.append((idx, key, emb))
|
||||
return out
|
||||
|
||||
async def _call_with_retry(self, texts: list[str], **kwargs) -> list[list[float] | None] | None:
|
||||
"""Call provider with exponential backoff on transient errors."""
|
||||
for attempt in range(self.max_retries):
|
||||
try:
|
||||
result = await self._get_embeddings(texts, **kwargs)
|
||||
if result and len(result) == len(texts):
|
||||
return result
|
||||
except (TimeoutError, ConnectionError, OSError):
|
||||
if attempt < self.max_retries - 1:
|
||||
await asyncio.sleep(2**attempt)
|
||||
except Exception:
|
||||
self.logger.exception("Embedding request failed")
|
||||
return None
|
||||
return None
|
||||
|
||||
def _normalize_dim(self, emb: np.ndarray) -> np.ndarray:
|
||||
if len(emb) == self.dimensions:
|
||||
return emb
|
||||
if len(emb) < self.dimensions:
|
||||
return np.pad(emb, (0, self.dimensions - len(emb)))
|
||||
return emb[: self.dimensions]
|
||||
|
||||
# -- Cache --
|
||||
|
||||
def _cache_key(self, text: str) -> str:
|
||||
return hashlib.sha256(text.encode() + self._key_suffix).hexdigest()
|
||||
|
||||
def _cache_get(self, key: str) -> np.ndarray | None:
|
||||
if not self.enable_cache or key not in self._cache:
|
||||
return None
|
||||
self._cache.move_to_end(key)
|
||||
return self._cache[key]
|
||||
|
||||
def _cache_put(self, key: str, embedding: np.ndarray) -> None:
|
||||
if not self.enable_cache or self.max_cache_size <= 0 or len(embedding) != self.dimensions:
|
||||
return
|
||||
key = self._get_cache_key(text)
|
||||
if len(self._embedding_cache) >= self.max_cache_size and key not in self._embedding_cache:
|
||||
self._embedding_cache.popitem(last=False)
|
||||
self._embedding_cache[key] = embedding
|
||||
self._embedding_cache.move_to_end(key)
|
||||
cache = self._cache
|
||||
if key in cache:
|
||||
cache.move_to_end(key)
|
||||
cache[key] = embedding
|
||||
return
|
||||
if len(cache) >= self.max_cache_size:
|
||||
cache.popitem(last=False)
|
||||
cache[key] = embedding
|
||||
|
||||
def _get_cache_key(self, text: str) -> str:
|
||||
"""Generate cache key from text, model name, and dimensions."""
|
||||
return hashlib.sha256(f"{text}|{self.model_name}|{self.dimensions}".encode()).hexdigest()
|
||||
|
||||
# -- Cache Persistence --
|
||||
# -- Persistence --
|
||||
|
||||
async def load(self) -> None:
|
||||
"""Load cached embeddings from disk (npz format); replaces in-memory cache."""
|
||||
self._embedding_cache.clear()
|
||||
"""Load cached embeddings from disk (npz); replaces in-memory cache."""
|
||||
self._cache.clear()
|
||||
if not self.enable_cache or not self.cache_path.exists():
|
||||
return
|
||||
await asyncio.to_thread(self._load_sync)
|
||||
|
||||
def _load_sync(self) -> None:
|
||||
try:
|
||||
data = np.load(self.cache_path)
|
||||
except Exception:
|
||||
self.logger.exception("Failed to load embedding cache, removing")
|
||||
self.cache_path.unlink(missing_ok=True)
|
||||
return
|
||||
|
||||
for key, emb in zip(data["keys"], data["embeddings"]):
|
||||
if len(emb) != self.dimensions:
|
||||
continue
|
||||
if len(self._embedding_cache) >= self.max_cache_size:
|
||||
if len(self._cache) >= self.max_cache_size:
|
||||
break
|
||||
self._embedding_cache[str(key)] = emb.astype(np.float16)
|
||||
self.logger.info(f"Loaded {len(self._embedding_cache)} embeddings from {self.cache_path}")
|
||||
self._cache[str(key)] = emb.astype(np.float16)
|
||||
self.logger.info(f"Loaded {len(self._cache)} embeddings from {self.cache_path}")
|
||||
|
||||
async def dump(self) -> None:
|
||||
"""Persist in-memory cache to disk (npz format)."""
|
||||
if not self.enable_cache or not self._embedding_cache:
|
||||
"""Persist in-memory cache to disk (npz)."""
|
||||
if not self.enable_cache or not self._cache:
|
||||
return
|
||||
await asyncio.to_thread(self._dump_sync)
|
||||
|
||||
def _dump_sync(self) -> None:
|
||||
self.cache_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
keys = list(self._embedding_cache.keys())
|
||||
embeddings = np.stack(list(self._embedding_cache.values()))
|
||||
keys = np.array(list(self._cache.keys()), dtype=str)
|
||||
embeddings = np.stack(list(self._cache.values()))
|
||||
try:
|
||||
np.savez(self.cache_path, keys=np.array(keys, dtype=str), embeddings=embeddings)
|
||||
self.logger.info(f"Saved {len(self._embedding_cache)} embeddings to {self.cache_path}")
|
||||
np.savez(self.cache_path, keys=keys, embeddings=embeddings)
|
||||
self.logger.info(f"Saved {len(self._cache)} embeddings to {self.cache_path}")
|
||||
except Exception:
|
||||
self.logger.exception("Failed to save embedding cache")
|
||||
|
|
|
|||
|
|
@ -1,7 +1,4 @@
|
|||
"""Abstract base for file-catalog backends."""
|
||||
|
||||
from abc import abstractmethod
|
||||
from pathlib import Path
|
||||
|
||||
from ..base_component import BaseComponent
|
||||
from ...enumeration import ComponentEnum
|
||||
|
|
@ -9,22 +6,10 @@ from ...schema import FileNode
|
|||
|
||||
|
||||
class BaseFileCatalog(BaseComponent):
|
||||
"""Abstract base for file-catalog backends.
|
||||
|
||||
A catalog records FileNode entries keyed by path — a lightweight
|
||||
counterpart to FileStore that drops chunk/embedding/keyword/link
|
||||
machinery and exposes only node upsert / delete / lookup.
|
||||
"""
|
||||
"""File-catalog backend recording FileNode entries keyed by path."""
|
||||
|
||||
component_type = ComponentEnum.FILE_CATALOG
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.catalog_path: Path = self.vault_metadata_path / self.component_type.value
|
||||
self.catalog_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# -- Lifecycle ---------------------------------------------------------
|
||||
|
||||
async def _start(self) -> None:
|
||||
await super()._start()
|
||||
await self.load()
|
||||
|
|
@ -34,12 +19,10 @@ class BaseFileCatalog(BaseComponent):
|
|||
await super()._close()
|
||||
|
||||
async def load(self) -> None:
|
||||
"""Load persisted state. No-op for backends without local files."""
|
||||
"""Load persisted state. No-op without local files."""
|
||||
|
||||
async def dump(self) -> None:
|
||||
"""Persist state. No-op for backends without local files."""
|
||||
|
||||
# -- CRUD --------------------------------------------------------------
|
||||
"""Persist state. No-op without local files."""
|
||||
|
||||
@abstractmethod
|
||||
async def upsert(self, nodes: list[FileNode]) -> None:
|
||||
|
|
@ -51,4 +34,4 @@ class BaseFileCatalog(BaseComponent):
|
|||
|
||||
@abstractmethod
|
||||
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
|
||||
"""Return nodes by paths; None = all nodes; [] = []; missing paths are skipped."""
|
||||
"""Return nodes by paths; None = all; missing paths are skipped."""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,3 @@
|
|||
"""In-memory file catalog with JSONL persistence."""
|
||||
|
||||
import aiofiles
|
||||
|
||||
from .base_file_catalog import BaseFileCatalog
|
||||
|
|
@ -9,34 +7,27 @@ from ...schema import FileNode
|
|||
|
||||
@R.register("local")
|
||||
class LocalFileCatalog(BaseFileCatalog):
|
||||
"""Dict-backed file catalog persisted as JSONL on close."""
|
||||
"""Dict-backed catalog persisted as JSONL."""
|
||||
|
||||
def __init__(self, encoding: str = "utf-8", **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.encoding = encoding
|
||||
self._nodes: dict[str, FileNode] = {}
|
||||
self._catalog_file = self.catalog_path / f"{self.name}.jsonl"
|
||||
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
|
||||
self._catalog_file = self.component_metadata_path / f"{self.name}.jsonl"
|
||||
|
||||
async def load(self) -> None:
|
||||
if not self._catalog_file.exists():
|
||||
return
|
||||
try:
|
||||
async with aiofiles.open(self._catalog_file, encoding=self.encoding) as f:
|
||||
async for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
node = FileNode.model_validate_json(line)
|
||||
self._nodes[node.path] = node
|
||||
await self._read_jsonl()
|
||||
self.logger.info(f"Loaded {len(self._nodes)} nodes from {self._catalog_file}")
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to load {self._catalog_file}: {e}")
|
||||
|
||||
async def dump(self) -> None:
|
||||
try:
|
||||
tmp = self._catalog_file.with_suffix(".tmp")
|
||||
async with aiofiles.open(tmp, "w", encoding=self.encoding) as f:
|
||||
await f.write("\n".join(n.model_dump_json() for n in self._nodes.values()))
|
||||
tmp.replace(self._catalog_file)
|
||||
await self._write_jsonl()
|
||||
self.logger.info(f"Saved {len(self._nodes)} nodes to {self._catalog_file}")
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to write {self._catalog_file}: {e}")
|
||||
|
|
@ -54,3 +45,16 @@ class LocalFileCatalog(BaseFileCatalog):
|
|||
if paths is None:
|
||||
return list(self._nodes.values())
|
||||
return [self._nodes[p] for p in paths if p in self._nodes]
|
||||
|
||||
async def _read_jsonl(self) -> None:
|
||||
async with aiofiles.open(self._catalog_file, encoding=self.encoding) as f:
|
||||
async for line in f:
|
||||
if stripped := line.strip():
|
||||
node = FileNode.model_validate_json(stripped)
|
||||
self._nodes[node.path] = node
|
||||
|
||||
async def _write_jsonl(self) -> None:
|
||||
tmp = self._catalog_file.with_suffix(".tmp")
|
||||
async with aiofiles.open(tmp, "w", encoding=self.encoding) as f:
|
||||
await f.write("\n".join(n.model_dump_json() for n in self._nodes.values()))
|
||||
tmp.replace(self._catalog_file)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue