11import heapq
22import logging
33from abc import ABC , abstractmethod
4- from collections import defaultdict
5- from typing import Any , DefaultDict , Optional , Sequence , cast
4+ from typing import Any
65
76import numpy
8- from chromadb .api .types import QueryResult
97
8+ from vectorcode .chunking import Chunk
109from vectorcode .cli_utils import Config , QueryInclude
10+ from vectorcode .subcommands .query .types import QueryResult
1111
1212logger = logging .getLogger (name = __name__ )
1313
@@ -29,7 +29,7 @@ def __init__(self, configs: Config, **kwargs: Any):
2929 "'configs' should contain the query messages."
3030 )
3131 self .n_result = configs .n_result
32- self ._raw_results : Optional [QueryResult ] = None
32+ self ._raw_results : list [QueryResult ] = []
3333
3434 @classmethod
3535 def create (cls , configs : Config , ** kwargs : Any ):
@@ -46,53 +46,31 @@ def create(cls, configs: Config, **kwargs: Any):
4646 raise
4747
4848 @abstractmethod
49- async def compute_similarity (
50- self , results : list [str ], query_message : str
51- ) -> Sequence [float ]: # pragma: nocover
52- """Given a list of n results and 1 query message,
53- return a list-like object of length n that contains the similarity scores between
54- each item in `results` and the `query_message`.
55-
56- A high similarity score means the strings are semantically similar to each other.
57- `query_message` will be loaded in the same order as they appear in `self.configs.query`.
58-
59- If you need the raw query results from chromadb,
60- it'll be saved in `self._raw_results` before this method is called.
49+ async def compute_similarity (self , results : list [QueryResult ]): # pragma: nocover
50+ """
51+ Modify the `QueryResult.scores` field IN-PLACE so that they contain the correct scores.
6152 """
6253 raise NotImplementedError
6354
64- async def rerank (self , results : QueryResult | dict ) -> list [str ]:
65- if len (results [ "ids" ] ) == 0 or all ( len ( i ) == 0 for i in results [ "ids" ]) :
55+ async def rerank (self , results : list [ QueryResult ] ) -> list [str ]:
56+ if len (results ) == 0 :
6657 return []
58+ results = await self .compute_similarity (results )
6759
68- self ._raw_results = cast (QueryResult , results )
69- query_chunks = self .configs .query
70- assert query_chunks
71- assert results ["metadatas" ] is not None
72- assert results ["documents" ] is not None
73- documents : DefaultDict [str , list [float ]] = defaultdict (list )
74- for query_chunk_idx in range (len (query_chunks )):
75- chunk_ids = results ["ids" ][query_chunk_idx ]
76- chunk_metas = results ["metadatas" ][query_chunk_idx ]
77- chunk_docs = results ["documents" ][query_chunk_idx ]
78- scores = await self .compute_similarity (
79- chunk_docs , query_chunks [query_chunk_idx ]
80- )
81- for i , score in enumerate (scores ):
82- if QueryInclude .chunk in self .configs .include :
83- documents [chunk_ids [i ]].append (float (score ))
84- else :
85- documents [str (chunk_metas [i ]["path" ])].append (float (score ))
60+ group_by = "path"
61+ if QueryInclude .chunk in self .configs .include :
62+ group_by = "chunk"
63+ grouped_results = QueryResult .group (* results , key = group_by , top_k = "auto" )
8664
87- logger .debug ("Document scores: %s" , documents )
88- top_k = int (numpy .mean (tuple (len (i ) for i in documents .values ())))
89- for key in documents .keys ():
90- documents [key ] = heapq .nlargest (top_k , documents [key ])
91-
92- self ._raw_results = None
65+ scores : dict [Chunk | str , float ] = {}
66+ for key in grouped_results .keys ():
67+ scores [key ] = float (
68+ numpy .mean (tuple (i .mean_score () for i in grouped_results [key ]))
69+ )
9370
94- return heapq .nlargest (
95- self .n_result ,
96- documents .keys (),
97- key = lambda x : float (numpy .mean (documents [x ])),
71+ return list (
72+ str (i )
73+ for i in heapq .nlargest (
74+ self .configs .n_result , grouped_results .keys (), key = lambda x : scores [x ]
75+ )
9876 )
0 commit comments