Team Ai
Apppublic

aphilippov/python-server-api

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
qdrant_retrieve_user_proxy_agent.py301 linesDownload Raw Back to contrib
1from typing import Callable, Dict, List, Optional2 3from autogen.agentchat.contrib.retrieve_user_proxy_agent import RetrieveUserProxyAgent4from autogen.retrieve_utils import get_files_from_dir, split_files_to_chunks, TEXT_FORMATS5import logging6 7logger = logging.getLogger(__name__)8 9try:10    from qdrant_client import QdrantClient, models11    from qdrant_client.fastembed_common import QueryResponse12    import fastembed13except ImportError as e:14    logging.fatal("Failed to import qdrant_client with fastembed. Try running 'pip install qdrant_client[fastembed]'")15    raise e16 17 18class QdrantRetrieveUserProxyAgent(RetrieveUserProxyAgent):19    def __init__(20        self,21        name="RetrieveChatAgent",  # default set to RetrieveChatAgent22        human_input_mode: Optional[str] = "ALWAYS",23        is_termination_msg: Optional[Callable[[Dict], bool]] = None,24        retrieve_config: Optional[Dict] = None,  # config for the retrieve agent25        **kwargs,26    ):27        """28        Args:29            name (str): name of the agent.30            human_input_mode (str): whether to ask for human inputs every time a message is received.31                Possible values are "ALWAYS", "TERMINATE", "NEVER".32                (1) When "ALWAYS", the agent prompts for human input every time a message is received.33                    Under this mode, the conversation stops when the human input is "exit",34                    or when is_termination_msg is True and there is no human input.35                (2) When "TERMINATE", the agent only prompts for human input only when a termination message is received or36                    the number of auto reply reaches the max_consecutive_auto_reply.37                (3) When "NEVER", the agent will never prompt for human input. Under this mode, the conversation stops38                    when the number of auto reply reaches the max_consecutive_auto_reply or when is_termination_msg is True.39            is_termination_msg (function): a function that takes a message in the form of a dictionary40                and returns a boolean value indicating if this received message is a termination message.41                The dict can contain the following keys: "content", "role", "name", "function_call".42            retrieve_config (dict or None): config for the retrieve agent.43                To use default config, set to None. Otherwise, set to a dictionary with the following keys:44                - task (Optional, str): the task of the retrieve chat. Possible values are "code", "qa" and "default". System45                    prompt will be different for different tasks. The default value is `default`, which supports both code and qa.46                - client (Optional, qdrant_client.QdrantClient(":memory:")): A QdrantClient instance. If not provided, an in-memory instance will be assigned. Not recommended for production.47                    will be used. If you want to use other vector db, extend this class and override the `retrieve_docs` function.48                - docs_path (Optional, Union[str, List[str]]): the path to the docs directory. It can also be the path to a single file,49                    the url to a single file or a list of directories, files and urls. Default is None, which works only if the collection is already created.50                - extra_docs (Optional, bool): when true, allows adding documents with unique IDs without overwriting existing ones; when false, it replaces existing documents using default IDs, risking collection overwrite.,51                    when set to true it enables the system to assign unique IDs starting from "length+i" for new document chunks, preventing the replacement of existing documents and facilitating the addition of more content to the collection..52                    By default, "extra_docs" is set to false, starting document IDs from zero. This poses a risk as new documents might overwrite existing ones, potentially causing unintended loss or alteration of data in the collection.53                - collection_name (Optional, str): the name of the collection.54                    If key not provided, a default name `autogen-docs` will be used.55                - model (Optional, str): the model to use for the retrieve chat.56                    If key not provided, a default model `gpt-4` will be used.57                - chunk_token_size (Optional, int): the chunk token size for the retrieve chat.58                    If key not provided, a default size `max_tokens * 0.4` will be used.59                - context_max_tokens (Optional, int): the context max token size for the retrieve chat.60                    If key not provided, a default size `max_tokens * 0.8` will be used.61                - chunk_mode (Optional, str): the chunk mode for the retrieve chat. Possible values are62                    "multi_lines" and "one_line". If key not provided, a default mode `multi_lines` will be used.63                - must_break_at_empty_line (Optional, bool): chunk will only break at empty line if True. Default is True.64                    If chunk_mode is "one_line", this parameter will be ignored.65                - embedding_model (Optional, str): the embedding model to use for the retrieve chat.66                    If key not provided, a default model `BAAI/bge-small-en-v1.5` will be used. All available models67                    can be found at `https://qdrant.github.io/fastembed/examples/Supported_Models/`.68                - customized_prompt (Optional, str): the customized prompt for the retrieve chat. Default is None.69                - customized_answer_prefix (Optional, str): the customized answer prefix for the retrieve chat. Default is "".70                    If not "" and the customized_answer_prefix is not in the answer, `Update Context` will be triggered.71                - update_context (Optional, bool): if False, will not apply `Update Context` for interactive retrieval. Default is True.72                - custom_token_count_function (Optional, Callable): a custom function to count the number of tokens in a string.73                    The function should take a string as input and return three integers (token_count, tokens_per_message, tokens_per_name).74                    Default is None, tiktoken will be used and may not be accurate for non-OpenAI models.75                - custom_text_split_function (Optional, Callable): a custom function to split a string into a list of strings.76                    Default is None, will use the default function in `autogen.retrieve_utils.split_text_to_chunks`.77                - custom_text_types (Optional, List[str]): a list of file types to be processed. Default is `autogen.retrieve_utils.TEXT_FORMATS`.78                    This only applies to files under the directories in `docs_path`. Explicitly included files and urls will be chunked regardless of their types.79                - recursive (Optional, bool): whether to search documents recursively in the docs_path. Default is True.80                - parallel (Optional, int): How many parallel workers to use for embedding. Defaults to the number of CPU cores.81                - on_disk (Optional, bool): Whether to store the collection on disk. Default is False.82                - quantization_config: Quantization configuration. If None, quantization will be disabled.83                - hnsw_config: HNSW configuration. If None, default configuration will be used.84                  You can find more info about the hnsw configuration options at https://qdrant.tech/documentation/concepts/indexing/#vector-index.85                  API Reference: https://qdrant.github.io/qdrant/redoc/index.html#tag/collections/operation/create_collection86                - payload_indexing: Whether to create a payload index for the document field. Default is False.87                  You can find more info about the payload indexing options at https://qdrant.tech/documentation/concepts/indexing/#payload-index88                  API Reference: https://qdrant.github.io/qdrant/redoc/index.html#tag/collections/operation/create_field_index89             **kwargs (dict): other kwargs in [UserProxyAgent](../user_proxy_agent#__init__).90 91        """92        super().__init__(name, human_input_mode, is_termination_msg, retrieve_config, **kwargs)93        self._client = self._retrieve_config.get("client", QdrantClient(":memory:"))94        self._embedding_model = self._retrieve_config.get("embedding_model", "BAAI/bge-small-en-v1.5")95        # Uses all available CPU cores to encode data when set to 096        self._parallel = self._retrieve_config.get("parallel", 0)97        self._on_disk = self._retrieve_config.get("on_disk", False)98        self._quantization_config = self._retrieve_config.get("quantization_config", None)99        self._hnsw_config = self._retrieve_config.get("hnsw_config", None)100        self._payload_indexing = self._retrieve_config.get("payload_indexing", False)101 102    def retrieve_docs(self, problem: str, n_results: int = 20, search_string: str = ""):103        """104        Args:105            problem (str): the problem to be solved.106            n_results (int): the number of results to be retrieved. Default is 20.107            search_string (str): only docs that contain an exact match of this string will be retrieved. Default is "".108        """109        if not self._collection:110            print("Trying to create collection.")111            create_qdrant_from_dir(112                dir_path=self._docs_path,113                max_tokens=self._chunk_token_size,114                client=self._client,115                collection_name=self._collection_name,116                chunk_mode=self._chunk_mode,117                must_break_at_empty_line=self._must_break_at_empty_line,118                embedding_model=self._embedding_model,119                custom_text_split_function=self.custom_text_split_function,120                custom_text_types=self._custom_text_types,121                recursive=self._recursive,122                extra_docs=self._extra_docs,123                parallel=self._parallel,124                on_disk=self._on_disk,125                quantization_config=self._quantization_config,126                hnsw_config=self._hnsw_config,127                payload_indexing=self._payload_indexing,128            )129            self._collection = True130 131        results = query_qdrant(132            query_texts=problem,133            n_results=n_results,134            search_string=search_string,135            client=self._client,136            collection_name=self._collection_name,137            embedding_model=self._embedding_model,138        )139        self._results = results140 141 142def create_qdrant_from_dir(143    dir_path: str,144    max_tokens: int = 4000,145    client: QdrantClient = None,146    collection_name: str = "all-my-documents",147    chunk_mode: str = "multi_lines",148    must_break_at_empty_line: bool = True,149    embedding_model: str = "BAAI/bge-small-en-v1.5",150    custom_text_split_function: Callable = None,151    custom_text_types: List[str] = TEXT_FORMATS,152    recursive: bool = True,153    extra_docs: bool = False,154    parallel: int = 0,155    on_disk: bool = False,156    quantization_config: Optional[models.QuantizationConfig] = None,157    hnsw_config: Optional[models.HnswConfigDiff] = None,158    payload_indexing: bool = False,159    qdrant_client_options: Optional[Dict] = {},160):161    """Create a Qdrant collection from all the files in a given directory, the directory can also be a single file or a162      url to a single file.163 164    Args:165        dir_path (str): the path to the directory, file or url.166        max_tokens (Optional, int): the maximum number of tokens per chunk. Default is 4000.167        client (Optional, QdrantClient): the QdrantClient instance. Default is None.168        collection_name (Optional, str): the name of the collection. Default is "all-my-documents".169        chunk_mode (Optional, str): the chunk mode. Default is "multi_lines".170        must_break_at_empty_line (Optional, bool): Whether to break at empty line. Default is True.171        embedding_model (Optional, str): the embedding model to use. Default is "BAAI/bge-small-en-v1.5".172            The list of all the available models can be at https://qdrant.github.io/fastembed/examples/Supported_Models/.173        custom_text_split_function (Optional, Callable): a custom function to split a string into a list of strings.174            Default is None, will use the default function in `autogen.retrieve_utils.split_text_to_chunks`.175        custom_text_types (Optional, List[str]): a list of file types to be processed. Default is TEXT_FORMATS.176        recursive (Optional, bool): whether to search documents recursively in the dir_path. Default is True.177        extra_docs (Optional, bool): whether to add more documents in the collection. Default is False178        parallel (Optional, int): How many parallel workers to use for embedding. Defaults to the number of CPU cores179        on_disk (Optional, bool): Whether to store the collection on disk. Default is False.180        quantization_config: Quantization configuration. If None, quantization will be disabled.181            Ref: https://qdrant.github.io/qdrant/redoc/index.html#tag/collections/operation/create_collection182        hnsw_config: HNSW configuration. If None, default configuration will be used.183            Ref: https://qdrant.github.io/qdrant/redoc/index.html#tag/collections/operation/create_collection184        payload_indexing: Whether to create a payload index for the document field. Default is False.185        qdrant_client_options: (Optional, dict): the options for instantiating the qdrant client.186            Ref: https://github.com/qdrant/qdrant-client/blob/master/qdrant_client/qdrant_client.py#L36-L58.187    """188    if client is None:189        client = QdrantClient(**qdrant_client_options)190        client.set_model(embedding_model)191 192    if custom_text_split_function is not None:193        chunks = split_files_to_chunks(194            get_files_from_dir(dir_path, custom_text_types, recursive),195            custom_text_split_function=custom_text_split_function,196        )197    else:198        chunks = split_files_to_chunks(199            get_files_from_dir(dir_path, custom_text_types, recursive), max_tokens, chunk_mode, must_break_at_empty_line200        )201    logger.info(f"Found {len(chunks)} chunks.")202 203    collection = None204    # Check if collection by same name exists, if not, create it with custom options205    try:206        collection = client.get_collection(collection_name=collection_name)207    except Exception:208        client.create_collection(209            collection_name=collection_name,210            vectors_config=client.get_fastembed_vector_params(211                on_disk=on_disk, quantization_config=quantization_config, hnsw_config=hnsw_config212            ),213        )214        collection = client.get_collection(collection_name=collection_name)215 216    length = 0217    if extra_docs:218        length = len(collection.get()["ids"])219 220    # Upsert in batch of 100 or less if the total number of chunks is less than 100221    for i in range(0, len(chunks), min(100, len(chunks))):222        end_idx = i + min(100, len(chunks) - i)223        client.add(224            collection_name,225            documents=chunks[i:end_idx],226            ids=[(j + length) for j in range(i, end_idx)],227            parallel=parallel,228        )229 230    # Create a payload index for the document field231    # Enables highly efficient payload filtering. Reference: https://qdrant.tech/documentation/concepts/indexing/#indexing232    # Creating an index requires additional computational resources and memory.233    # If filtering performance is critical, we can consider creating an index.234    if payload_indexing:235        client.create_payload_index(236            collection_name=collection_name,237            field_name="document",238            field_schema=models.TextIndexParams(239                type="text",240                tokenizer=models.TokenizerType.WORD,241                min_token_len=2,242                max_token_len=15,243            ),244        )245 246 247def query_qdrant(248    query_texts: List[str],249    n_results: int = 10,250    client: QdrantClient = None,251    collection_name: str = "all-my-documents",252    search_string: str = "",253    embedding_model: str = "BAAI/bge-small-en-v1.5",254    qdrant_client_options: Optional[Dict] = {},255) -> List[List[QueryResponse]]:256    """Perform a similarity search with filters on a Qdrant collection257 258    Args:259        query_texts (List[str]): the query texts.260        n_results (Optional, int): the number of results to return. Default is 10.261        client (Optional, API): the QdrantClient instance. A default in-memory client will be instantiated if None.262        collection_name (Optional, str): the name of the collection. Default is "all-my-documents".263        search_string (Optional, str): the search string. Default is "".264        embedding_model (Optional, str): the embedding model to use. Default is "all-MiniLM-L6-v2". Will be ignored if embedding_function is not None.265        qdrant_client_options: (Optional, dict): the options for instantiating the qdrant client. Reference: https://github.com/qdrant/qdrant-client/blob/master/qdrant_client/qdrant_client.py#L36-L58.266 267    Returns:268        List[List[QueryResponse]]: the query result. The format is:269            class QueryResponse(BaseModel, extra="forbid"):  # type: ignore270                id: Union[str, int]271                embedding: Optional[List[float]]272                metadata: Dict[str, Any]273                document: str274                score: float275    """276    if client is None:277        client = QdrantClient(**qdrant_client_options)278        client.set_model(embedding_model)279 280    results = client.query_batch(281        collection_name,282        query_texts,283        limit=n_results,284        query_filter=models.Filter(285            must=[286                models.FieldCondition(287                    key="document",288                    match=models.MatchText(text=search_string),289                )290            ]291        )292        if search_string293        else None,294    )295 296    data = {297        "ids": [[result.id for result in sublist] for sublist in results],298        "documents": [[result.document for result in sublist] for sublist in results],299    }300    return data301