DuckDBRetriever 来了
在上篇文章简要地用DuckDB fts 和BM25Retriever 的结果做了下对比, BM25Retriever 需要把索引的内容全部加载到nodes(内存里),数据小的时候,还好办,如果数据比较大,就会比较麻烦, llamaindex 仓库很多人提出相关issue。于是我尝试着利用DuckDB fts 封装了DuckDBRetriever(假设文件名为duckdb_retriever.py) ,代码如下,
import logging
import os
from typing import List, Optionalfrom llama_index.core.base.base_retriever import BaseRetriever
from llama_index.core.callbacks.base import CallbackManager
from llama_index.core.constants import DEFAULT_SIMILARITY_TOP_K
from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle
logger = logging.getLogger(__name__)
import_err_msg = "`duckdb` package not found, please run `pip install duckdb`"
class DuckDBLocalContext:
def __init__(self, database_path: str):
self.database_path = database_path
self._conn = None
self._home_dir = os.path.expanduser("~")
def __enter__(self) -> "duckdb.DuckDBPyConnection":
try:
import duckdb
except ImportError:
raise ImportError(import_err_msg)
if not os.path.exists(os.path.dirname(self.database_path)):
raise ValueError(
f"Directory {os.path.dirname(self.database_path)} does not exist."
)
# if not os.path.isfile(self.database_path):
# raise ValueError(f"Database path {self.database_path} is not a valid file.")
self._conn = duckdb.connect(self.database_path)
self._conn.execute(f"SET home_directory='{self._home_dir}';")
self._conn.install_extension("fts")
self._conn.load_extension("fts")
return self._conn
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
self._conn.close()
if self._conn:
self._conn.close()
class DuckDBRetriever(BaseRetriever):
def __init__(
self,
database_name: Optional[str] = ":memory:",
table_name: Optional[str] = "documents",
text_search_config: Optional[dict] = {
"stemmer": "english",
"stopwords": "english",
"ignore": r"(\\.|[^a-z])+",
"strip_accents": True,
"lower": True,
"overwrite": True,
},
persist_dir: Optional[str] = "./storage",
node_id_column: Optional[str] = "node_id",
text_column: Optional[str] = "text",
# TODO: Add more options for FTS index creation
similarity_top_k: int = DEFAULT_SIMILARITY_TOP_K,
callback_manager: Optional[CallbackManager] = None,
verbose: bool = False,
) -> None:
self._similarity_top_k = similarity_top_k
self._callback_manager = callback_manager
self._verbose = verbose
self._table_name = table_name
self._node_id_column = node_id_column
self._text_column = text_column
# TODO: Check if the vector store already has data
# Create an FTS index on the 'text' column if it doesn't already exist
if database_name == ':memory:':
self._database_path = ':memory:'
else:
self._database_path = os.path.join(persist_dir, database_name)
strip_accents = 1 if text_search_config["strip_accents"] else 0
lower = 1 if text_search_config["lower"] else 0
overwrite = 1 if text_search_config["overwrite"] else 0
ignore = text_search_config["ignore"]
sql = f"""
PRAGMA create_fts_index({self._table_name}, {self._node_id_column}, {self._text_column},
stemmer = '{text_search_config["stemmer"]}',
stopwords = '{text_search_config["stopwords"]}', ignore = '{ignore}',
strip_accents = {strip_accents}, lower = {lower}, overwrite = {overwrite})
"""
with DuckDBLocalContext(self._database_path) as conn:
conn.execute(sql)
def _retrieve(self, query_bundle: QueryBundle) -> List[NodeWithScore]:
if self._verbose:
logger.info(f"Searching for: {query_bundle.query_str}")
query = query_bundle.query_str
sql = f"""
SELECT
fts_main_{self._table_name}.match_bm25({self._node_id_column}, '{query}') AS score,
{self._node_id_column}, {self._text_column}
FROM {self._table_name}
WHERE score IS NOT NULL
ORDER BY score DESC
LIMIT {self._similarity_top_k};
"""
with DuckDBLocalContext(self._database_path) as conn:
query_result = conn.execute(sql).fetchall()
# Convert query result to NodeWithScore objects
retrieve_nodes = []
for row in query_result:
score, node_id, text = row
node = TextNode(id=node_id, text=text)
retrieve_nodes.append(NodeWithScore(node=node, score=float(score)))
return retrieve_nodes
使用方法
与DuckDBVectorStore 配合使用
上篇文章的代码,将BM25Retriever 实例化部分
bm25_retriever = BM25Retriever.from_defaults(nodes=nodes, similarity_top_k=5)替换成 DuckDBRetriever即可,
from duckdb_retriever import DuckDBRetrieverbm25_retriever = DuckDBRetriever(database_name="paul.duck",
persist_dir="duckdb",similarity_top_k=5)
其中database_name 和 persist_dir 和DuckDBVectorStore的参数保持一致。
仅仅使用DuckDBRetriever
• 建表及数据准备
import duckdb
with duckdb.connect('fts.db') as conn:
conn.sql("""
CREATE TABLE documents (
node_id VARCHAR,
text VARCHAR,
author VARCHAR,
doc_version INTEGER
);
INSERT INTO documents
VALUES ('doc1',
'The mallard is a dabbling duck that breeds throughout the temperate.',
'Hannes Mühleisen',
3),
('doc2',
'The cat is a domestic species of small carnivorous mammal.',
'Laurens Kuiper',
2
);
""")
# conn.close()• 索引及搜索
from duckdb_retriever import DuckDBRetriever
bm25_retriever = DuckDBRetriever(database_name="fts.db",
table_name='documents',
node_id_column='node_id',
text_column='text',
persist_dir=".",
similarity_top_k=5)
query = "small cat"
from llama_index.core.response.notebook_utils import display_source_nodenodes_bm25 = bm25_retriever.retrieve(query)
for node in nodes_bm25:
display_source_node(node)