-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_vector_sql.py
More file actions
67 lines (55 loc) · 2.37 KB
/
Copy pathtest_vector_sql.py
File metadata and controls
67 lines (55 loc) · 2.37 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
from google.cloud import spanner
from vertexai.language_models import TextEmbeddingModel
import vertexai
PROJECT_ID = "alloydbtest-374215"
INSTANCE_ID = "virtual-retail-instance"
DATABASE_ID = "catalog"
LOCATION = "us-central1"
vertexai.init(project=PROJECT_ID, location=LOCATION)
def get_embedding():
model = TextEmbeddingModel.from_pretrained("text-embedding-005")
embeddings = model.get_embeddings(["test query"])
return embeddings[0].values
def test_query():
client = spanner.Client(project=PROJECT_ID)
instance = client.instance(INSTANCE_ID)
database = instance.database(DATABASE_ID)
vector = get_embedding()
# Test 1: APPROX_COSINE_DISTANCE without options (Expected Failure)
sql1 = """
SELECT id, APPROX_COSINE_DISTANCE(embedding, @query_vector) as distance
FROM products@{FORCE_INDEX=ProductsVectorIndex}
WHERE embedding IS NOT NULL
ORDER BY distance
LIMIT 1
"""
# Test 2: APPROX_COSINE_DISTANCE WITH options (Expected Success)
# Options for ScaNN are typically '{"num_leaves_to_search": 10}' BUT Spanner might want a string literal that parses as a proto?
# Actually, let's try just empty options first or standard JSON if the error was just about the keys?
# The error "Expected identifier, got: {" suggests it wants a proto text format, NOT JSON.
# Proto text format: 'num_leaves_to_search: 10'
sql2 = """
SELECT id, APPROX_COSINE_DISTANCE(embedding, @query_vector, options => 'num_leaves_to_search: 10') as distance
FROM products@{FORCE_INDEX=ProductsVectorIndex}
WHERE embedding IS NOT NULL
ORDER BY distance
LIMIT 1
"""
params = {"query_vector": vector}
param_types = {"query_vector": spanner.param_types.Array(spanner.param_types.FLOAT64)}
print("Testing SQL 1 (No options)...")
try:
with database.snapshot() as snapshot:
results = list(snapshot.execute_sql(sql1, params=params, param_types=param_types))
print("SQL 1 Success")
except Exception as e:
print(f"SQL 1 Failed: {e}")
print("\nTesting SQL 2 (With options)...")
try:
with database.snapshot() as snapshot:
results = list(snapshot.execute_sql(sql2, params=params, param_types=param_types))
print("SQL 2 Success")
except Exception as e:
print(f"SQL 2 Failed: {e}")
if __name__ == "__main__":
test_query()