MCPcopy Create free account
hub / github.com/cosdata/cosdata / calculate_recall

Function calculate_recall

tests/test-sparse-vector.py:369–422  ·  view source on GitHub ↗

Calculate recall metrics

(brute_force_results, server_results, top_k=10)

Source from the content-addressed store, hash-verified

367 return [[ind, val] for ind, val in zip(vector["indices"], vector["values"])]
368
369def calculate_recall(brute_force_results, server_results, top_k=10):
370 """Calculate recall metrics"""
371 recalls = []
372
373 for i, (bf_result, server_result) in enumerate(
374 zip(brute_force_results, server_results)
375 ):
376 # Debug: Print first few server results to understand structure
377 if i < 3:
378 print(f"\nDebug - Query {i}:")
379 print(f"Brute force top 3: {bf_result['top_results'][:3]}")
380 print(f"Server result keys: {list(server_result.keys())}")
381 print(f"Server result: {server_result}")
382
383 bf_ids = set(item["id"] for item in bf_result["top_results"])
384
385 # Handle different possible server response structures
386 if "results" in server_result:
387 server_items = server_result["results"]
388 elif "vectors" in server_result:
389 server_items = server_result["vectors"]
390 else:
391 print(f"Warning: Unexpected server result structure: {server_result}")
392 server_items = []
393
394 # Extract IDs from server results - handle different possible structures
395 server_ids = set()
396 for item in server_items:
397 if isinstance(item, dict):
398 if "id" in item:
399 server_ids.add(item["id"])
400 elif "vector_id" in item:
401 server_ids.add(item["vector_id"])
402 else:
403 print(f"Warning: Unknown server item structure: {item}")
404 else:
405 print(f"Warning: Server item is not a dict: {item}")
406
407 if not bf_ids:
408 continue # Skip if brute force found no results
409
410 intersection = bf_ids.intersection(server_ids)
411 recall = len(intersection) / len(bf_ids)
412 recalls.append(recall)
413
414 # Debug: Print recall for first few queries
415 if i < 3:
416 print(f"BF IDs (first 5): {list(bf_ids)[:5]}")
417 print(f"Server IDs (first 5): {list(server_ids)[:5]}")
418 print(f"Intersection: {len(intersection)}")
419 print(f"Recall: {recall:.2%}")
420
421 avg_recall = sum(recalls) / len(recalls) if recalls else 0
422 return avg_recall, recalls
423
424
425def run_rps_tests(

Callers 1

mainFunction · 0.70

Calls 2

addMethod · 0.80
appendMethod · 0.45

Tested by

no test coverage detected