From e403a30952801c6ded79dccf5920c36a6408e7f8 Mon Sep 17 00:00:00 2001 From: Lucas Koontz Date: Thu, 4 Sep 2025 23:55:34 -0700 Subject: [PATCH 1/5] Configures LLM reranker with adjustable parameters Allows configuration of the LLM reranker with parameters like `n`, `logprobs`, `top_logprobs`, and `max_tokens` to control completion generation. These parameters are configurable via the `RerankerConfig` and environment variables, providing flexibility in controlling the reranking process. Also, this change improves logprobs handling to better find and use class tokens. --- .../utilities/rag/rerankers/base_reranker.py | 64 +++++- .../integrations/utilities/rag/settings.py | 186 +++++------------- mindsdb/utilities/config.py | 40 ++++ 3 files changed, 150 insertions(+), 140 deletions(-) diff --git a/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py b/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py index 5bf826789ed..c6b5c105e7d 100644 --- a/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py +++ b/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py @@ -13,7 +13,14 @@ from openai import AsyncOpenAI, AsyncAzureOpenAI from pydantic import BaseModel -from mindsdb.integrations.utilities.rag.settings import DEFAULT_RERANKING_MODEL, DEFAULT_LLM_ENDPOINT +from mindsdb.integrations.utilities.rag.settings import ( + DEFAULT_RERANKING_MODEL, + DEFAULT_LLM_ENDPOINT, + DEFAULT_RERANKER_N, + DEFAULT_RERANKER_LOGPROBS, + DEFAULT_RERANKER_TOP_LOGPROBS, + DEFAULT_RERANKER_MAX_TOKENS, +) from mindsdb.integrations.libs.base import BaseMLEngine log = logging.getLogger(__name__) @@ -38,6 +45,10 @@ class BaseLLMReranker(BaseModel, ABC): request_timeout: float = 20.0 # Timeout for API requests early_stop: bool = True # Whether to enable early stopping early_stop_threshold: float = 0.8 # Confidence threshold for early stopping + n: int = DEFAULT_RERANKER_N # Number of completions to generate + logprobs: bool = DEFAULT_RERANKER_LOGPROBS # Whether to include log probabilities + top_logprobs: int = DEFAULT_RERANKER_TOP_LOGPROBS # Number of top log probabilities to include + max_tokens: int = DEFAULT_RERANKER_MAX_TOKENS # Maximum tokens to generate class Config: arbitrary_types_allowed = True @@ -234,6 +245,28 @@ async def search_relevancy_no_logprob(self, query: str, document: str) -> Any: return rerank_data async def search_relevancy_score(self, query: str, document: str) -> Any: + """ + This method is used to score the relevance of a document to a query. + + Args: + query: The query to score the relevance of. + document: The document to score the relevance of. + + Returns: + A dictionary with the document and the relevance score. + """ + + log.debug("Start search_relevancy_score") + + log.debug(f"Reranker query: {query[:5]}") + log.debug(f"Reranker document: {document[:50]}") + log.debug(f"Reranker model: {self.model}") + log.debug(f"Reranker temperature: {self.temperature}") + log.debug(f"Reranker n: {self.n}") + log.debug(f"Reranker logprobs: {self.logprobs}") + log.debug(f"Reranker top_logprobs: {self.top_logprobs}") + log.debug(f"Reranker max_tokens: {self.max_tokens}") + response = await self.client.chat.completions.create( model=self.model, messages=[ @@ -306,17 +339,30 @@ async def search_relevancy_score(self, query: str, document: str) -> Any: }, ], temperature=self.temperature, - n=1, - logprobs=True, - top_logprobs=4, - max_tokens=3, + n=self.n, + logprobs=self.logprobs, + top_logprobs=self.top_logprobs, + max_tokens=self.max_tokens, ) # Extract response and logprobs token_logprobs = response.choices[0].logprobs.content - # Reconstruct the prediction and extract the top logprobs from the final token (e.g., "1") - final_token_logprob = token_logprobs[-1] - top_logprobs = final_token_logprob.top_logprobs + + # Find the token that contains the class number (1, 2, 3, or 4) + # Instead of just taking the last token, search for the actual class number token + class_token_logprob = None + for token_logprob in reversed(token_logprobs): + if token_logprob.token in ["1", "2", "3", "4"]: + class_token_logprob = token_logprob + break + + # If we couldn't find a class token, fall back to the last non-empty token + if class_token_logprob is None: + log.warning("No class token logprob found, using the last token as fallback") + class_token_logprob = token_logprobs[-1] + + top_logprobs = class_token_logprob.top_logprobs + # Create a map of 'class_1' -> probability, using token combinations class_probs = {} for top_token in top_logprobs: @@ -337,6 +383,8 @@ async def search_relevancy_score(self, query: str, document: str) -> Any: score = 0.0 rerank_data = {"document": document, "relevance_score": score} + log.debug(f"Reranker score: {score}") + log.debug("End search_relevancy_score") return rerank_data def get_scores(self, query: str, documents: list[str]): diff --git a/mindsdb/integrations/utilities/rag/settings.py b/mindsdb/integrations/utilities/rag/settings.py index 04824ef94e5..210a81a7996 100644 --- a/mindsdb/integrations/utilities/rag/settings.py +++ b/mindsdb/integrations/utilities/rag/settings.py @@ -32,6 +32,10 @@ DEFAULT_RERANKER_FLAG = False DEFAULT_RERANKING_MODEL = "gpt-4o" DEFAULT_LLM_ENDPOINT = "https://api.openai.com/v1" +DEFAULT_RERANKER_N = 1 +DEFAULT_RERANKER_LOGPROBS = True +DEFAULT_RERANKER_TOP_LOGPROBS = 4 +DEFAULT_RERANKER_MAX_TOKENS = 3 DEFAULT_AUTO_META_PROMPT_TEMPLATE = """ Below is a json representation of a table with information about {description}. Return a JSON list with an entry for each column. Each entry should have @@ -366,9 +370,7 @@ class LLMConfig(BaseModel): - model_name: str = Field( - default=DEFAULT_LLM_MODEL, description="LLM model to use for generation" - ) + model_name: str = Field(default=DEFAULT_LLM_MODEL, description="LLM model to use for generation") provider: str = Field( default=DEFAULT_LLM_MODEL_PROVIDER, description="LLM model provider to use for generation", @@ -430,9 +432,7 @@ class SearchType(Enum): class SearchKwargs(BaseModel): k: int = Field(default=DEFAULT_K, description="Amount of documents to return", ge=1) - filter: Optional[Dict[str, Any]] = Field( - default=None, description="Filter by document metadata" - ) + filter: Optional[Dict[str, Any]] = Field(default=None, description="Filter by document metadata") # For similarity_score_threshold search type score_threshold: Optional[float] = Field( default=None, @@ -441,9 +441,7 @@ class SearchKwargs(BaseModel): le=1.0, ) # For MMR search type - fetch_k: Optional[int] = Field( - default=None, description="Amount of documents to pass to MMR algorithm", ge=1 - ) + fetch_k: Optional[int] = Field(default=None, description="Amount of documents to pass to MMR algorithm", ge=1) lambda_mult: Optional[float] = Field( default=None, description="Diversity of results returned by MMR (1=min diversity, 0=max)", @@ -459,9 +457,7 @@ def model_dump(self, *args, **kwargs): class LLMExample(BaseModel): input: str = Field(description="User input for the example") - output: str = Field( - description="What the LLM should generate for this example's input" - ) + output: str = Field(description="What the LLM should generate for this example's input") class ValueSchema(BaseModel): @@ -502,41 +498,25 @@ class ValueSchema(BaseModel): class MetadataConfig(BaseModel): """Class to configure metadata for retrieval. Only supports very basic document name lookup at the moment.""" - table: str = Field( - description="Source table for metadata." - ) + + table: str = Field(description="Source table for metadata.") max_document_context: int = Field( # To work well with models with context window of 32768. default=16384, - description="Truncate a document before using as context with an LLM if it exceeds this amount of tokens" - ) - embeddings_table: str = Field( - default="embeddings", - description="Source table for embeddings" - ) - id_column: str = Field( - default="Id", - description="Name of ID column in metadata table" - ) - name_column: str = Field( - default="Title", - description="Name of column containing name or title of document" - ) - name_column_index: Optional[str] = Field( - default=None, - description="Name of GIN index to use when looking up name." + description="Truncate a document before using as context with an LLM if it exceeds this amount of tokens", ) + embeddings_table: str = Field(default="embeddings", description="Source table for embeddings") + id_column: str = Field(default="Id", description="Name of ID column in metadata table") + name_column: str = Field(default="Title", description="Name of column containing name or title of document") + name_column_index: Optional[str] = Field(default=None, description="Name of GIN index to use when looking up name.") content_column: str = Field( - default="content", - description="Name of column in embeddings table containing chunk content" + default="content", description="Name of column in embeddings table containing chunk content" ) embeddings_metadata_column: str = Field( - default="metadata", - description="Name of column in embeddings table containing chunk metadata" + default="metadata", description="Name of column in embeddings table containing chunk metadata" ) doc_id_key: str = Field( - default="original_row_id", - description="Metadata field that links an embedded chunk back to source document ID" + default="original_row_id", description="Metadata field that links an embedded chunk back to source document ID" ) @@ -552,14 +532,12 @@ class ColumnSchema(BaseModel): ] ] = Field( default=None, - description="One of the following. A dict or ordered dict of {schema_value: ValueSchema, ...}, where schema value is the name given for this value description in the schema." + description="One of the following. A dict or ordered dict of {schema_value: ValueSchema, ...}, where schema value is the name given for this value description in the schema.", ) example_questions: Optional[List[LLMExample]] = Field( default=None, description="Example questions where this table is useful." ) - max_filters: Optional[int] = Field( - default=1, description="Maximum number of filters to generate for this column." - ) + max_filters: Optional[int] = Field(default=1, description="Maximum number of filters to generate for this column.") filter_threshold: Optional[float] = Field( default=0.0, description="Minimum relevance threshold to include metadata filters from this column.", @@ -578,9 +556,7 @@ class TableSchema(BaseModel): table: str = Field(description="Name of table in the database") description: str = Field(description="Description of what the table represents") usage: str = Field(description="How and when to use this Table for search.") - columns: Optional[ - Union[OrderedDict[str, ColumnSchema], Dict[str, ColumnSchema]] - ] = Field( + columns: Optional[Union[OrderedDict[str, ColumnSchema], Dict[str, ColumnSchema]]] = Field( description="Dict or Ordered Dict of {column_name: ColumnSchemas} describing the metadata columns available for the table" ) example_questions: Optional[List[LLMExample]] = Field( @@ -590,9 +566,7 @@ class TableSchema(BaseModel): description="SQL join string to join this table with source documents table", default="", ) - max_filters: Optional[int] = Field( - default=1, description="Maximum number of filters to generate for this table." - ) + max_filters: Optional[int] = Field(default=1, description="Maximum number of filters to generate for this table.") filter_threshold: Optional[float] = Field( default=0.0, description="Minimum relevance required to use this table to generate filters.", @@ -675,12 +649,8 @@ class SQLRetrieverConfig(BaseModel): source_table: str = Field( description="Name of the source table containing the original documents that were embedded" ) - source_id_column: str = Field( - description="Name of the column containing the UUID.", default="Id" - ) - max_filters: Optional[int] = Field( - description="Maximum number of filters to generate for sql queries.", default=10 - ) + source_id_column: str = Field(description="Name of the column containing the UUID.", default="Id") + max_filters: Optional[int] = Field(description="Maximum number of filters to generate for sql queries.", default=10) filter_threshold: Optional[float] = Field( description="Minimum relevance required to use this Database to generate filters.", default=0.0, @@ -728,6 +698,10 @@ class RerankerConfig(BaseModel): retry_delay: float = 1.0 early_stop: bool = True # Whether to enable early stopping early_stop_threshold: float = 0.8 # Confidence threshold for early stopping + n: int = DEFAULT_RERANKER_N # Number of completions to generate + logprobs: bool = DEFAULT_RERANKER_LOGPROBS # Whether to include log probabilities + top_logprobs: int = DEFAULT_RERANKER_TOP_LOGPROBS # Number of top log probabilities to include + max_tokens: int = DEFAULT_RERANKER_MAX_TOKENS # Maximum tokens to generate class MultiHopRetrieverConfig(BaseModel): @@ -737,9 +711,7 @@ class MultiHopRetrieverConfig(BaseModel): default=RetrieverType.VECTOR_STORE, description="Type of base retriever to use for multi-hop retrieval", ) - max_hops: int = Field( - default=3, description="Maximum number of follow-up questions to generate", ge=1 - ) + max_hops: int = Field(default=3, description="Maximum number of follow-up questions to generate", ge=1) reformulation_template: str = Field( default=DEFAULT_QUESTION_REFORMULATION_TEMPLATE, description="Template for reformulating questions", @@ -751,48 +723,29 @@ class MultiHopRetrieverConfig(BaseModel): class RAGPipelineModel(BaseModel): - documents: Optional[List[Document]] = Field( - default=None, description="List of documents" - ) + documents: Optional[List[Document]] = Field(default=None, description="List of documents") vector_store_config: VectorStoreConfig = Field( default_factory=VectorStoreConfig, description="Vector store configuration" ) llm: Optional[BaseChatModel] = Field(default=None, description="Language model") - llm_model_name: str = Field( - default=DEFAULT_LLM_MODEL, description="Language model name" - ) - llm_provider: Optional[str] = Field( - default=None, description="Language model provider" - ) + llm_model_name: str = Field(default=DEFAULT_LLM_MODEL, description="Language model name") + llm_provider: Optional[str] = Field(default=None, description="Language model provider") vector_store: VectorStore = Field( default_factory=lambda: vector_store_map[VectorStoreConfig().vector_store_type], description="Vector store", ) - db_connection_string: Optional[str] = Field( - default=None, description="Database connection string" - ) + db_connection_string: Optional[str] = Field(default=None, description="Database connection string") metadata_config: Optional[MetadataConfig] = Field( - default=None, - description="Configuration for metadata to be used for retrieval" + default=None, description="Configuration for metadata to be used for retrieval" ) table_name: str = Field(default=DEFAULT_TEST_TABLE_NAME, description="Table name") - embedding_model: Optional[Embeddings] = Field( - default=None, description="Embedding model" - ) - rag_prompt_template: str = Field( - default=DEFAULT_RAG_PROMPT_TEMPLATE, description="RAG prompt template" - ) - retriever_prompt_template: Optional[Union[str, dict]] = Field( - default=None, description="Retriever prompt template" - ) - retriever_type: RetrieverType = Field( - default=RetrieverType.VECTOR_STORE, description="Retriever type" - ) - search_type: SearchType = Field( - default=SearchType.SIMILARITY, description="Type of search to perform" - ) + embedding_model: Optional[Embeddings] = Field(default=None, description="Embedding model") + rag_prompt_template: str = Field(default=DEFAULT_RAG_PROMPT_TEMPLATE, description="RAG prompt template") + retriever_prompt_template: Optional[Union[str, dict]] = Field(default=None, description="Retriever prompt template") + retriever_type: RetrieverType = Field(default=RetrieverType.VECTOR_STORE, description="Retriever type") + search_type: SearchType = Field(default=SearchType.SIMILARITY, description="Type of search to perform") search_kwargs: SearchKwargs = Field( default_factory=SearchKwargs, description="Search configuration for the retriever", @@ -811,39 +764,23 @@ class RAGPipelineModel(BaseModel): multi_retriever_mode: MultiVectorRetrieverMode = Field( default=MultiVectorRetrieverMode.BOTH, description="Multi retriever mode" ) - max_concurrency: int = Field( - default=DEFAULT_MAX_CONCURRENCY, description="Maximum concurrency" - ) + max_concurrency: int = Field(default=DEFAULT_MAX_CONCURRENCY, description="Maximum concurrency") id_key: int = Field(default=DEFAULT_ID_KEY, description="ID key") parent_store: Optional[BaseStore] = Field(default=None, description="Parent store") - text_splitter: Optional[TextSplitter] = Field( - default=None, description="Text splitter" - ) + text_splitter: Optional[TextSplitter] = Field(default=None, description="Text splitter") chunk_size: int = Field(default=DEFAULT_CHUNK_SIZE, description="Chunk size") - chunk_overlap: int = Field( - default=DEFAULT_CHUNK_OVERLAP, description="Chunk overlap" - ) + chunk_overlap: int = Field(default=DEFAULT_CHUNK_OVERLAP, description="Chunk overlap") # Auto retriever specific - auto_retriever_filter_columns: Optional[List[str]] = Field( - default=None, description="Filter columns" - ) - cardinality_threshold: int = Field( - default=DEFAULT_CARDINALITY_THRESHOLD, description="Cardinality threshold" - ) + auto_retriever_filter_columns: Optional[List[str]] = Field(default=None, description="Filter columns") + cardinality_threshold: int = Field(default=DEFAULT_CARDINALITY_THRESHOLD, description="Cardinality threshold") content_column_name: str = Field( default=DEFAULT_CONTENT_COLUMN_NAME, description="Content column name (the column we will get embeddings)", ) - dataset_description: str = Field( - default=DEFAULT_DATASET_DESCRIPTION, description="Description of the dataset" - ) - reranker: bool = Field( - default=DEFAULT_RERANKER_FLAG, description="Whether to use reranker" - ) - reranker_config: RerankerConfig = Field( - default_factory=RerankerConfig, description="Reranker configuration" - ) + dataset_description: str = Field(default=DEFAULT_DATASET_DESCRIPTION, description="Description of the dataset") + reranker: bool = Field(default=DEFAULT_RERANKER_FLAG, description="Whether to use reranker") + reranker_config: RerankerConfig = Field(default_factory=RerankerConfig, description="Reranker configuration") multi_hop_config: Optional[MultiHopRetrieverConfig] = Field( default=None, @@ -856,9 +793,7 @@ def validate_multi_hop_config(cls, v: Optional[MultiHopRetrieverConfig], info): """Validate that multi_hop_config is set when using multi-hop retrieval.""" values = info.data if values.get("retriever_type") == RetrieverType.MULTI_HOP and v is None: - raise ValueError( - "multi_hop_config must be set when using multi-hop retrieval" - ) + raise ValueError("multi_hop_config must be set when using multi-hop retrieval") return v class Config: @@ -889,13 +824,9 @@ def validate_search_kwargs(cls, v: SearchKwargs, info) -> SearchKwargs: if v.lambda_mult is not None and (v.lambda_mult < 0 or v.lambda_mult > 1): raise ValueError("lambda_mult must be between 0 and 1") if v.fetch_k is None and v.lambda_mult is not None: - raise ValueError( - "fetch_k is required when using lambda_mult with MMR search type" - ) + raise ValueError("fetch_k is required when using lambda_mult with MMR search type") if v.lambda_mult is None and v.fetch_k is not None: - raise ValueError( - "lambda_mult is required when using fetch_k with MMR search type" - ) + raise ValueError("lambda_mult is required when using fetch_k with MMR search type") elif search_type != SearchType.MMR: if v.fetch_k is not None: raise ValueError("fetch_k is only valid for MMR search type") @@ -904,20 +835,11 @@ def validate_search_kwargs(cls, v: SearchKwargs, info) -> SearchKwargs: # Validate similarity_score_threshold parameters if search_type == SearchType.SIMILARITY_SCORE_THRESHOLD: - if v.score_threshold is not None and ( - v.score_threshold < 0 or v.score_threshold > 1 - ): + if v.score_threshold is not None and (v.score_threshold < 0 or v.score_threshold > 1): raise ValueError("score_threshold must be between 0 and 1") if v.score_threshold is None: - raise ValueError( - "score_threshold is required for similarity_score_threshold search type" - ) - elif ( - search_type != SearchType.SIMILARITY_SCORE_THRESHOLD - and v.score_threshold is not None - ): - raise ValueError( - "score_threshold is only valid for similarity_score_threshold search type" - ) + raise ValueError("score_threshold is required for similarity_score_threshold search type") + elif search_type != SearchType.SIMILARITY_SCORE_THRESHOLD and v.score_threshold is not None: + raise ValueError("score_threshold is only valid for similarity_score_threshold search type") return v diff --git a/mindsdb/utilities/config.py b/mindsdb/utilities/config.py index 4b2c9e1893f..a69f81a4d75 100644 --- a/mindsdb/utilities/config.py +++ b/mindsdb/utilities/config.py @@ -300,6 +300,46 @@ def prepare_env_config(self) -> None: self._env_config["default_reranking_model"] = { "api_key": os.environ["MINDSDB_DEFAULT_RERANKING_MODEL_API_KEY"] } + + # Reranker configuration from environment variables + reranker_config = {} + if os.environ.get("MINDSDB_RERANKER_N", "") != "": + try: + reranker_config["n"] = int(os.environ["MINDSDB_RERANKER_N"]) + except ValueError: + raise ValueError(f"MINDSDB_RERANKER_N must be an integer, got: {os.environ['MINDSDB_RERANKER_N']}") + + if os.environ.get("MINDSDB_RERANKER_LOGPROBS", "") != "": + logprobs_value = os.environ["MINDSDB_RERANKER_LOGPROBS"].lower() + if logprobs_value in ("true", "1", "yes", "y"): + reranker_config["logprobs"] = True + elif logprobs_value in ("false", "0", "no", "n"): + reranker_config["logprobs"] = False + else: + raise ValueError( + f"MINDSDB_RERANKER_LOGPROBS must be a boolean value, got: {os.environ['MINDSDB_RERANKER_LOGPROBS']}" + ) + + if os.environ.get("MINDSDB_RERANKER_TOP_LOGPROBS", "") != "": + try: + reranker_config["top_logprobs"] = int(os.environ["MINDSDB_RERANKER_TOP_LOGPROBS"]) + except ValueError: + raise ValueError( + f"MINDSDB_RERANKER_TOP_LOGPROBS must be an integer, got: {os.environ['MINDSDB_RERANKER_TOP_LOGPROBS']}" + ) + + if os.environ.get("MINDSDB_RERANKER_MAX_TOKENS", "") != "": + try: + reranker_config["max_tokens"] = int(os.environ["MINDSDB_RERANKER_MAX_TOKENS"]) + except ValueError: + raise ValueError( + f"MINDSDB_RERANKER_MAX_TOKENS must be an integer, got: {os.environ['MINDSDB_RERANKER_MAX_TOKENS']}" + ) + + if reranker_config: + if "default_reranking_model" not in self._env_config: + self._env_config["default_reranking_model"] = {} + self._env_config["default_reranking_model"].update(reranker_config) if os.environ.get("MINDSDB_DATA_CATALOG_ENABLED", "").lower() in ("1", "true"): self._env_config["data_catalog"] = {"enabled": True} From 5d8379d079c64ce89771c351f2462f0bade05c2b Mon Sep 17 00:00:00 2001 From: Lucas Koontz Date: Thu, 4 Sep 2025 23:57:34 -0700 Subject: [PATCH 2/5] Adds valid class tokens to reranker config Allows users to configure valid class tokens for the reranker, improving flexibility and adaptability to different models and scenarios. This change introduces a new `valid_class_tokens` parameter to the reranker configuration, enabling users to specify the tokens that the reranker should consider as valid class labels. This allows to customize the tokens searched for within the reranker response. --- .../integrations/utilities/rag/rerankers/base_reranker.py | 7 +++++-- mindsdb/integrations/utilities/rag/settings.py | 2 ++ mindsdb/utilities/config.py | 8 ++++++++ 3 files changed, 15 insertions(+), 2 deletions(-) diff --git a/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py b/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py index c6b5c105e7d..2a12ab3a744 100644 --- a/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py +++ b/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py @@ -20,6 +20,7 @@ DEFAULT_RERANKER_LOGPROBS, DEFAULT_RERANKER_TOP_LOGPROBS, DEFAULT_RERANKER_MAX_TOKENS, + DEFAULT_VALID_CLASS_TOKENS, ) from mindsdb.integrations.libs.base import BaseMLEngine @@ -49,6 +50,7 @@ class BaseLLMReranker(BaseModel, ABC): logprobs: bool = DEFAULT_RERANKER_LOGPROBS # Whether to include log probabilities top_logprobs: int = DEFAULT_RERANKER_TOP_LOGPROBS # Number of top log probabilities to include max_tokens: int = DEFAULT_RERANKER_MAX_TOKENS # Maximum tokens to generate + valid_class_tokens: List[str] = DEFAULT_VALID_CLASS_TOKENS class Config: arbitrary_types_allowed = True @@ -266,6 +268,7 @@ async def search_relevancy_score(self, query: str, document: str) -> Any: log.debug(f"Reranker logprobs: {self.logprobs}") log.debug(f"Reranker top_logprobs: {self.top_logprobs}") log.debug(f"Reranker max_tokens: {self.max_tokens}") + log.debug(f"Reranker valid_class_tokens: {self.valid_class_tokens}") response = await self.client.chat.completions.create( model=self.model, @@ -348,11 +351,11 @@ async def search_relevancy_score(self, query: str, document: str) -> Any: # Extract response and logprobs token_logprobs = response.choices[0].logprobs.content - # Find the token that contains the class number (1, 2, 3, or 4) + # Find the token that contains the class number # Instead of just taking the last token, search for the actual class number token class_token_logprob = None for token_logprob in reversed(token_logprobs): - if token_logprob.token in ["1", "2", "3", "4"]: + if token_logprob.token in self.valid_class_tokens: class_token_logprob = token_logprob break diff --git a/mindsdb/integrations/utilities/rag/settings.py b/mindsdb/integrations/utilities/rag/settings.py index 210a81a7996..f7e2d829ce7 100644 --- a/mindsdb/integrations/utilities/rag/settings.py +++ b/mindsdb/integrations/utilities/rag/settings.py @@ -36,6 +36,7 @@ DEFAULT_RERANKER_LOGPROBS = True DEFAULT_RERANKER_TOP_LOGPROBS = 4 DEFAULT_RERANKER_MAX_TOKENS = 3 +DEFAULT_VALID_CLASS_TOKENS = ["1", "2", "3", "4"] DEFAULT_AUTO_META_PROMPT_TEMPLATE = """ Below is a json representation of a table with information about {description}. Return a JSON list with an entry for each column. Each entry should have @@ -702,6 +703,7 @@ class RerankerConfig(BaseModel): logprobs: bool = DEFAULT_RERANKER_LOGPROBS # Whether to include log probabilities top_logprobs: int = DEFAULT_RERANKER_TOP_LOGPROBS # Number of top log probabilities to include max_tokens: int = DEFAULT_RERANKER_MAX_TOKENS # Maximum tokens to generate + valid_class_tokens: List[str] = DEFAULT_VALID_CLASS_TOKENS # Valid class tokens to look for in the response class MultiHopRetrieverConfig(BaseModel): diff --git a/mindsdb/utilities/config.py b/mindsdb/utilities/config.py index a69f81a4d75..8681d12cd07 100644 --- a/mindsdb/utilities/config.py +++ b/mindsdb/utilities/config.py @@ -336,6 +336,14 @@ def prepare_env_config(self) -> None: f"MINDSDB_RERANKER_MAX_TOKENS must be an integer, got: {os.environ['MINDSDB_RERANKER_MAX_TOKENS']}" ) + if os.environ.get("MINDSDB_RERANKER_VALID_CLASS_TOKENS", "") != "": + try: + reranker_config["valid_class_tokens"] = os.environ["MINDSDB_RERANKER_VALID_CLASS_TOKENS"].split(",") + except ValueError: + raise ValueError( + f"MINDSDB_RERANKER_VALID_CLASS_TOKENS must be a comma-separated list of strings, got: {os.environ['MINDSDB_RERANKER_VALID_CLASS_TOKENS']}" + ) + if reranker_config: if "default_reranking_model" not in self._env_config: self._env_config["default_reranking_model"] = {} From b7576a1cfb889868306b4d39913277388e006faa Mon Sep 17 00:00:00 2001 From: Lucas Koontz Date: Fri, 5 Sep 2025 00:14:56 -0700 Subject: [PATCH 3/5] Cleans up debug logging in reranker Removes an unnecessary newline that was cluttering the debug logs in the reranker. This improves the readability of debug output when troubleshooting reranking issues. --- mindsdb/integrations/utilities/rag/rerankers/base_reranker.py | 1 - 1 file changed, 1 deletion(-) diff --git a/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py b/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py index 2a12ab3a744..32afef7af2a 100644 --- a/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py +++ b/mindsdb/integrations/utilities/rag/rerankers/base_reranker.py @@ -259,7 +259,6 @@ async def search_relevancy_score(self, query: str, document: str) -> Any: """ log.debug("Start search_relevancy_score") - log.debug(f"Reranker query: {query[:5]}") log.debug(f"Reranker document: {document[:50]}") log.debug(f"Reranker model: {self.model}") From bdd52b42294de8f14891b4bde32aa093e2457685 Mon Sep 17 00:00:00 2001 From: Lucas Koontz Date: Fri, 5 Sep 2025 11:40:22 -0700 Subject: [PATCH 4/5] Increases max tokens for re-ranker. Increases the `DEFAULT_RERANKER_MAX_TOKENS` from 3 to 100. This allows the re-ranker to process longer text snippets for improved accuracy. --- mindsdb/integrations/utilities/rag/settings.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mindsdb/integrations/utilities/rag/settings.py b/mindsdb/integrations/utilities/rag/settings.py index f7e2d829ce7..6eab66fe585 100644 --- a/mindsdb/integrations/utilities/rag/settings.py +++ b/mindsdb/integrations/utilities/rag/settings.py @@ -35,7 +35,7 @@ DEFAULT_RERANKER_N = 1 DEFAULT_RERANKER_LOGPROBS = True DEFAULT_RERANKER_TOP_LOGPROBS = 4 -DEFAULT_RERANKER_MAX_TOKENS = 3 +DEFAULT_RERANKER_MAX_TOKENS = 100 DEFAULT_VALID_CLASS_TOKENS = ["1", "2", "3", "4"] DEFAULT_AUTO_META_PROMPT_TEMPLATE = """ Below is a json representation of a table with information about {description}. From 798b385f7382274b3e18c4861663f56267a69703 Mon Sep 17 00:00:00 2001 From: Lucas Koontz Date: Fri, 12 Sep 2025 10:05:06 -0700 Subject: [PATCH 5/5] Update option to skip data (#11533) Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- .../knowledge_bases/insert_data.mdx | 19 +++++++++++++- .../libs/vectordatabase_handler.py | 26 +++++++++++++++++++ .../interfaces/knowledge_base/controller.py | 20 ++++++++++++++ 3 files changed, 64 insertions(+), 1 deletion(-) diff --git a/docs/mindsdb_sql/knowledge_bases/insert_data.mdx b/docs/mindsdb_sql/knowledge_bases/insert_data.mdx index 8bb47bec0a7..ed54c90c7c5 100644 --- a/docs/mindsdb_sql/knowledge_bases/insert_data.mdx +++ b/docs/mindsdb_sql/knowledge_bases/insert_data.mdx @@ -35,14 +35,25 @@ FROM information_schema.queries; -To speed up the initial loading of a knowledge base with data, run the `INSERT INTO` command with the `kb_no_upsert` flag. It is recommended to use this flag only when the knowledge base is empty, that is, only for the initial data insertion. +To speed up data insertion, you can use these performance optimization flags: +**Skip duplicate checking (kb_no_upsert)** ```sql INSERT INTO my_kb SELECT * FROM table_name USING kb_no_upsert = true; ``` +This skips all duplicate checking and directly inserts data. Use only when the knowledge base is empty (initial data load). + +**Skip existing items (kb_skip_existing)** +```sql +INSERT INTO my_kb +SELECT * +FROM table_name +USING kb_skip_existing = true; +``` +This checks for existing items and skips them entirely, including avoiding embedding calculation for existing content. More efficient than upsert when you only want to insert new items. @@ -86,6 +97,12 @@ Knowledge bases uniquely identify data rows using an ID column, which prevents f Ensure the `id_column` uniquely identifies each row to avoid unintentional data loss due to duplicate ID skipping. +**Performance optimization for duplicate handling** + +For better performance when handling duplicates, you can use: +- `kb_skip_existing = true`: Checks for existing IDs and skips them completely (no embedding calculation, more efficient) +- `kb_no_upsert = true`: Skips duplicate checking entirely (fastest, use only for initial load into empty KB) + ### Update Existing Data diff --git a/mindsdb/integrations/libs/vectordatabase_handler.py b/mindsdb/integrations/libs/vectordatabase_handler.py index 64e7cf61dea..ac691729a6e 100644 --- a/mindsdb/integrations/libs/vectordatabase_handler.py +++ b/mindsdb/integrations/libs/vectordatabase_handler.py @@ -593,6 +593,32 @@ def hybrid_search( """ raise NotImplementedError(f"Hybrid search not supported for VectorStoreHandler {self.name}") + def check_existing_ids(self, table_name: str, ids: List[str]) -> List[str]: + """ + Check which IDs from the provided list already exist in the table. + + Args: + table_name (str): Name of the table to check + ids (List[str]): List of IDs to check for existence + + Returns: + List[str]: List of IDs that already exist in the table + """ + if not ids: + return [] + + try: + # Query existing IDs + df_existing = self.select( + table_name, + columns=[TableField.ID.value], + conditions=[FilterCondition(column=TableField.ID.value, op=FilterOperator.IN, value=ids)], + ) + return list(df_existing[TableField.ID.value]) if not df_existing.empty else [] + except Exception: + # If select fails for any reason, return empty list to be safe + return [] + def create_index(self, *args, **kwargs): """ Create an index on the specified table. diff --git a/mindsdb/interfaces/knowledge_base/controller.py b/mindsdb/interfaces/knowledge_base/controller.py index 31cc2ab1534..8ca003904be 100644 --- a/mindsdb/interfaces/knowledge_base/controller.py +++ b/mindsdb/interfaces/knowledge_base/controller.py @@ -56,6 +56,7 @@ class KnowledgeBaseInputParams(BaseModel): content_columns: List[str] | None = None id_column: str | None = None kb_no_upsert: bool = False + kb_skip_existing: bool = False embedding_model: Dict[Text, Any] | None = None is_sparse: bool = False vector_size: int | None = None @@ -678,6 +679,25 @@ def insert(self, df: pd.DataFrame, params: dict = None): logger.warning("No valid content found in any content columns") return + # Check if we should skip existing items (before calculating embeddings) + if params is not None and params.get("kb_skip_existing", False): + logger.debug(f"Checking for existing items to skip before processing {len(df)} items") + db_handler = self.get_vector_db() + + # Get list of IDs from current batch + current_ids = df[TableField.ID.value].dropna().astype(str).tolist() + if current_ids: + # Check which IDs already exist + existing_ids = db_handler.check_existing_ids(self._kb.vector_database_table, current_ids) + if existing_ids: + # Filter out existing items + df = df[~df[TableField.ID.value].astype(str).isin(existing_ids)] + logger.info(f"Skipped {len(existing_ids)} existing items, processing {len(df)} new items") + + if df.empty: + logger.info("All items already exist, nothing to insert") + return + # add embeddings and send to vector db df_emb = self._df_to_embeddings(df) df = pd.concat([df, df_emb], axis=1)