|
8 | 8 | delete, |
9 | 9 | func, |
10 | 10 | literal_column, |
11 | | - values, |
12 | | - column, |
13 | | - Integer, |
14 | | - String, |
15 | 11 | cast, |
16 | 12 | Float, |
17 | 13 | ) |
18 | 14 | from pgvector.sqlalchemy import Vector |
19 | | -from typing import List, Set, Dict, Any |
| 15 | +from typing import List, Set |
20 | 16 | import dataclasses |
21 | 17 |
|
22 | 18 |
|
@@ -64,72 +60,23 @@ async def delete_event_data(self, event_code: str) -> None: |
64 | 60 | async def find_matches( |
65 | 61 | self, |
66 | 62 | event_code: str, |
67 | | - encodings: List[List[float]], |
| 63 | + encoding: List[float], |
68 | 64 | threshold: float, |
69 | | - min_matches: int, |
70 | 65 | ) -> List[str]: |
71 | | - ref_encodings = ( |
72 | | - values( |
73 | | - column("id", Integer), column("embedding", String), name="ref_encodings" |
74 | | - ) |
75 | | - .data([(i + 1, str(emb)) for i, emb in enumerate(encodings)]) |
76 | | - .cte("ref_encodings") |
77 | | - ) |
78 | | - |
79 | 66 | distance_op = EventEncodingModel.embedding.op("<=>", return_type=Float())( |
80 | | - cast(ref_encodings.c.embedding, Vector(512)) |
| 67 | + cast(str(encoding), Vector(512)) |
81 | 68 | ) |
82 | 69 |
|
83 | 70 | stmt = ( |
84 | 71 | select( |
85 | 72 | EventEncodingModel.image_path, |
86 | | - func.count(ref_encodings.c.id).label("match_count"), |
87 | 73 | func.min(distance_op).label("best_distance"), |
88 | 74 | ) |
89 | | - .join(ref_encodings, literal_column("true")) |
90 | 75 | .where(EventEncodingModel.event_code == event_code, distance_op < threshold) |
91 | 76 | .group_by(EventEncodingModel.image_path) |
92 | | - .having(func.count(ref_encodings.c.id) >= min_matches) |
93 | | - .order_by( |
94 | | - literal_column("match_count").desc(), |
95 | | - literal_column("best_distance").asc(), |
96 | | - ) |
| 77 | + .order_by(literal_column("best_distance").asc()) |
97 | 78 | ) |
98 | 79 |
|
99 | 80 | result = await self.session.execute(stmt) |
100 | 81 | rows = result.all() |
101 | 82 | return [row[0] for row in rows] |
102 | | - |
103 | | - async def get_closest_matches_debug( |
104 | | - self, event_code: str, encodings: List[List[float]], limit: int = 5 |
105 | | - ) -> List[Dict[str, Any]]: |
106 | | - ref_encodings = ( |
107 | | - values( |
108 | | - column("id", Integer), column("embedding", String), name="ref_encodings" |
109 | | - ) |
110 | | - .data([(i + 1, str(emb)) for i, emb in enumerate(encodings)]) |
111 | | - .cte("ref_encodings") |
112 | | - ) |
113 | | - |
114 | | - distance_op = EventEncodingModel.embedding.op("<=>", return_type=Float())( |
115 | | - cast(ref_encodings.c.embedding, Vector(512)) |
116 | | - ) |
117 | | - |
118 | | - stmt = ( |
119 | | - select( |
120 | | - EventEncodingModel.image_path, |
121 | | - func.count(ref_encodings.c.id).label("match_count"), |
122 | | - func.min(distance_op).label("best_distance"), |
123 | | - ) |
124 | | - .join(ref_encodings, literal_column("true")) |
125 | | - .where(EventEncodingModel.event_code == event_code) |
126 | | - .group_by(EventEncodingModel.image_path) |
127 | | - .order_by(literal_column("best_distance").asc()) |
128 | | - .limit(limit) |
129 | | - ) |
130 | | - |
131 | | - result = await self.session.execute(stmt) |
132 | | - return [ |
133 | | - {"image_path": row[0], "match_count": row[1], "best_distance": row[2]} |
134 | | - for row in result.all() |
135 | | - ] |
0 commit comments