@@ -147,35 +147,42 @@ def _print_text():
147147 def get_all_collections (self ) -> dict [str , _CollectionConfigSimple ]:
148148 return self .client .collections .list_all ()
149149
150+ _DISTANCE_METRIC_MAP = {
151+ "cosine" : wvc .VectorDistances .COSINE ,
152+ "dot" : wvc .VectorDistances .DOT ,
153+ "l2-squared" : wvc .VectorDistances .L2_SQUARED ,
154+ "hamming" : wvc .VectorDistances .HAMMING ,
155+ "manhattan" : wvc .VectorDistances .MANHATTAN ,
156+ }
157+
158+ def _resolve_distance_metric (
159+ self , distance_metric : Optional [str ]
160+ ) -> Optional [wvc .VectorDistances ]:
161+ """Convert a distance metric string to its VectorDistances enum value."""
162+ if distance_metric is None :
163+ return None
164+ if distance_metric not in self ._DISTANCE_METRIC_MAP :
165+ raise ValueError (
166+ f"Invalid distance_metric: '{ distance_metric } '. "
167+ f"Must be one of: { list (self ._DISTANCE_METRIC_MAP .keys ())} "
168+ )
169+ return self ._DISTANCE_METRIC_MAP [distance_metric ]
170+
150171 def _build_hfresh_config (
151172 self ,
152173 max_posting_size_kb : Optional [int ] = None ,
153- distance_metric : Optional [str ] = "cosine" ,
174+ distance_metric : Optional [wvc . VectorDistances ] = None ,
154175 rescore_limit : Optional [int ] = None ,
155176 replicas : Optional [int ] = None ,
156177 search_probe : Optional [int ] = None ,
157178 ):
158179 """Build hfresh configuration with provided parameters."""
159- # Explicit mapping of distance metric strings to enum values
160- distance_metric_map = {
161- "cosine" : wvc .VectorDistances .COSINE ,
162- "dot" : wvc .VectorDistances .DOT ,
163- "l2-squared" : wvc .VectorDistances .L2_SQUARED ,
164- "hamming" : wvc .VectorDistances .HAMMING ,
165- "manhattan" : wvc .VectorDistances .MANHATTAN ,
166- }
167-
168180 kwargs = {}
169181
170182 if max_posting_size_kb is not None :
171183 kwargs ["max_posting_size_kb" ] = max_posting_size_kb
172184 if distance_metric is not None :
173- if distance_metric not in distance_metric_map :
174- raise ValueError (
175- f"Invalid distance_metric: '{ distance_metric } '. "
176- f"Must be one of: { list (distance_metric_map .keys ())} "
177- )
178- kwargs ["distance_metric" ] = distance_metric_map [distance_metric ]
185+ kwargs ["distance_metric" ] = distance_metric
179186 if replicas is not None :
180187 kwargs ["replicas" ] = replicas
181188 if search_probe is not None :
@@ -248,56 +255,62 @@ def create_collection(
248255 "Error: Named vector name is only supported with named vectors. Please use --named_vector to enable named vectors."
249256 )
250257
258+ distance_metric_enum = self ._resolve_distance_metric (distance_metric )
259+
251260 vector_index_map : Dict [str , wvc .VectorIndexConfig ] = {
252- "hnsw" : wvc .Configure .VectorIndex .hnsw (distance_metric = distance_metric ),
253- "flat" : wvc .Configure .VectorIndex .flat (distance_metric = distance_metric ),
261+ "hnsw" : wvc .Configure .VectorIndex .hnsw (
262+ distance_metric = distance_metric_enum
263+ ),
264+ "flat" : wvc .Configure .VectorIndex .flat (
265+ distance_metric = distance_metric_enum
266+ ),
254267 "dynamic" : wvc .Configure .VectorIndex .dynamic (),
255268 "dynamic_flat_bq" : wvc .Configure .VectorIndex .dynamic (
256269 flat = wvc .Configure .VectorIndex .flat (
257270 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (),
258- distance_metric = distance_metric ,
271+ distance_metric = distance_metric_enum ,
259272 )
260273 ),
261274 "dynamic_flat_bq_hnsw_pq" : wvc .Configure .VectorIndex .dynamic (
262275 flat = wvc .Configure .VectorIndex .flat (
263276 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
264277 rescore_limit = rescore_limit
265278 ),
266- distance_metric = distance_metric ,
279+ distance_metric = distance_metric_enum ,
267280 ),
268281 hnsw = wvc .Configure .VectorIndex .hnsw (
269282 quantizer = wvc .Configure .VectorIndex .Quantizer .pq (
270283 training_limit = training_limit
271284 ),
272- distance_metric = distance_metric ,
285+ distance_metric = distance_metric_enum ,
273286 ),
274287 ),
275288 "dynamic_flat_bq_hnsw_sq" : wvc .Configure .VectorIndex .dynamic (
276289 flat = wvc .Configure .VectorIndex .flat (
277290 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
278291 rescore_limit = rescore_limit
279292 ),
280- distance_metric = distance_metric ,
293+ distance_metric = distance_metric_enum ,
281294 ),
282295 hnsw = wvc .Configure .VectorIndex .hnsw (
283296 quantizer = wvc .Configure .VectorIndex .Quantizer .sq (
284297 rescore_limit = rescore_limit , training_limit = training_limit
285298 ),
286- distance_metric = distance_metric ,
299+ distance_metric = distance_metric_enum ,
287300 ),
288301 ),
289302 "dynamic_flat_bq_hnsw_bq" : wvc .Configure .VectorIndex .dynamic (
290303 flat = wvc .Configure .VectorIndex .flat (
291304 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
292305 rescore_limit = rescore_limit
293306 ),
294- distance_metric = distance_metric ,
307+ distance_metric = distance_metric_enum ,
295308 ),
296309 hnsw = wvc .Configure .VectorIndex .hnsw (
297310 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
298311 rescore_limit = rescore_limit
299312 ),
300- distance_metric = distance_metric ,
313+ distance_metric = distance_metric_enum ,
301314 ),
302315 ),
303316 "dynamic_hnsw_pq" : wvc .Configure .VectorIndex .dynamic (
@@ -312,70 +325,70 @@ def create_collection(
312325 quantizer = wvc .Configure .VectorIndex .Quantizer .sq (
313326 rescore_limit = rescore_limit , training_limit = training_limit
314327 ),
315- distance_metric = distance_metric ,
328+ distance_metric = distance_metric_enum ,
316329 )
317330 ),
318331 "dynamic_hnsw_bq" : wvc .Configure .VectorIndex .dynamic (
319332 hnsw = wvc .Configure .VectorIndex .hnsw (
320333 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
321334 rescore_limit = rescore_limit
322335 ),
323- distance_metric = distance_metric ,
336+ distance_metric = distance_metric_enum ,
324337 )
325338 ),
326339 "hnsw_pq" : wvc .Configure .VectorIndex .hnsw (
327340 quantizer = wvc .Configure .VectorIndex .Quantizer .pq (
328341 training_limit = training_limit
329342 ),
330- distance_metric = distance_metric ,
343+ distance_metric = distance_metric_enum ,
331344 ),
332345 "hnsw_bq" : wvc .Configure .VectorIndex .hnsw (
333346 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
334347 rescore_limit = rescore_limit
335348 ),
336- distance_metric = distance_metric ,
349+ distance_metric = distance_metric_enum ,
337350 ),
338351 "hnsw_bq_cache" : wvc .Configure .VectorIndex .hnsw (
339352 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
340353 cache = True , rescore_limit = rescore_limit
341354 ),
342- distance_metric = distance_metric ,
355+ distance_metric = distance_metric_enum ,
343356 ),
344357 "hnsw_sq" : wvc .Configure .VectorIndex .hnsw (
345358 quantizer = wvc .Configure .VectorIndex .Quantizer .sq (
346359 rescore_limit = rescore_limit , training_limit = training_limit
347360 ),
348- distance_metric = distance_metric ,
361+ distance_metric = distance_metric_enum ,
349362 ),
350363 "hnsw_rq" : wvc .Configure .VectorIndex .hnsw (
351364 quantizer = wvc .Configure .VectorIndex .Quantizer .rq (
352365 rescore_limit = rescore_limit
353366 ),
354- distance_metric = distance_metric ,
367+ distance_metric = distance_metric_enum ,
355368 ),
356369 "hnsw_acorn" : wvc .Configure .VectorIndex .hnsw (
357370 filter_strategy = VectorFilterStrategy .ACORN ,
358- distance_metric = distance_metric ,
371+ distance_metric = distance_metric_enum ,
359372 ),
360373 "hnsw_multivector" : wvc .Configure .VectorIndex .hnsw (
361374 multi_vector = wvc .Configure .VectorIndex .MultiVector .multi_vector (),
362- distance_metric = distance_metric ,
375+ distance_metric = distance_metric_enum ,
363376 ),
364377 "flat_bq" : wvc .Configure .VectorIndex .flat (
365378 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
366379 rescore_limit = rescore_limit
367380 ),
368- distance_metric = distance_metric ,
381+ distance_metric = distance_metric_enum ,
369382 ),
370383 "flat_bq_cache" : wvc .Configure .VectorIndex .flat (
371384 quantizer = wvc .Configure .VectorIndex .Quantizer .bq (
372385 cache = True , rescore_limit = rescore_limit
373386 ),
374- distance_metric = distance_metric ,
387+ distance_metric = distance_metric_enum ,
375388 ),
376389 "hfresh" : self ._build_hfresh_config (
377390 max_posting_size_kb = hfresh_max_posting_size_kb ,
378- distance_metric = distance_metric ,
391+ distance_metric = distance_metric_enum ,
379392 rescore_limit = rescore_limit ,
380393 replicas = hfresh_replicas ,
381394 search_probe = hfresh_search_probe ,
0 commit comments