mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-29 01:41:38 +00:00
up
This commit is contained in:
parent
964137b474
commit
aef40098bf
5 changed files with 142 additions and 122 deletions
|
|
@ -10,27 +10,33 @@ from ...schema import FileChunk, FileLink, FileNode
|
|||
class BaseFileStore(BaseComponent):
|
||||
"""Abstract base for file store backends.
|
||||
|
||||
Defines the *semantic* contract a file store must offer: write (upsert / delete / clear),
|
||||
retrieve (vector / keyword), and graph queries (nodes / links). Sub-component composition
|
||||
(embedding model, keyword index, file graph) is each backend's implementation choice and
|
||||
is not part of the base contract.
|
||||
Defines the *semantic* contract a file store must offer: write (upsert / delete /
|
||||
clear), retrieve (vector / keyword), and graph queries (nodes / links). How the
|
||||
backend composes sub-components (embedding model, keyword index, file graph) is
|
||||
an implementation detail outside this contract.
|
||||
"""
|
||||
|
||||
component_type = ComponentEnum.FILE_STORE
|
||||
|
||||
# -- CRUD ------------------------------------------------------------
|
||||
# -- CRUD -----------------------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
async def upsert(self, files: list[tuple[FileNode, list[FileChunk]]]) -> None:
|
||||
"""Upsert files and their chunks into the store."""
|
||||
"""Upsert files and their chunks; existing chunks for the same path are replaced."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete(self, path: str | list[str]) -> None:
|
||||
"""Delete files by path from the store."""
|
||||
"""Delete the given path(s) and all their chunks; unknown paths are skipped."""
|
||||
|
||||
@abstractmethod
|
||||
async def clear(self) -> None:
|
||||
"""Drop every file and chunk in the store."""
|
||||
|
||||
# -- graph queries --------------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
|
||||
"""Return file nodes; None = all nodes; missing paths are skipped."""
|
||||
"""Return file nodes; ``None`` = all; missing paths are skipped."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_outlinks(
|
||||
|
|
@ -38,7 +44,7 @@ class BaseFileStore(BaseComponent):
|
|||
path: str,
|
||||
scope: LinkScopeEnum = LinkScopeEnum.REAL,
|
||||
) -> list[FileLink]:
|
||||
"""Return outgoing links for *path*. See ``BaseFileGraph.get_outlinks`` for scope semantics."""
|
||||
"""Outgoing links for *path*; scope semantics match ``BaseFileGraph.get_outlinks``."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_inlinks(
|
||||
|
|
@ -46,18 +52,14 @@ class BaseFileStore(BaseComponent):
|
|||
path: str,
|
||||
scope: LinkScopeEnum = LinkScopeEnum.REAL,
|
||||
) -> list[FileLink]:
|
||||
"""Return incoming links for *path*. See ``BaseFileGraph.get_inlinks`` for scope semantics."""
|
||||
"""Incoming links for *path*; scope semantics match ``BaseFileGraph.get_inlinks``."""
|
||||
|
||||
@abstractmethod
|
||||
async def clear(self) -> None:
|
||||
"""Clear the store of all files and chunks."""
|
||||
|
||||
# -- Search -----------------------------------------------------------
|
||||
# -- search ---------------------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
|
||||
"""Perform vector similarity search."""
|
||||
"""Vector similarity search over chunk embeddings."""
|
||||
|
||||
@abstractmethod
|
||||
async def keyword_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
|
||||
"""Perform full-text keyword search."""
|
||||
"""Full-text keyword search over chunk text."""
|
||||
|
|
|
|||
|
|
@ -31,22 +31,26 @@ class FaissLocalFileStore(LocalFileStore):
|
|||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
try:
|
||||
import faiss
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"faiss is required for FaissLocalFileStore. Install with `pip install faiss-cpu`.",
|
||||
) from e
|
||||
self._faiss = faiss
|
||||
self._faiss = self._import_faiss()
|
||||
self.normalize = normalize
|
||||
self.max_tombstones = max_tombstones
|
||||
self.faiss_path = self.component_metadata_path / f"faiss_index_{self.name}_{self.store_version}.bin"
|
||||
self.faiss_idmap_path = self.component_metadata_path / f"faiss_idmap_{self.name}_{self.store_version}.json"
|
||||
self._faiss_index = None # faiss.Index | None
|
||||
self._id_map: list[str] = [] # row -> chunk_id
|
||||
self._id_to_row: dict[str, int] = {} # chunk_id -> row
|
||||
self._id_to_row: dict[str, int] = {} # chunk_id -> row (live entries only)
|
||||
self._tombstones: set[int] = set() # rows whose chunk_id was deleted
|
||||
|
||||
@staticmethod
|
||||
def _import_faiss():
|
||||
try:
|
||||
import faiss
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"faiss is required for FaissLocalFileStore. Install with `pip install faiss-cpu`.",
|
||||
) from e
|
||||
return faiss
|
||||
|
||||
# -- helpers ----------------------------------------------------------
|
||||
|
||||
@property
|
||||
|
|
@ -105,57 +109,64 @@ class FaissLocalFileStore(LocalFileStore):
|
|||
# -- persistence ------------------------------------------------------
|
||||
|
||||
async def load(self) -> None:
|
||||
"""Load chunks (parent), then load FAISS sidecar; rebuild from chunks on miss/corruption."""
|
||||
"""Load chunks via the parent, then attach FAISS state (sidecar or rebuild)."""
|
||||
await super().load()
|
||||
if self.embedding_model is None or self._dim == 0:
|
||||
self._faiss_index = None
|
||||
return
|
||||
|
||||
loaded = False
|
||||
if self.faiss_path.exists() and self.faiss_idmap_path.exists():
|
||||
try:
|
||||
index = self._faiss.read_index(str(self.faiss_path))
|
||||
if index.d != self._dim:
|
||||
raise ValueError(f"FAISS dim {index.d} != embedding dim {self._dim}")
|
||||
async with aiofiles.open(self.faiss_idmap_path, encoding=self.encoding) as f:
|
||||
data = json.loads(await f.read())
|
||||
id_map = list(data.get("id_map", []))
|
||||
if len(id_map) != index.ntotal:
|
||||
raise ValueError(f"id_map size {len(id_map)} != index ntotal {index.ntotal}")
|
||||
self._faiss_index = index
|
||||
self._id_map = id_map
|
||||
self._tombstones = set(data.get("tombstones", []))
|
||||
self._id_to_row = {cid: i for i, cid in enumerate(self._id_map) if i not in self._tombstones}
|
||||
self.logger.info(f"Loaded FAISS index: {index.ntotal} vectors from {self.faiss_path}")
|
||||
loaded = True
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to load FAISS index, will rebuild: {e}")
|
||||
self.faiss_path.unlink(missing_ok=True)
|
||||
self.faiss_idmap_path.unlink(missing_ok=True)
|
||||
|
||||
if not loaded:
|
||||
if not await self._try_load_sidecar():
|
||||
self._rebuild_index()
|
||||
|
||||
async def _try_load_sidecar(self) -> bool:
|
||||
"""Read the binary index plus id-map sidecar. On any mismatch or read error,
|
||||
wipe the partial files so the caller can rebuild from chunks cleanly.
|
||||
"""
|
||||
if not (self.faiss_path.exists() and self.faiss_idmap_path.exists()):
|
||||
return False
|
||||
try:
|
||||
index = self._faiss.read_index(str(self.faiss_path))
|
||||
if index.d != self._dim:
|
||||
raise ValueError(f"FAISS dim {index.d} != embedding dim {self._dim}")
|
||||
async with aiofiles.open(self.faiss_idmap_path, encoding=self.encoding) as f:
|
||||
data = json.loads(await f.read())
|
||||
id_map = list(data.get("id_map", []))
|
||||
if len(id_map) != index.ntotal:
|
||||
raise ValueError(f"id_map size {len(id_map)} != index ntotal {index.ntotal}")
|
||||
self._faiss_index = index
|
||||
self._id_map = id_map
|
||||
self._tombstones = set(data.get("tombstones", []))
|
||||
self._id_to_row = {cid: i for i, cid in enumerate(self._id_map) if i not in self._tombstones}
|
||||
self.logger.info(f"Loaded FAISS index: {index.ntotal} vectors from {self.faiss_path}")
|
||||
return True
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to load FAISS index, will rebuild: {e}")
|
||||
self.faiss_path.unlink(missing_ok=True)
|
||||
self.faiss_idmap_path.unlink(missing_ok=True)
|
||||
return False
|
||||
|
||||
async def dump(self) -> None:
|
||||
"""Persist chunks JSONL (parent) plus FAISS sidecar via atomic rename."""
|
||||
"""Persist chunks JSONL via the parent, then write the FAISS sidecar atomically."""
|
||||
await super().dump()
|
||||
if self._faiss_index is None or self.embedding_model is None:
|
||||
return
|
||||
try:
|
||||
self._compact_if_needed()
|
||||
tmp_index = self.faiss_path.with_suffix(".tmp")
|
||||
self._faiss.write_index(self._faiss_index, str(tmp_index))
|
||||
tmp_index.replace(self.faiss_path)
|
||||
|
||||
tmp_idmap = self.faiss_idmap_path.with_suffix(".tmp")
|
||||
payload = json.dumps({"id_map": self._id_map, "tombstones": sorted(self._tombstones)})
|
||||
async with aiofiles.open(tmp_idmap, "w", encoding=self.encoding) as f:
|
||||
await f.write(payload)
|
||||
tmp_idmap.replace(self.faiss_idmap_path)
|
||||
await self._write_sidecar()
|
||||
self.logger.info(f"Saved FAISS index: {self._faiss_index.ntotal} vectors to {self.faiss_path}")
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to write FAISS index: {e}")
|
||||
|
||||
async def _write_sidecar(self) -> None:
|
||||
tmp_index = self.faiss_path.with_suffix(".tmp")
|
||||
self._faiss.write_index(self._faiss_index, str(tmp_index))
|
||||
tmp_index.replace(self.faiss_path)
|
||||
|
||||
tmp_idmap = self.faiss_idmap_path.with_suffix(".tmp")
|
||||
payload = json.dumps({"id_map": self._id_map, "tombstones": sorted(self._tombstones)})
|
||||
async with aiofiles.open(tmp_idmap, "w", encoding=self.encoding) as f:
|
||||
await f.write(payload)
|
||||
tmp_idmap.replace(self.faiss_idmap_path)
|
||||
|
||||
# -- CRUD overrides ---------------------------------------------------
|
||||
|
||||
async def upsert(self, files: list[tuple[FileNode, list[FileChunk]]]) -> None:
|
||||
|
|
@ -163,23 +174,27 @@ class FaissLocalFileStore(LocalFileStore):
|
|||
return
|
||||
assert self.file_graph is not None
|
||||
|
||||
# Snapshot the chunk_ids the file_graph currently holds for these paths,
|
||||
# so we can compute add/delete deltas after super finishes.
|
||||
# Snapshot pre-upsert chunk_ids so we can diff against the post-upsert state.
|
||||
old_ids_by_path = {
|
||||
n.path: set(n.chunk_ids) for n in await self.file_graph.get_nodes([node.path for node, _ in files])
|
||||
}
|
||||
|
||||
await super().upsert(files)
|
||||
|
||||
if self._faiss_index is None or self.embedding_model is None:
|
||||
return
|
||||
self._sync_index_after_upsert(files, old_ids_by_path)
|
||||
|
||||
def _sync_index_after_upsert(
|
||||
self,
|
||||
files: list[tuple[FileNode, list[FileChunk]]],
|
||||
old_ids_by_path: dict[str, set[str]],
|
||||
) -> None:
|
||||
"""Apply add/tombstone deltas to FAISS based on chunk_id set differences."""
|
||||
existing = set(self._id_to_row)
|
||||
to_add: list[FileChunk] = []
|
||||
for node, _ in files:
|
||||
new_ids = set(node.chunk_ids)
|
||||
old_ids = old_ids_by_path.get(node.path, set())
|
||||
for cid in old_ids - new_ids:
|
||||
for cid in old_ids_by_path.get(node.path, set()) - new_ids:
|
||||
self._tombstone(cid)
|
||||
for cid in new_ids - existing:
|
||||
chunk = self.file_chunks.get(cid)
|
||||
|
|
@ -228,13 +243,16 @@ class FaissLocalFileStore(LocalFileStore):
|
|||
if query_embedding is None:
|
||||
return []
|
||||
|
||||
# Over-fetch by len(tombstones) so dropped rows can't starve the result set.
|
||||
q = self._prepare(query_embedding)
|
||||
# Over-fetch to compensate for tombstoned rows.
|
||||
k = min(self._faiss_index.ntotal, limit + len(self._tombstones))
|
||||
scores, rows = self._faiss_index.search(q, k)
|
||||
return self._collect_hits(rows[0].tolist(), scores[0].tolist(), limit)
|
||||
|
||||
def _collect_hits(self, rows: list[int], scores: list[float], limit: int) -> list[FileChunk]:
|
||||
"""Map raw FAISS rows back to chunks, skipping tombstones and stale ids."""
|
||||
results: list[FileChunk] = []
|
||||
for raw_row, score in zip(rows[0].tolist(), scores[0].tolist()):
|
||||
for raw_row, score in zip(rows, scores):
|
||||
row = int(raw_row)
|
||||
if row < 0 or row in self._tombstones or row >= len(self._id_map):
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -52,7 +52,7 @@ class LocalFileStore(BaseFileStore):
|
|||
self.file_chunks: dict[str, FileChunk] = {}
|
||||
self.chunks_path = self.component_metadata_path / f"file_chunks_{self.name}_{self.store_version}.jsonl"
|
||||
|
||||
# Lifecycle
|
||||
# -- lifecycle ------------------------------------------------------------
|
||||
|
||||
async def _start(self) -> None:
|
||||
await super()._start()
|
||||
|
|
@ -73,8 +73,10 @@ class LocalFileStore(BaseFileStore):
|
|||
self.logger.error(f"{self.name}: embedding disabled, {reason}")
|
||||
self.embedding_model = None
|
||||
|
||||
# -- persistence ----------------------------------------------------------
|
||||
|
||||
async def load(self) -> None:
|
||||
"""Load chunks from JSONL file into memory."""
|
||||
"""Load chunks from the JSONL file into memory; missing file is a no-op."""
|
||||
if not self.chunks_path.exists():
|
||||
return
|
||||
try:
|
||||
|
|
@ -89,7 +91,7 @@ class LocalFileStore(BaseFileStore):
|
|||
self.logger.exception(f"Failed to load {self.chunks_path}: {e}")
|
||||
|
||||
async def dump(self) -> None:
|
||||
"""Persist chunks to JSONL via atomic rename, then cascade to keyword_index and file_graph."""
|
||||
"""Atomically rewrite the JSONL, then cascade dump into keyword_index and file_graph."""
|
||||
assert self.file_graph is not None
|
||||
try:
|
||||
tmp = self.chunks_path.with_suffix(".tmp")
|
||||
|
|
@ -103,7 +105,7 @@ class LocalFileStore(BaseFileStore):
|
|||
await self.keyword_index.dump()
|
||||
await self.file_graph.dump()
|
||||
|
||||
# CRUD
|
||||
# -- CRUD -----------------------------------------------------------------
|
||||
|
||||
async def upsert(self, files: list[tuple[FileNode, list[FileChunk]]]) -> None:
|
||||
if not files:
|
||||
|
|
@ -216,7 +218,7 @@ class LocalFileStore(BaseFileStore):
|
|||
await self.keyword_index.clear()
|
||||
await self.file_graph.clear()
|
||||
|
||||
# Search
|
||||
# -- search ---------------------------------------------------------------
|
||||
|
||||
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
|
||||
if self.embedding_model is None or not query:
|
||||
|
|
@ -261,7 +263,7 @@ class LocalFileStore(BaseFileStore):
|
|||
|
||||
return results
|
||||
|
||||
# Extensions
|
||||
# -- extensions -----------------------------------------------------------
|
||||
|
||||
async def rebuild_links(self) -> None:
|
||||
"""Rebuild graph links via the underlying file graph."""
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from ...enumeration import ComponentEnum
|
|||
|
||||
|
||||
class BaseKeywordIndex(BaseComponent):
|
||||
"""关键词索引基类:定义增、删、查、清的统一接口,由具体实现(如 BM25)继承。"""
|
||||
"""Common interface for keyword indexes (add / delete / retrieve / clear)."""
|
||||
|
||||
component_type = ComponentEnum.KEYWORD_INDEX
|
||||
|
||||
|
|
@ -14,7 +14,6 @@ class BaseKeywordIndex(BaseComponent):
|
|||
super().__init__(**kwargs)
|
||||
from ..tokenizer import RegexTokenizer
|
||||
|
||||
# 绑定分词器,未显式指定时回落到 RegexTokenizer
|
||||
self.tokenizer = self.bind(tokenizer, BaseTokenizer, default_factory=RegexTokenizer)
|
||||
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
|
@ -25,7 +24,7 @@ class BaseKeywordIndex(BaseComponent):
|
|||
await self.dump()
|
||||
|
||||
def _tokenize(self, text: str) -> list[str]:
|
||||
"""对单段文本调用分词器,返回 token 列表。"""
|
||||
"""Tokenize a single text into a list of tokens."""
|
||||
if self.tokenizer is None:
|
||||
raise RuntimeError("Tokenizer not initialized. Call start() first.")
|
||||
return self.tokenizer.tokenize([text])[0]
|
||||
|
|
@ -43,11 +42,11 @@ class BaseKeywordIndex(BaseComponent):
|
|||
async def clear(self) -> None: ...
|
||||
|
||||
async def reset_index(self, docs_dict: dict[str, str]) -> None:
|
||||
"""清空索引后重新构建,并立即落盘。"""
|
||||
"""Wipe the index, rebuild it from `docs_dict`, and persist the result."""
|
||||
await self.clear()
|
||||
await self.add_docs(docs_dict)
|
||||
await self.dump()
|
||||
|
||||
async def optimize_index(self) -> None:
|
||||
"""对索引进行物理压缩或重建;基类默认无操作,由子类按需重载。"""
|
||||
"""Compact or rebuild the index. No-op by default; override as needed."""
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
"""基于 BM25 的倒排索引实现,支持持久化。
|
||||
"""BM25 inverted index with on-disk persistence.
|
||||
|
||||
核心存储结构(落盘真相源):
|
||||
On-disk truth source (see `_snapshot` / `_restore`):
|
||||
vocab : dict[token, token_id]
|
||||
_doc_ids : list[doc_id],按 doc_idx 索引
|
||||
_doc_ids : list[doc_id], indexed by doc_idx
|
||||
_doc_id_to_idx : dict[doc_id, doc_idx]
|
||||
_doc_lens : np.ndarray[int32],按 doc_idx 索引
|
||||
_deleted : np.ndarray[bool],按 doc_idx 索引(懒删除标记)
|
||||
_doc_token_ids : list[np.ndarray[int32]],每篇文档去重后的 token_id
|
||||
_posting_doc_idxs : dict[token_id, np.ndarray[int32]],倒排表的 doc_idx
|
||||
_posting_tfs : dict[token_id, np.ndarray[int32]],与上方一一对应的词频
|
||||
_doc_lens : np.int32[n], indexed by doc_idx
|
||||
_deleted : np.bool[n], lazy-delete flag per doc_idx
|
||||
_doc_token_ids : list[np.int32[]], unique token_ids per doc
|
||||
_posting_doc_idxs : dict[token_id, np.int32[]], posting list (doc_idx)
|
||||
_posting_tfs : dict[token_id, np.int32[]], aligned term frequencies
|
||||
|
||||
删除采用懒标记:_deleted[idx] = True 即视为删除,倒排表中的物理回收由
|
||||
optimize_index 统一完成;更新已存在的 doc_id 时,先把旧槽位标记删除,再
|
||||
分配新的 idx。
|
||||
Deletion is lazy: setting `_deleted[idx] = True` retires the slot. The posting
|
||||
lists keep the stale entries until `optimize_index` rewrites them. Updating an
|
||||
existing doc_id retires the old slot first, then allocates a fresh idx.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
|
@ -35,7 +35,6 @@ class BM25Index(BaseKeywordIndex):
|
|||
self.b = b
|
||||
self.index_version = index_version
|
||||
|
||||
# 词表与文档元数据
|
||||
self.vocab: dict[str, int] = {}
|
||||
self._doc_ids: list[str] = []
|
||||
self._doc_id_to_idx: dict[str, int] = {}
|
||||
|
|
@ -43,18 +42,17 @@ class BM25Index(BaseKeywordIndex):
|
|||
self._deleted: np.ndarray = np.zeros(0, dtype=bool)
|
||||
self._doc_token_ids: list[np.ndarray] = []
|
||||
|
||||
# 倒排表:token_id -> (doc_idxs, tfs)
|
||||
self._posting_doc_idxs: dict[int, np.ndarray] = {}
|
||||
self._posting_tfs: dict[int, np.ndarray] = {}
|
||||
|
||||
# IDF 缓存,对增删与重建索引时失效
|
||||
# IDF cache; invalidated whenever live-doc count or postings change.
|
||||
self._idf_cache: dict[int, float] = {}
|
||||
|
||||
# -- Properties -----------------------------------------------------------
|
||||
|
||||
@property
|
||||
def index_file(self) -> Path:
|
||||
"""落盘文件路径,包含分词器名与索引版本,便于区分不同配置。"""
|
||||
"""Path of the persisted index, namespaced by tokenizer and version."""
|
||||
if self.tokenizer is None:
|
||||
raise RuntimeError("Tokenizer not initialized. Call start() first.")
|
||||
name = type(self.tokenizer).__name__.replace("Tokenizer", "").lower()
|
||||
|
|
@ -62,23 +60,23 @@ class BM25Index(BaseKeywordIndex):
|
|||
|
||||
@property
|
||||
def n_docs(self) -> int:
|
||||
"""当前存活文档数(排除懒删除)。"""
|
||||
"""Number of live (non-deleted) documents."""
|
||||
return 0 if self._deleted.size == 0 else int((~self._deleted).sum())
|
||||
|
||||
@property
|
||||
def total_len(self) -> int:
|
||||
"""所有存活文档的 token 总数。"""
|
||||
"""Sum of token counts across live documents."""
|
||||
return 0 if self._deleted.size == 0 else int(self._doc_lens[~self._deleted].sum())
|
||||
|
||||
@property
|
||||
def avg_len(self) -> float:
|
||||
"""存活文档的平均长度,用于 BM25 长度归一化。"""
|
||||
"""Average length of live documents, used for BM25 length normalization."""
|
||||
n = self.n_docs
|
||||
return self.total_len / n if n > 0 else 0.0
|
||||
|
||||
@property
|
||||
def doc_meta(self) -> dict[str, dict]:
|
||||
"""对外暴露每篇存活文档的长度与去重后的 token_id 集合。"""
|
||||
"""Per-live-doc length and unique token_id set, keyed by doc_id."""
|
||||
return {
|
||||
self._doc_ids[idx]: {
|
||||
"len": int(self._doc_lens[idx]),
|
||||
|
|
@ -90,7 +88,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
|
||||
@property
|
||||
def inverted_index(self) -> dict[int, dict[str, int]]:
|
||||
"""重建可读形式的倒排表:token_id -> {doc_id: tf},跳过已删除文档。"""
|
||||
"""Readable view of postings: token_id -> {doc_id: tf}, deleted skipped."""
|
||||
out: dict[int, dict[str, int]] = {}
|
||||
for tid, doc_idxs in self._posting_doc_idxs.items():
|
||||
tfs = self._posting_tfs[tid]
|
||||
|
|
@ -106,7 +104,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
# -- Internal helpers -----------------------------------------------------
|
||||
|
||||
def _tokens_to_ids(self, tokens: list[str]) -> list[int]:
|
||||
"""将 token 转为 id;遇到新词时自动分配新的 token_id。"""
|
||||
"""Map tokens to ids, allocating a fresh id for any unseen token."""
|
||||
vocab = self.vocab
|
||||
ids: list[int] = []
|
||||
for token in tokens:
|
||||
|
|
@ -121,7 +119,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
return ids
|
||||
|
||||
def _remove_doc(self, doc_id: str) -> None:
|
||||
"""懒删除:仅置 _deleted 位并解除 doc_id 映射,不动倒排表。"""
|
||||
"""Lazy-delete a doc: flip `_deleted` and drop the id mapping."""
|
||||
idx = self._doc_id_to_idx.get(doc_id)
|
||||
if idx is None or self._deleted[idx]:
|
||||
return
|
||||
|
|
@ -130,7 +128,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
self._idf_cache = {}
|
||||
|
||||
def _get_idf(self, token_id: int, n_docs: int | None = None) -> float:
|
||||
"""计算并缓存 token 的 IDF;存活文档数发生变化时缓存会被清空。"""
|
||||
"""Return the cached IDF for a token, computing it on miss."""
|
||||
if token_id in self._idf_cache:
|
||||
return self._idf_cache[token_id]
|
||||
doc_idxs = self._posting_doc_idxs.get(token_id)
|
||||
|
|
@ -145,7 +143,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
return idf
|
||||
|
||||
def _prepare_doc(self, doc_id: str, content: str) -> tuple[np.ndarray, int, Counter] | None:
|
||||
"""分词并统计词频;若 doc_id 已存在则先标记旧版为删除。空文档返回 None。"""
|
||||
"""Tokenize and count terms; retire any prior version of `doc_id`."""
|
||||
self._remove_doc(doc_id)
|
||||
token_ids = self._tokens_to_ids(self._tokenize(content))
|
||||
if not token_ids:
|
||||
|
|
@ -157,7 +155,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
def _append_doc_arrays(
|
||||
self, new_doc_ids: list[str], new_doc_lens: list[int], new_doc_token_ids: list[np.ndarray]
|
||||
) -> None:
|
||||
"""把一批新文档的元数据一次性追加到文档数组中。"""
|
||||
"""Append metadata for a batch of new docs to the doc-level arrays."""
|
||||
if not new_doc_ids:
|
||||
return
|
||||
self._doc_ids.extend(new_doc_ids)
|
||||
|
|
@ -166,7 +164,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
self._deleted = np.concatenate([self._deleted, np.zeros(len(new_doc_ids), dtype=bool)])
|
||||
|
||||
def _extend_postings(self, pending: dict[int, list[tuple[int, int]]]) -> None:
|
||||
"""把待写入的 (doc_idx, tf) 增量按 token 追加到倒排表。"""
|
||||
"""Append pending (doc_idx, tf) pairs to each token's posting list."""
|
||||
for tid, items in pending.items():
|
||||
n = len(items)
|
||||
new_idxs = np.fromiter((idx for idx, _ in items), dtype=np.int32, count=n)
|
||||
|
|
@ -179,12 +177,12 @@ class BM25Index(BaseKeywordIndex):
|
|||
self._posting_tfs[tid] = new_tfs
|
||||
|
||||
def _encode_query(self, query: str) -> list[int]:
|
||||
"""切词、过滤未登录词并去重,返回查询的 token_id 列表。"""
|
||||
"""Tokenize query; drop OOV terms; deduplicate while preserving order."""
|
||||
vocab = self.vocab
|
||||
return list(dict.fromkeys(vocab[t] for t in self._tokenize(query) if t in vocab))
|
||||
|
||||
def _top_k(self, scores: np.ndarray, limit: int) -> np.ndarray:
|
||||
"""挑出得分前 limit 名(且严格大于 0)的索引,按得分降序排列。"""
|
||||
"""Indices of the top `limit` strictly-positive scores, descending."""
|
||||
if limit <= 0:
|
||||
return np.empty(0, dtype=np.int64)
|
||||
positive_count = int((scores > 0).sum())
|
||||
|
|
@ -199,7 +197,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
# -- Public API -----------------------------------------------------------
|
||||
|
||||
async def add_docs(self, docs_dict: dict[str, str]) -> None:
|
||||
"""批量加入文档;已存在的 doc_id 会被替换为新版本。"""
|
||||
"""Add or replace documents in batch (existing doc_ids are overwritten)."""
|
||||
if not docs_dict:
|
||||
return
|
||||
|
||||
|
|
@ -229,13 +227,13 @@ class BM25Index(BaseKeywordIndex):
|
|||
self._idf_cache = {}
|
||||
|
||||
async def delete_docs(self, doc_ids: list[str]) -> None:
|
||||
"""批量懒删除;倒排表中的物理回收由 optimize_index 完成。"""
|
||||
"""Lazy-delete a batch of doc_ids; physical reclaim happens in optimize_index."""
|
||||
for doc_id in doc_ids:
|
||||
self._remove_doc(doc_id)
|
||||
self._idf_cache = {}
|
||||
|
||||
def _score_query(self, query_ids: list[int], n_docs: int) -> np.ndarray:
|
||||
"""对所有文档计算 BM25 得分,已删除文档置 0。"""
|
||||
"""Compute BM25 scores across all docs; deleted docs zeroed out."""
|
||||
avg_len = self.total_len / n_docs
|
||||
k1, b = self.k1, self.b
|
||||
denom_base = k1 * (1.0 - b)
|
||||
|
|
@ -251,8 +249,9 @@ class BM25Index(BaseKeywordIndex):
|
|||
continue
|
||||
tfs = self._posting_tfs[tid].astype(np.float32)
|
||||
d_lens = self._doc_lens[doc_idxs].astype(np.float32)
|
||||
# 同一倒排表中每个 doc_idx 至多出现一次:Counter 已在文档内去重,
|
||||
# 文档更新也会分配新的 idx,因此可安全使用花式索引累加。
|
||||
# Each doc_idx appears at most once per posting list (Counter dedups
|
||||
# within a doc, and updates allocate fresh idxs), so fancy-index
|
||||
# accumulation is safe here.
|
||||
scores[doc_idxs] += idf * tfs * (k1 + 1.0) / (tfs + denom_base + denom_norm * d_lens)
|
||||
|
||||
if self._deleted.any():
|
||||
|
|
@ -260,7 +259,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
return scores
|
||||
|
||||
async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]:
|
||||
"""对查询做 BM25 召回,返回 {doc_id: score},按得分降序。"""
|
||||
"""BM25 retrieval; returns {doc_id: score} sorted by score descending."""
|
||||
n_docs = self.n_docs
|
||||
if n_docs == 0:
|
||||
return {}
|
||||
|
|
@ -275,7 +274,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
# -- Persistence ----------------------------------------------------------
|
||||
|
||||
def _snapshot(self) -> dict:
|
||||
"""收集需要落盘的全部字段,集中在一处以便与 _restore 对齐。"""
|
||||
"""Bundle every persistent field; mirrors `_restore`."""
|
||||
return {
|
||||
"vocab": self.vocab,
|
||||
"doc_ids": self._doc_ids,
|
||||
|
|
@ -290,7 +289,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
}
|
||||
|
||||
def _restore(self, data: dict) -> None:
|
||||
"""从 _snapshot 产生的字典还原索引内部状态。"""
|
||||
"""Restore index state from a `_snapshot` dict."""
|
||||
self.vocab = data["vocab"]
|
||||
self._doc_ids = data["doc_ids"]
|
||||
self._doc_id_to_idx = data["doc_id_to_idx"]
|
||||
|
|
@ -304,7 +303,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
self._idf_cache = {}
|
||||
|
||||
async def dump(self) -> None:
|
||||
"""通过临时文件 + 原子替换的方式持久化索引,避免半写状态。"""
|
||||
"""Persist the index via temp file + atomic rename to avoid torn writes."""
|
||||
try:
|
||||
tmp = self.index_file.with_suffix(".tmp")
|
||||
with open(tmp, "wb") as f:
|
||||
|
|
@ -315,7 +314,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
self.logger.exception(f"Failed to write {self.index_file}: {e}")
|
||||
|
||||
async def load(self) -> None:
|
||||
"""读取持久化文件并还原索引;文件不存在则不做事,损坏则清空。"""
|
||||
"""Load from disk; missing file is a no-op, corrupt file resets state."""
|
||||
if not self.index_file.exists():
|
||||
return
|
||||
try:
|
||||
|
|
@ -329,7 +328,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
await self.clear()
|
||||
|
||||
async def clear(self) -> None:
|
||||
"""清空内存中的索引并删除持久化文件。"""
|
||||
"""Reset in-memory state and remove the persisted file."""
|
||||
self.vocab = {}
|
||||
self._doc_ids = []
|
||||
self._doc_id_to_idx = {}
|
||||
|
|
@ -344,7 +343,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
# -- Compaction -----------------------------------------------------------
|
||||
|
||||
def _build_idx_remap(self, active_mask: np.ndarray) -> tuple[np.ndarray, int]:
|
||||
"""构造 old_idx → new_idx 的映射数组(被删槽位为 -1),并返回存活数量。"""
|
||||
"""Build an old_idx -> new_idx array (-1 for retired slots)."""
|
||||
active_old_idxs = np.where(active_mask)[0]
|
||||
n_active = int(active_old_idxs.size)
|
||||
remap = -np.ones(self._deleted.size, dtype=np.int32)
|
||||
|
|
@ -352,7 +351,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
return remap, n_active
|
||||
|
||||
def _compact_vocab(self, active_mask: np.ndarray) -> tuple[dict[str, int], dict[int, int]]:
|
||||
"""只保留仍被任意存活文档引用的 token,重排成连续的新 token_id。"""
|
||||
"""Keep only tokens still referenced by a live doc; renumber contiguously."""
|
||||
used_tids = {
|
||||
tid for tid, doc_idxs in self._posting_doc_idxs.items()
|
||||
if active_mask[doc_idxs].any()
|
||||
|
|
@ -372,7 +371,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
old_to_new_idx: np.ndarray,
|
||||
old_tid_to_new: dict[int, int],
|
||||
) -> tuple[dict[int, np.ndarray], dict[int, np.ndarray]]:
|
||||
"""剔除删除文档并按新 idx/tid 重写倒排表。"""
|
||||
"""Drop deleted entries and rewrite postings under new idx/tid numbering."""
|
||||
new_idxs: dict[int, np.ndarray] = {}
|
||||
new_tfs: dict[int, np.ndarray] = {}
|
||||
for tid, doc_idxs in self._posting_doc_idxs.items():
|
||||
|
|
@ -387,7 +386,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
def _compact_docs(
|
||||
self, active_mask: np.ndarray, old_tid_to_new: dict[int, int]
|
||||
) -> tuple[list[str], list[np.ndarray]]:
|
||||
"""在压缩后的词表下重建存活文档的 doc_id 列表与去重 token_id 数组。"""
|
||||
"""Rebuild doc_id list and unique-token arrays under the new vocab."""
|
||||
active_old_idxs = np.where(active_mask)[0]
|
||||
new_doc_ids = [self._doc_ids[int(i)] for i in active_old_idxs]
|
||||
new_doc_token_ids = [
|
||||
|
|
@ -400,7 +399,7 @@ class BM25Index(BaseKeywordIndex):
|
|||
return new_doc_ids, new_doc_token_ids
|
||||
|
||||
async def optimize_index(self) -> None:
|
||||
"""物理回收懒删除的文档与未被引用的词表项,重建紧凑索引。"""
|
||||
"""Physically reclaim deleted docs and unused vocab entries."""
|
||||
if self._deleted.size == 0:
|
||||
return
|
||||
active_mask = ~self._deleted
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue