Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 49 additions & 51 deletions sdk/python/feast/infra/online_stores/sqlite.py
Original file line number Diff line number Diff line change
Expand Up @@ -568,30 +568,7 @@ def retrieve_online_documents(
)
vector_field = _get_vector_field(table)

cur.execute(
f"""
CREATE VIRTUAL TABLE vec_table using vec0(
vector_value float[{vector_field_length}]
);
"""
)

# Currently I can only insert the embedding value without crashing SQLite, will report a bug
cur.execute(
f"""
INSERT INTO vec_table(rowid, vector_value)
select rowid, vector_value from {_quote_id(table_name)}
where feature_name = ?
""",
(vector_field,),
)
cur.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS vec_table using vec0(
vector_value float[{vector_field_length}]
);
"""
)
_build_vec_table(conn, table_name, vector_field, vector_field_length)

# Have to join this with the main table to get the feature name and entity_key
# Also the `top_k` doesn't appear to be working for some reason
Expand Down Expand Up @@ -720,21 +697,7 @@ def retrieve_online_documents_v2(

if online_store.vector_enabled:
query_embedding_bin = serialize_f32(embedding, vector_field_length) # type: ignore
cur.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS vec_table using vec0(
vector_value float[{vector_field_length}]
);
"""
)
cur.execute(
f"""
INSERT INTO vec_table (rowid, vector_value)
select rowid, vector_value from {_quote_id(table_name)}
where feature_name = ?
""",
(vector_field,),
)
_build_vec_table(conn, table_name, vector_field, vector_field_length)
elif online_store.text_search_enabled:
string_field_list = [
f.name for f in table.features if f.dtype == PrimitiveFeastType.STRING
Expand All @@ -747,17 +710,7 @@ def retrieve_online_documents_v2(
if f.dtype == PrimitiveFeastType.STRING
]
)
cur.execute(
f"""
CREATE VIRTUAL TABLE IF NOT EXISTS search_table using fts5(
entity_key, fv_rowid, {string_fields}, tokenize="porter unicode61"
);
"""
)
insert_query = _generate_bm25_search_insert_query(
table_name, string_field_list
)
cur.execute(insert_query)
_build_search_table(conn, table_name, string_field_list)
filter_clause, filter_params = SqliteFilterTranslator(
table_name, alias="fv"
).translate(filters)
Expand Down Expand Up @@ -1065,6 +1018,51 @@ def _get_vector_field(table: FeatureView) -> str:
return vector_field


def _build_vec_table(
conn: sqlite3.Connection,
table_name: str,
vector_field: str,
vector_length: int,
) -> None:
"""
Index the stored vectors of a feature view table in `temp.vec_table`.

The index is rebuilt from the current rows for every search, so rows copied
by earlier searches are never searched again. It lives in the connection's
temp schema and is committed right away, so a search neither writes to the
online store database nor keeps it locked.
"""
with conn:
conn.execute("DROP TABLE IF EXISTS temp.vec_table")
conn.execute(
f"CREATE VIRTUAL TABLE temp.vec_table "
f"USING vec0(vector_value float[{vector_length}])"
)
conn.execute(
f"INSERT INTO temp.vec_table (rowid, vector_value) "
f"SELECT rowid, vector_value FROM {_quote_id(table_name)} "
f"WHERE feature_name = ?",
(vector_field,),
)


def _build_search_table(
conn: sqlite3.Connection, table_name: str, string_field_list: List[str]
) -> None:
"""
Index the stored string features of a feature view table in
`temp.search_table`, rebuilt for every search like `_build_vec_table`.
"""
with conn:
conn.execute("DROP TABLE IF EXISTS temp.search_table")
conn.execute(
f"CREATE VIRTUAL TABLE temp.search_table USING fts5("
f"entity_key, fv_rowid, {', '.join(string_field_list)}, "
f'tokenize="porter unicode61")'
)
conn.execute(_generate_bm25_search_insert_query(table_name, string_field_list))


def _generate_bm25_search_insert_query(
table_name: str, string_field_list: List[str]
) -> str:
Expand All @@ -1079,7 +1077,7 @@ def _generate_bm25_search_insert_query(
str: The generated SQL insertion query.
"""
_string_fields = ", ".join(string_field_list)
query = f"INSERT INTO search_table (entity_key, fv_rowid, {_string_fields})\nSELECT\n\tDISTINCT fv0.entity_key,\n\tfv0.rowid as fv_rowid"
query = f"INSERT INTO temp.search_table (entity_key, fv_rowid, {_string_fields})\nSELECT\n\tDISTINCT fv0.entity_key,\n\tfv0.rowid as fv_rowid"
quoted_table = _quote_id(table_name)
from_query = f"\nFROM (select rowid, * from {quoted_table} where feature_name = '{string_field_list[0]}') fv0"

Expand Down
116 changes: 116 additions & 0 deletions sdk/python/tests/unit/online_store/test_online_retrieval.py
Original file line number Diff line number Diff line change
Expand Up @@ -1045,6 +1045,122 @@ def test_sqlite_get_online_documents_v2_search() -> None:
assert result["distance"] == [-1.8458267450332642, -1.8458267450332642]


def test_sqlite_get_online_documents_v2_search_is_repeatable() -> None:
"""Every keyword search ranks each stored document once, as last written."""
runner = CliRunner()
with runner.local_repo(
get_example_repo("example_feature_repo_1.py"), "file"
) as store:
store.config.online_store.text_search_enabled = True
document_embeddings_fv = store.get_feature_view(name="document_embeddings")
provider = store._get_provider()

def write_document(item_id: int, content: str) -> None:
provider.online_write_batch(
config=store.config,
table=document_embeddings_fv,
data=[
(
EntityKeyProto(
join_keys=["item_id"],
entity_values=[ValueProto(int64_val=item_id)],
),
{
"content": ValueProto(string_val=content),
"title": ValueProto(string_val=f"Title {item_id}"),
},
_utc_now(),
_utc_now(),
)
],
progress=None,
)

def search(query_string: str) -> list[str]:
result = store.retrieve_online_documents_v2(
features=["document_embeddings:content", "document_embeddings:title"],
query_string=query_string,
top_k=3,
).to_dict()
return sorted(result["title"])

# The more often a document repeats the word, the better it ranks.
write_document(0, "sentence sentence sentence")
write_document(1, "sentence sentence")
write_document(2, "sentence")
write_document(3, "some other text")

for _ in range(3):
assert search("sentence") == ["Title 0", "Title 1", "Title 2"]

write_document(0, "a rewritten document")
assert search("sentence") == ["Title 1", "Title 2"]
assert search("rewritten") == ["Title 0"]


@pytest.mark.skipif(
sys.version_info[0:2] != (3, 10),
reason="Only works on Python 3.10",
)
def test_sqlite_get_online_documents_is_repeatable() -> None:
"""Vector searches keep working after the first one."""
vector_length = 8
runner = CliRunner()
with runner.local_repo(
get_example_repo("example_feature_repo_1.py"), "file"
) as store:
store.config.online_store.vector_enabled = True
document_embeddings_fv = store.get_feature_view(name="document_embeddings")
provider = store._get_provider()
provider.online_write_batch(
config=store.config,
table=document_embeddings_fv,
data=[
(
EntityKeyProto(
join_keys=["item_id"], entity_values=[ValueProto(int64_val=i)]
),
{
"Embeddings": ValueProto(
float_list_val=FloatListProto(
val=[float(i)] * vector_length
)
),
"content": ValueProto(string_val=f"the {i}th sentence"),
"title": ValueProto(string_val=f"Title {i}"),
},
_utc_now(),
_utc_now(),
)
for i in range(5)
],
progress=None,
)
query = [0.0] * vector_length

for _ in range(2):
result = store.retrieve_online_documents_v2(
features=[
"document_embeddings:Embeddings",
"document_embeddings:title",
],
query=query,
top_k=3,
).to_dict()
assert sorted(result["title"]) == ["Title 0", "Title 1", "Title 2"]

for _ in range(2):
result = store.retrieve_online_documents(
features=[
"document_embeddings:Embeddings",
"document_embeddings:distance",
],
query=query,
top_k=3,
).to_dict()
assert len(result["distance"]) == 3


@pytest.mark.skip(reason="Skipping this test as CI struggles with it")
def test_local_milvus() -> None:
import random
Expand Down
Loading