From 1d90e4b65c4e3b6058e20bee48592ef7f1956734 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Sat, 14 Jun 2025 11:36:20 -0500 Subject: [PATCH 01/14] Retriever cache --- syftr/configuration.py | 3 + syftr/flows.py | 72 ++++++++++++--- syftr/optimization.py | 2 +- syftr/retrievers/cached_retriever.py | 126 +++++++++++++++++++++++++++ syftr/retrievers/storage.py | 2 +- syftr/studies.py | 59 ++++++++----- syftr/tuner/qa_tuner.py | 2 + 7 files changed, 227 insertions(+), 39 deletions(-) create mode 100644 syftr/retrievers/cached_retriever.py diff --git a/syftr/configuration.py b/syftr/configuration.py index 3edd177b..e95fe52f 100644 --- a/syftr/configuration.py +++ b/syftr/configuration.py @@ -133,6 +133,9 @@ class Paths(BaseModel): tmp_dir / "huggingface" ) index_cache: Annotated[Path, Field(validate_default=True)] = tmp_dir / "indexcache" + retrieval_cache: Annotated[Path, Field(validate_default=True)] = ( + tmp_dir / "retrieval_cache" + ) onnx_dir: Annotated[Path, Field(validate_default=True)] = tmp_dir / "onnx" sota_dir: Annotated[Path, Field(validate_default=True)] = data_dir / "sota" lock_dir: Annotated[Path, Field(validate_default=True)] = tmp_dir / "syftr-locks" diff --git a/syftr/flows.py b/syftr/flows.py index c0d22bda..39b377d2 100644 --- a/syftr/flows.py +++ b/syftr/flows.py @@ -54,7 +54,12 @@ from syftr.instrumentation.tokens import LLMCallData, TokenTrackingEventHandler from syftr.llm import get_llm_name, get_tokenizer from syftr.logger import logger -from syftr.studies import get_critique_template, get_react_template +from syftr.retrievers.cached_retriever import ( + get_retrieval_cache, + get_retrieval_cache_key, + put_retrieval_cache, +) +from syftr.studies import ParamDict, get_critique_template, get_react_template dispatcher = instrument.get_dispatcher() _event_handler = TokenTrackingEventHandler() @@ -175,6 +180,8 @@ class RetrieverFlow(Flow): hyde_llm: LLM | None = None additional_context_num_nodes: int = 0 name: str = "Retriever Only Flow" + # A unique fingerprint of the retriever, used for caching retrieved chunks + retriever_fingerprint: ParamDict | None = None def __repr__(self): return f"{self.name}: {self.params}" @@ -217,29 +224,40 @@ async def agenerate(self, query: str, *args, **kwargs): @dispatcher.span def retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: + if self.retriever_fingerprint: + return self.cached_retrieve(query) start_time = time.perf_counter() qb = QueryBundle(query) - if isinstance(self.query_engine, TransformQueryEngine): - response = self.query_engine.query(qb) - assert isinstance(response, Response), ( - f"Expected Response, got {type(response)=}" - ) - retrieval_result = response.source_nodes - else: - retrieval_result = self.query_engine.retrieve(qb) + retrieval_result = self.query_engine.retrieve(qb) + duration = time.perf_counter() - start_time + return retrieval_result, duration + + @dispatcher.span + def cached_retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: + assert self.retriever_fingerprint, ( + "Retriever fingerprint must be set to use cache." + ) + start_time = time.perf_counter() + with get_retrieval_cache_key(query, self.retriever_fingerprint) as key: + if (retrieval_result := get_retrieval_cache(key)) is not None: + logger.info(f"Cache hit for question: {query}") + else: + logger.info(f"Cache miss for question: {query}") + qb = QueryBundle(query) + retrieval_result = self.query_engine.retrieve(qb) + put_retrieval_cache(key, retrieval_result) duration = time.perf_counter() - start_time return retrieval_result, duration @dispatcher.span async def aretrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: + if self.retriever_fingerprint: + return await self.cached_aretrieve(query) start_time = time.perf_counter() qb = QueryBundle(query) if isinstance(self.query_engine, TransformQueryEngine): - response = await self.query_engine.aquery(qb) - assert isinstance(response, Response), ( - f"Expected Response, got {type(response)=}" - ) - retrieval_result = response.source_nodes + # TransformQueryEngine does not have aretrieve method + retrieval_result = self.query_engine.retrieve(qb) else: assert hasattr(self.query_engine, "aretrieve"), ( f"{self.query_engine} does not have 'aretrieve' method" @@ -248,6 +266,32 @@ async def aretrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: duration = time.perf_counter() - start_time return retrieval_result, duration + @dispatcher.span + async def cached_aretrieve( + self, query: str + ) -> T.Tuple[T.List[NodeWithScore], float]: + assert self.retriever_fingerprint, ( + "Retriever fingerprint must be set to use cache." + ) + start_time = time.perf_counter() + with get_retrieval_cache_key(query, self.retriever_fingerprint) as key: + if (retrieval_result := get_retrieval_cache(key)) is not None: + logger.info(f"Cache hit for question: {query}") + else: + logger.info(f"Cache miss for question: {query}") + qb = QueryBundle(query) + if isinstance(self.query_engine, TransformQueryEngine): + # TransformQueryEngine does not have aretrieve method + retrieval_result = self.query_engine.retrieve(qb) + else: + assert hasattr(self.query_engine, "aretrieve"), ( + f"{self.query_engine} does not have 'aretrieve' method" + ) + retrieval_result = await self.query_engine.aretrieve(qb) + put_retrieval_cache(key, retrieval_result) + duration = time.perf_counter() - start_time + return retrieval_result, duration + @dataclass(kw_only=True) class RAGFlow(Flow): diff --git a/syftr/optimization.py b/syftr/optimization.py index a42edb3f..614e90b4 100644 --- a/syftr/optimization.py +++ b/syftr/optimization.py @@ -307,7 +307,7 @@ def run(self) -> optuna.Study: block.name, defaults, ) - study_config.search_space.update_defaults(defaults) + study_config.search_space.custom_defaults.update(defaults) future.result() # Wait for seeding to finish diff --git a/syftr/retrievers/cached_retriever.py b/syftr/retrievers/cached_retriever.py new file mode 100644 index 00000000..76401b97 --- /dev/null +++ b/syftr/retrievers/cached_retriever.py @@ -0,0 +1,126 @@ +import hashlib +import io +import json +from contextlib import contextmanager +from typing import Any, Dict, Optional + +import cloudpickle +import diskcache +from lz4.frame import compress, decompress + +from syftr.amazon import get_file_from_s3 +from syftr.configuration import cfg +from syftr.logger import logger +from syftr.studies import ParamDict, StudyConfig +from syftr.utils.locks import distributed_lock + +# Retrieval cache constants and key builder +RETRIEVAL_CACHE_PREFIX = "retrieval_cache" +RETRIEVER_CACHE_VERSION = 1 + + +@contextmanager +def get_retrieval_cache_key(question: str, retriever_params_dict: Dict[str, Any]): + """ + Build a cache key from question text and retriever params. + """ + raw_dict = {**retriever_params_dict, "question": question} + raw = json.dumps(raw_dict, sort_keys=True).encode("utf-8") + cache_key = hashlib.sha1(raw).hexdigest() + host_only = not cfg.storage.s3_cache_enabled + with distributed_lock(cache_key, host_only=host_only): + yield cache_key + + +def get_retriever_fingerprint( + study_config: StudyConfig, params: ParamDict +) -> Dict[str, Any]: + param_names = [ + "hyde_enabled", + "additional_context_enabled", + "rag_method", + "rag_query_decomposition_enabled", + "rag_top_k", + "rag_embedding_model", + "rag_query_decomposition_llm_name", + "rag_query_decomposition_num_queries", + "rag_fusion_mode", + "splitter_method", + "splitter_chunk_exp", + "splitter_chunk_overlap_frac", + "hyde_llm_name", + "additional_context_num_nodes", + ] + retriever_params = { + key: value for key, value in params.items() if key in param_names + } + dataset_param_names = ["xname", "partition_map", "subset", "grounding_data_path"] + dataset_params = { + key: value + for key, value in study_config.dataset.model_dump().items() + if key in dataset_param_names + } + key_params = { + "retriever": retriever_params, + "dataset": dataset_params, + "cache_version": RETRIEVER_CACHE_VERSION, + } + return key_params + + +@contextmanager +def local_retrieval_cache(): + with diskcache.Cache( + cfg.paths.retrieval_cache, + size_limit=cfg.storage.local_cache_max_size_gb * 1024**3, + ) as cache: + yield cache + + +def put_retrieval_cache(cache_key: str, obj: Any, local_only: bool = False): + """ + Mirror to both diskcache & S3 under “retrieval_cache/{cache_key}.pkl”. + """ + serialized = compress(cloudpickle.dumps(obj)) + # Local diskcache + with local_retrieval_cache() as cache: + logger.info(f"Storing object to {cache.directory} under key {cache_key}") + cache.set(cache_key, serialized) + + # S3 mirror + if not local_only and cfg.storage.s3_cache_enabled: + s3_key = f"{RETRIEVAL_CACHE_PREFIX}/{cache_key}.pkl" + try: + import boto3 + from boto3.s3.transfer import TransferConfig + except ImportError: + logger.info("Skipping S3 cache - install boto3 to cache objects in S3") + return + s3 = boto3.client("s3") + config = TransferConfig(multipart_threshold=5 * 1024**3) + fileobj = io.BytesIO(serialized) + logger.info(f"Storing object to S3: {s3_key}") + s3.upload_fileobj(fileobj, cfg.storage.cache_bucket, s3_key, Config=config) + logger.info("Done storing object to S3") + + +def get_retrieval_cache(cache_key: str) -> Optional[Any]: + """ + First check diskcache, then fall back to S3. Returns the retrieved list + (NodeWithScore) or None if missing. + """ + s3_key = f"{RETRIEVAL_CACHE_PREFIX}/{cache_key}.pkl" + # Try local + with local_retrieval_cache() as cache: + if (data := cache.get(cache_key)) is not None: + logger.info(f"Loading cached object from {cache.directory}") + return cloudpickle.loads(decompress(data)) + # Try S3 + if cfg.storage.s3_cache_enabled: + if (data := get_file_from_s3(s3_key)) is not None: + logger.info(f"Loading cached object from S3: {s3_key}") + obj = cloudpickle.loads(decompress(data)) + # populate local cache + put_retrieval_cache(cache_key, obj, local_only=True) + return obj + return None diff --git a/syftr/retrievers/storage.py b/syftr/retrievers/storage.py index a4cbffde..afe89c00 100644 --- a/syftr/retrievers/storage.py +++ b/syftr/retrievers/storage.py @@ -85,7 +85,7 @@ def put_cache(cache_key, index, local_only: bool = False) -> None: with local_cache() as cache: logger.info(f"Storing index to {cache.directory}") - cache.add(cache_key, serialized_obj) + cache.set(cache_key, serialized_obj) logger.info(f"Done storing index to {cache.directory}") if not local_only and cfg.storage.s3_cache_enabled: diff --git a/syftr/studies.py b/syftr/studies.py index 95595678..b683f5c5 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -945,9 +945,12 @@ class SearchSpace(BaseModel): default_factory=LATSRagAgent, description="Configuration for the LATS RAG agent.", ) - _custom_defaults: ParamDict = {} + custom_defaults: ParamDict = Field( + default_factory=dict, + description="Override default parameters for the search space.", + ) - def _defaults(self) -> ParamDict: + def defaults(self) -> ParamDict: return { "rag_mode": self.rag_modes[0], "template_name": self.template_names[0], @@ -964,15 +967,6 @@ def _defaults(self) -> ParamDict: **self.lats_rag_agent.defaults(), } - def update_defaults(self, defaults: ParamDict) -> None: - self._custom_defaults.update(defaults) - - def defaults(self) -> ParamDict: - return { - **self._defaults(), - **self._custom_defaults, - } - def param_names( self, params: T.Dict[str, T.Any] | T.List[str] | None = None ) -> T.List[str]: @@ -1022,10 +1016,9 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic for param in parameters: assert param in PARAMETERS, f"Invalid parameter: {param}" - params: ParamDict = { - "few_shot_enabled": False, - } + params: ParamDict = {} defaults = self.defaults() + defaults.update(self.custom_defaults) if "rag_mode" in parameters: params["rag_mode"] = trial.suggest_categorical("rag_mode", self.rag_modes) @@ -1046,7 +1039,7 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic else: params["response_synthesizer_llm"] = defaults["response_synthesizer_llm"] - # No-RAG general parameters + # RAG general parameters if params["rag_mode"] != "no_rag": if "rag_retriever" in parameters: params.update(**self.rag_retriever.sample(trial)) @@ -1065,7 +1058,9 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic if params["reranker_enabled"]: params.update(**self.reranker.sample(trial)) else: - params["reranker_enabled"] = False + params["reranker_enabled"] = defaults["reranker_enabled"] + if params["reranker_enabled"]: + params.update(**self.reranker.defaults()) if "hyde" in parameters: params["hyde_enabled"] = trial.suggest_categorical( @@ -1074,7 +1069,9 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic if params["hyde_enabled"]: params.update(**self.hyde.sample(trial)) else: - params["hyde_enabled"] = False + params["hyde_enabled"] = defaults["hyde_enabled"] + if params["hyde_enabled"]: + params.update(**self.hyde.defaults()) if "additional_context" in parameters: params["additional_context_enabled"] = trial.suggest_categorical( @@ -1083,7 +1080,11 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic if params["additional_context_enabled"]: params.update(**self.additional_context.sample(trial)) else: - params["additional_context_enabled"] = False + params["additional_context_enabled"] = defaults[ + "additional_context_enabled" + ] + if params["additional_context_enabled"]: + params.update(**self.additional_context.defaults()) if params["rag_mode"] == "react_rag_agent": if "react_rag_agent" in parameters: @@ -1105,13 +1106,21 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic params.update(**self.lats_rag_agent.sample(trial)) else: params.update(**self.lats_rag_agent.defaults()) - params.update(**self.lats_rag_agent.sample(trial)) - if few_shot_enabled := trial.suggest_categorical( - "few_shot_enabled", self.few_shot_enabled - ): + if "few_shot_retriever" in parameters: + few_shot_enabled = trial.suggest_categorical( + "few_shot_enabled", self.few_shot_enabled + ) params["few_shot_enabled"] = few_shot_enabled - params.update(**self.few_shot_retriever.sample(trial)) + if few_shot_enabled: + params.update(**self.few_shot_retriever.sample(trial)) + else: + params["few_shot_enabled"] = defaults["few_shot_enabled"] + if params["few_shot_enabled"]: + params.update(**self.few_shot_retriever.defaults()) + + # Use custom defaults to override any defaults + params.update(self.custom_defaults) return params @@ -1291,6 +1300,10 @@ class Block(BaseModel): components: T.List[str] = Field( default_factory=lambda: PARAMETERS, description="Block components" ) + custom_defaults: ParamDict = Field( + default_factory=dict, + description="Custom default parameters used for this block.", + ) class OptimizationConfig(BaseModel): diff --git a/syftr/tuner/qa_tuner.py b/syftr/tuner/qa_tuner.py index ee99b249..30ae22f5 100644 --- a/syftr/tuner/qa_tuner.py +++ b/syftr/tuner/qa_tuner.py @@ -41,6 +41,7 @@ ) from syftr.ray.utils import ray_init from syftr.retrievers.build import build_rag_retriever +from syftr.retrievers.cached_retriever import get_retriever_fingerprint from syftr.startup import prepare_worker from syftr.studies import ( RetrieverStudyConfig, @@ -139,6 +140,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: additional_context_num_nodes=params.get("additional_context_num_nodes", 0), params=params, enforce_full_evaluation=enforce_full_evaluation, + retriever_fingerprint=get_retriever_fingerprint(study_config, params), ) get_qa_examples = None From 931f970fd1e1c0a7c874c3ce423fc17ab6af0eca Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Sat, 14 Jun 2025 12:43:56 -0500 Subject: [PATCH 02/14] Updates --- syftr/studies.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/syftr/studies.py b/syftr/studies.py index b683f5c5..c5e7c480 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -1018,7 +1018,6 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic params: ParamDict = {} defaults = self.defaults() - defaults.update(self.custom_defaults) if "rag_mode" in parameters: params["rag_mode"] = trial.suggest_categorical("rag_mode", self.rag_modes) @@ -1119,7 +1118,7 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic if params["few_shot_enabled"]: params.update(**self.few_shot_retriever.defaults()) - # Use custom defaults to override any defaults + # Use custom defaults to override parameters params.update(self.custom_defaults) return params From 0e9c15484cbac8a88312c9b0070ffc9e586a30b3 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Sun, 15 Jun 2025 12:17:41 -0500 Subject: [PATCH 03/14] Cache updates --- syftr/flows.py | 109 +++++++++++---------------- syftr/logger.py | 3 + syftr/ray/utils.py | 61 ++++++++++++++- syftr/retrievers/cached_retriever.py | 18 +++++ syftr/tuner/qa_tuner.py | 8 +- 5 files changed, 132 insertions(+), 67 deletions(-) diff --git a/syftr/flows.py b/syftr/flows.py index 39b377d2..1880ebe8 100644 --- a/syftr/flows.py +++ b/syftr/flows.py @@ -48,6 +48,7 @@ from llama_index.core.storage.docstore.types import BaseDocumentStore from llama_index.core.tools import BaseTool, QueryEngineTool, ToolMetadata from numpy import ceil +from overrides import overrides from syftr.configuration import cfg from syftr.instrumentation.arize import instrument_arize @@ -181,26 +182,18 @@ class RetrieverFlow(Flow): additional_context_num_nodes: int = 0 name: str = "Retriever Only Flow" # A unique fingerprint of the retriever, used for caching retrieved chunks - retriever_fingerprint: ParamDict | None = None + retriever_cache_fingerprint: ParamDict | None = None def __repr__(self): return f"{self.name}: {self.params}" @cached_property def query_engine(self) -> BaseQueryEngine: - node_postprocessors: T.List[BaseNodePostprocessor] = [] - if self.additional_context_num_nodes > 0: - assert self.docstore is not None - node_postprocessors.append( - PrevNextNodePostprocessor( - docstore=self.docstore, - num_nodes=int(ceil(self.additional_context_num_nodes / 2)), - mode="both", - ) - ) + node_postprocessors = self._build_node_postprocessors() response_synthesizer = get_response_synthesizer( llm=self.response_synthesizer_llm, response_mode=ResponseMode.COMPACT, + text_qa_template=self.prompt_template, ) base_engine = RetrieverQueryEngine( retriever=self.retriever, @@ -212,19 +205,32 @@ def query_engine(self) -> BaseQueryEngine: return TransformQueryEngine(base_engine, query_transform=hyde) return base_engine + def _build_node_postprocessors(self) -> T.List[BaseNodePostprocessor]: + node_postprocessors: T.List[BaseNodePostprocessor] = [] + if self.additional_context_num_nodes > 0: + assert self.docstore is not None + node_postprocessors.append( + PrevNextNodePostprocessor( + docstore=self.docstore, + num_nodes=int(ceil(self.additional_context_num_nodes / 2)), + mode="both", + ) + ) + return node_postprocessors + @cached_property def tokenizer(self) -> T.Callable: return get_tokenizer(get_llm_name(self.response_synthesizer_llm)) - def generate(self, query: str, *args, **kwargs): + def _generate(self, query: str, *args, **kwargs): raise NotImplementedError("RetrieverFlow does not support generation.") - async def agenerate(self, query: str, *args, **kwargs): + async def _agenerate(self, query: str, *args, **kwargs): raise NotImplementedError("RetrieverFlow does not support generation.") @dispatcher.span def retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: - if self.retriever_fingerprint: + if self.retriever_cache_fingerprint: return self.cached_retrieve(query) start_time = time.perf_counter() qb = QueryBundle(query) @@ -234,15 +240,15 @@ def retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: @dispatcher.span def cached_retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: - assert self.retriever_fingerprint, ( + assert self.retriever_cache_fingerprint, ( "Retriever fingerprint must be set to use cache." ) start_time = time.perf_counter() - with get_retrieval_cache_key(query, self.retriever_fingerprint) as key: + with get_retrieval_cache_key(query, self.retriever_cache_fingerprint) as key: if (retrieval_result := get_retrieval_cache(key)) is not None: - logger.info(f"Cache hit for question: {query}") + logger.info(f"Retriever cache hit: {query}") else: - logger.info(f"Cache miss for question: {query}") + logger.info(f"Retriever cache miss: {query}") qb = QueryBundle(query) retrieval_result = self.query_engine.retrieve(qb) put_retrieval_cache(key, retrieval_result) @@ -251,7 +257,7 @@ def cached_retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: @dispatcher.span async def aretrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: - if self.retriever_fingerprint: + if self.retriever_cache_fingerprint: return await self.cached_aretrieve(query) start_time = time.perf_counter() qb = QueryBundle(query) @@ -270,15 +276,15 @@ async def aretrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: async def cached_aretrieve( self, query: str ) -> T.Tuple[T.List[NodeWithScore], float]: - assert self.retriever_fingerprint, ( + assert self.retriever_cache_fingerprint, ( "Retriever fingerprint must be set to use cache." ) start_time = time.perf_counter() - with get_retrieval_cache_key(query, self.retriever_fingerprint) as key: + with get_retrieval_cache_key(query, self.retriever_cache_fingerprint) as key: if (retrieval_result := get_retrieval_cache(key)) is not None: - logger.info(f"Cache hit for question: {query}") + logger.info(f"Retriever cache hit: {query}") else: - logger.info(f"Cache miss for question: {query}") + logger.info(f"Retriever cache miss: {query}") qb = QueryBundle(query) if isinstance(self.query_engine, TransformQueryEngine): # TransformQueryEngine does not have aretrieve method @@ -294,21 +300,17 @@ async def cached_aretrieve( @dataclass(kw_only=True) -class RAGFlow(Flow): - retriever: BaseRetriever - docstore: BaseDocumentStore | None = None - hyde_llm: LLM | None = None +class RAGFlow(RetrieverFlow): reranker_llm: LLM | None = None reranker_top_k: int | None = None name: str = "RAG Flow" - additional_context_num_nodes: int = 0 def __repr__(self): return f"{self.name}: {self.params}" - @cached_property - def query_engine(self) -> BaseQueryEngine: - node_postprocessors: T.List[BaseNodePostprocessor] = [] + @overrides + def _build_node_postprocessors(self) -> T.List[BaseNodePostprocessor]: + node_postprocessors = super()._build_node_postprocessors() if self.reranker_llm is not None: assert self.reranker_top_k, ( "Reranker enabled, need reranker_top_k param set" @@ -316,29 +318,7 @@ def query_engine(self) -> BaseQueryEngine: node_postprocessors.append( LLMRerank(top_n=self.reranker_top_k, llm=self.reranker_llm) ) - if self.additional_context_num_nodes > 0: - assert self.docstore is not None - node_postprocessors.append( - PrevNextNodePostprocessor( - docstore=self.docstore, - num_nodes=int(ceil(self.additional_context_num_nodes / 2)), - mode="both", - ) - ) - response_synthesizer = get_response_synthesizer( - llm=self.response_synthesizer_llm, - response_mode=ResponseMode.COMPACT, - text_qa_template=self.prompt_template, - ) - retriever = RetrieverQueryEngine( - retriever=self.retriever, - response_synthesizer=response_synthesizer, - node_postprocessors=node_postprocessors, - ) - if self.hyde_llm is not None: - hyde = HyDEQueryTransform(llm=self.hyde_llm, include_original=True) - retriever = TransformQueryEngine(retriever, query_transform=hyde) # type: ignore - return retriever + return node_postprocessors def get_prompt(self, query) -> str: if self.template is None: @@ -355,21 +335,16 @@ def get_prompt(self, query) -> str: few_shot_examples=examples, ) - def retrieve(self, query: str) -> T.List[NodeWithScore]: - return self.query_engine.retrieve(QueryBundle(query)) - - async def aretrieve(self, query: str) -> T.List[NodeWithScore]: - assert hasattr(self.query_engine, "aretrieve"), ( - f"{self.query_engine} does not have 'aretrieve' method" - ) - return await self.query_engine.aretrieve(QueryBundle(query)) - @dispatcher.span def _generate( self, query: str, invocation_id: str ) -> T.Tuple[CompletionResponse, float]: start_time = time.perf_counter() - response = self.query_engine.query(query) + nodes, _ = self.retrieve(query) + response = self.query_engine.synthesize( + query_bundle=QueryBundle(query), + nodes=nodes, + ) assert isinstance(response, Response), ( f"Expected Response, got {type(response)=}" ) @@ -388,7 +363,11 @@ async def _agenerate( self, query: str, invocation_id: str ) -> T.Tuple[CompletionResponse, float]: start_time = time.perf_counter() - response = await self.query_engine.aquery(query) + nodes, _ = await self.aretrieve(query) + response = await self.query_engine.asynthesize( + query_bundle=QueryBundle(query), + nodes=nodes, + ) assert isinstance(response, Response), ( f"Expected Response, got {type(response)=}" ) diff --git a/syftr/logger.py b/syftr/logger.py index df690922..fe632b26 100644 --- a/syftr/logger.py +++ b/syftr/logger.py @@ -10,6 +10,9 @@ else: logger = logging.getLogger(cfg.logging.name) +# Disable propagation so logs don't get double printed by Ray +logger.propagate = False + def _create_default_handler(handler_name: str) -> logging.Handler: """Gives the project default logging handler""" diff --git a/syftr/ray/utils.py b/syftr/ray/utils.py index aa3e069b..bbe1b14d 100644 --- a/syftr/ray/utils.py +++ b/syftr/ray/utils.py @@ -1,11 +1,15 @@ import getpass import os - +import functools +from typing import Any, Final import ray from syftr.configuration import cfg from syftr.logger import logger +NAMESPACE: Final = "syftr" +RAY_CACHE_ACTOR_NAME: Final = "ray_cache" + def ray_init(force_remote: bool = False): if ray.is_initialized(): @@ -26,4 +30,59 @@ def ray_init(force_remote: bool = False): ray.init( address=address, logging_level=cfg.logging.level, + namespace=NAMESPACE, ) + + +@ray.remote +class RayCacheActor: + def __init__(self): + self.cache = {} + + def get(self, key): + return self.cache.get(key, None) + + def set(self, key, value): + self.cache[key] = value + + def contains(self, key): + return key in self.cache + + def clear(self): + self.cache.clear() + + +@functools.lru_cache(maxsize=1) +def get_ray_cache(): + """Returns the ray cache actor.""" + if not ray.is_initialized(): + raise RuntimeError("Ray is not initialized. Cannot use Ray cache.") + try: + return ray.get_actor(RAY_CACHE_ACTOR_NAME) + except ValueError: + return RayCacheActor.options( + name=RAY_CACHE_ACTOR_NAME, lifetime="detached", namespace=NAMESPACE + ).remote() + + +def ray_cache_get(key: str): + """Get a value from the Ray cache. Raises a RuntimeError if Ray not initialized.""" + ray_cache = get_ray_cache() + return ray.get(ray_cache.get.remote(key)) + + +def ray_cache_put(key: str, obj: Any) -> None: + """Put a value in the ray cache. Raises a RuntimeError if Ray not initialized.""" + ray_cache = get_ray_cache() + ray_cache.set.remote(key, obj) + + +def ray_cache_restart(): + """Kills and restarts the otherwise persistent Ray cache actor.""" + try: + old = ray.get_actor(RAY_CACHE_ACTOR_NAME, namespace=NAMESPACE) + ray.kill(old) + except ValueError: + pass + get_ray_cache.cache_clear() + return get_ray_cache() diff --git a/syftr/retrievers/cached_retriever.py b/syftr/retrievers/cached_retriever.py index 76401b97..01c19f67 100644 --- a/syftr/retrievers/cached_retriever.py +++ b/syftr/retrievers/cached_retriever.py @@ -11,6 +11,7 @@ from syftr.amazon import get_file_from_s3 from syftr.configuration import cfg from syftr.logger import logger +from syftr.ray.utils import ray_cache_get, ray_cache_put from syftr.studies import ParamDict, StudyConfig from syftr.utils.locks import distributed_lock @@ -87,6 +88,13 @@ def put_retrieval_cache(cache_key: str, obj: Any, local_only: bool = False): logger.info(f"Storing object to {cache.directory} under key {cache_key}") cache.set(cache_key, serialized) + # Try ray cache + try: + logger.info(f"Storing {cache_key} to Ray cache") + ray_cache_put(cache_key, serialized) + except Exception as e: + logger.warning(f"Skipping Ray cache put due to error: {e}") + # S3 mirror if not local_only and cfg.storage.s3_cache_enabled: s3_key = f"{RETRIEVAL_CACHE_PREFIX}/{cache_key}.pkl" @@ -115,6 +123,16 @@ def get_retrieval_cache(cache_key: str) -> Optional[Any]: if (data := cache.get(cache_key)) is not None: logger.info(f"Loading cached object from {cache.directory}") return cloudpickle.loads(decompress(data)) + + # Try Ray cache + try: + data = ray_cache_get(cache_key) + if data is not None: + logger.info(f"Loading {cache_key} from Ray cache") + return cloudpickle.loads(decompress(data)) + except Exception as e: + logger.warning(f"Skipping Ray cache get due to error: {e}") + # Try S3 if cfg.storage.s3_cache_enabled: if (data := get_file_from_s3(s3_key)) is not None: diff --git a/syftr/tuner/qa_tuner.py b/syftr/tuner/qa_tuner.py index 30ae22f5..ff7812f5 100644 --- a/syftr/tuner/qa_tuner.py +++ b/syftr/tuner/qa_tuner.py @@ -140,7 +140,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: additional_context_num_nodes=params.get("additional_context_num_nodes", 0), params=params, enforce_full_evaluation=enforce_full_evaluation, - retriever_fingerprint=get_retriever_fingerprint(study_config, params), + retriever_cache_fingerprint=get_retriever_fingerprint(study_config, params), ) get_qa_examples = None @@ -174,6 +174,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: enforce_full_evaluation=enforce_full_evaluation, ) else: + retriever_cache_fingerprint = get_retriever_fingerprint(study_config, params) hyde_llm = reranker_llm = reranker_top_k = None if params.get("hyde_enabled"): hyde_llm = get_llm(params["hyde_llm_name"]) @@ -201,6 +202,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: additional_context_num_nodes=additional_context_num_nodes, enforce_full_evaluation=enforce_full_evaluation, params=params, + retriever_cache_fingerprint=retriever_cache_fingerprint, ) case "react_rag_agent": subquestion_engine_llm = get_llm(params["subquestion_engine_llm"]) @@ -225,6 +227,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: dataset_description=study_config.dataset.description, enforce_full_evaluation=enforce_full_evaluation, params=params, + retriever_cache_fingerprint=retriever_cache_fingerprint, ) case "critique_rag_agent": subquestion_engine_llm = get_llm(params["subquestion_engine_llm"]) @@ -253,6 +256,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: dataset_description=study_config.dataset.description, enforce_full_evaluation=enforce_full_evaluation, params=params, + retriever_cache_fingerprint=retriever_cache_fingerprint, ) case "sub_question_rag": subquestion_engine_llm = get_llm(params["subquestion_engine_llm"]) @@ -275,6 +279,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: dataset_description=study_config.dataset.description, enforce_full_evaluation=enforce_full_evaluation, params=params, + retriever_cache_fingerprint=retriever_cache_fingerprint, ) case "lats_rag_agent": flow = LATSAgentFlow( @@ -293,6 +298,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: max_rollouts=params["lats_max_rollouts"], enforce_full_evaluation=enforce_full_evaluation, params=params, + retriever_cache_fingerprint=retriever_cache_fingerprint, ) case _: raise ValueError(f"Invalid rag_mode: {params['rag_mode']}") From 549016bb814340aa71e04f2c8f442f8c87dd3b83 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 16 Jun 2025 10:17:23 -0500 Subject: [PATCH 04/14] Add no-retriever studies --- studies/no-retriever-financebench.yaml | 83 ++++++++++++++++++++++++++ studies/no-retriever-hotpotqa.yaml | 80 +++++++++++++++++++++++++ studies/no-retriever-multihoprag.yaml | 73 ++++++++++++++++++++++ 3 files changed, 236 insertions(+) create mode 100644 studies/no-retriever-financebench.yaml create mode 100644 studies/no-retriever-hotpotqa.yaml create mode 100644 studies/no-retriever-multihoprag.yaml diff --git a/studies/no-retriever-financebench.yaml b/studies/no-retriever-financebench.yaml new file mode 100644 index 00000000..8d4120fb --- /dev/null +++ b/studies/no-retriever-financebench.yaml @@ -0,0 +1,83 @@ +name: "no-retriever-financebench" +dataset: + dataset_dir: partitioned + description: Financial dataset that contains everything about finance, including + real-world financial documents, SEC filings, earning reports, call transcripts, + and much more. It has all the financial live data, historical data, just about + everything about finance, for instance, definitions and explanations of financial + term, insights on company revenues, mergers, founders, or stock performance, details + on financial laws, compliance, or government policies, information required to + evaluate finance risk, and information about banking operations, credit systems, + or loan structures. + examples_data_path: examples + grounding_data_path: aryn_html + load_examples_timeout_s: 3600 + load_grounding_data_timeout_s: 3600 + partition_map: + holdout: holdout + sample: pepsi + test: test + train: train + path_root: octo-syftr-benchmarking-data + storage_options: + cache_check: 60 + cache_storage: benchmarking/data/cache + check_files: true + expiry_time: 2592000 + protocol: filecache + same_names: true + target_protocol: s3 + storage_partitions: + - sample + - train + - test + - holdout + xname: financebench_hf +reuse_study: false +search_space: + custom_defaults: + rag_mode: rag + hyde_enabled: true + additional_context_enabled: true + rag_method: dense + rag_query_decomposition_enabled: true + rag_top_k: 19 + rag_embedding_model: BAAI/bge-small-en-v1.5 + rag_query_decomposition_llm_name: gpt-4o-mini + rag_query_decomposition_num_queries: 6 + rag_fusion_mode: simple + splitter_method: sentence + splitter_chunk_exp: 12 + splitter_chunk_overlap_frac: 0 + hyde_llm_name: gpt-4o-mini + additional_context_num_nodes: 16 +optimization: + baselines: [] + use_individual_baselines: false + use_agent_baselines: false + use_variations_of_baselines: false + shuffle_baselines: false + blocks: + - components: + - few_shot_retriever + - reranker + - sub_question_rag + - critique_rag_agent + - lats_rag_agent + - react_rag_agent + - response_synthesizer_llm + - template_name + name: no-retriever + num_trials: 500 + cpus_per_trial: 5 + embedding_device: cuda + # use_hf_embedding_models: true + gpus_per_trial: 0.2 + max_concurrent_trials: 10 + num_eval_samples: 100 + num_eval_batch: 100 + num_retries_unique_params: 10 + num_trials: 100 + use_cost_pruner: false + use_pareto_pruner: false + use_runtime_pruner: false diff --git a/studies/no-retriever-hotpotqa.yaml b/studies/no-retriever-hotpotqa.yaml new file mode 100644 index 00000000..b77dadc5 --- /dev/null +++ b/studies/no-retriever-hotpotqa.yaml @@ -0,0 +1,80 @@ +name: "no-retriever-hotpotqa" +dataset: + dataset_dir: partitioned + description: This dataset is a vast collection of all kind of information that you + can find on Wikipedia. It can be used, for instance, to retrieve straightforward + facts from one or more documents, compare two entities based on shared attributes, + identify relationships, roles, or attributes of entities, reason about dates, + timelines, or chronological order, determine geographical relationships or locations, + explain causes or sequences of events or processes, synthesize facts from multiple + documents to infer answers, and validate or refute premises in the context of + the question. + examples_data_path: examples + grounding_data_path: grounding_data + load_examples_timeout_s: 3600 + load_grounding_data_timeout_s: 3600 + partition_map: + holdout: holdout + sample: sample + test: test + train: train + path_root: octo-syftr-benchmarking-data + storage_options: + cache_check: 60 + cache_storage: benchmarking/data/cache + check_files: true + expiry_time: 2592000 + protocol: filecache + same_names: true + target_protocol: s3 + storage_partitions: + - sample + - train + - test + - holdout + subset: train_hard + xname: hotpotqa_hf +reuse_study: false +optimization: + baselines: [] + use_individual_baselines: false + use_agent_baselines: false + use_variations_of_baselines: false + shuffle_baselines: false + blocks: + - components: + - few_shot_retriever + - reranker + - rag_mode + - sub_question_rag + - critique_rag_agent + - lats_rag_agent + - react_rag_agent + - response_synthesizer_llm + - template_name + name: no-retriever + num_trials: 500 + custom_defaults: + hyde_enabled: false + additional_context_enabled: true + rag_method: hybrid + rag_query_decomposition_enabled: false + rag_top_k: 19 + rag_embedding_model: avsolatorio/GIST-large-Embedding-v0 + rag_hybrid_bm25_weight: 0.6 + rag_query_decomposition_num_queries: 6 + rag_fusion_mode: dist_based_score + splitter_method: html + splitter_chunk_exp: 10 + splitter_chunk_overlap_frac: 0.25 + hyde_llm_name: llama-33-70B + additional_context_num_nodes: 7 + cpus_per_trial: 5 + # embedding_device: cuda + # use_hf_embedding_models: true + gpus_per_trial: 0.2 + max_concurrent_trials: 50 + num_eval_samples: 100 + num_eval_batch: 10 + num_retries_unique_params: 10 + num_trials: 500 diff --git a/studies/no-retriever-multihoprag.yaml b/studies/no-retriever-multihoprag.yaml new file mode 100644 index 00000000..f45b092b --- /dev/null +++ b/studies/no-retriever-multihoprag.yaml @@ -0,0 +1,73 @@ +name: "no-retriever-multihoprag" +dataset: + dataset_dir: partitioned + description: multihoprag dataset + examples_data_path: examples + grounding_data_path: grounding_data + load_examples_timeout_s: 3600 + load_grounding_data_timeout_s: 3600 + partition_map: + holdout: holdout + sample: sample + test: test + train: train + path_root: octo-syftr-benchmarking-data + storage_options: + cache_check: 60 + cache_storage: benchmarking/data/cache + check_files: true + expiry_time: 2592000 + protocol: filecache + same_names: true + target_protocol: s3 + storage_partitions: + - sample + - train + - test + - holdout + xname: multihoprag_hf +reuse_study: true +toy_mode: false +search_space: + custom_defaults: + rag_mode: rag + hyde_enabled: false + additional_context_enabled: false + rag_method: hybrid + rag_query_decomposition_enabled: true + rag_top_k: 19 + rag_embedding_model: Labib11/MUG-B-1.6 + rag_hybrid_bm25_weight: 0.1s + rag_query_decomposition_llm_name: gpt-4o-mini + rag_query_decomposition_num_queries: 6 + rag_fusion_mode: reciprocal_rerank + splitter_method: sentence + splitter_chunk_exp: 12 + splitter_chunk_overlap_frac: 0 +optimization: + baselines: [] + use_individual_baselines: false + use_agent_baselines: false + use_variations_of_baselines: false + shuffle_baselines: false + blocks: + - components: + - few_shot_retriever + - reranker + - sub_question_rag + - critique_rag_agent + - lats_rag_agent + - react_rag_agent + - response_synthesizer_llm + - template_name + name: no-retriever + num_trials: 500 + cpus_per_trial: 5 + embedding_device: cuda + # use_hf_embedding_models: true + gpus_per_trial: 0.2 + max_concurrent_trials: 50 + num_eval_samples: 200 + num_eval_batch: 20 + num_retries_unique_params: 10 + num_trials: 500 From ec14d936589417d50b5d406fdc741ca32edcff8e Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 16 Jun 2025 10:23:18 -0500 Subject: [PATCH 05/14] typo --- studies/no-retriever-multihoprag.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/studies/no-retriever-multihoprag.yaml b/studies/no-retriever-multihoprag.yaml index f45b092b..1239f887 100644 --- a/studies/no-retriever-multihoprag.yaml +++ b/studies/no-retriever-multihoprag.yaml @@ -37,7 +37,7 @@ search_space: rag_query_decomposition_enabled: true rag_top_k: 19 rag_embedding_model: Labib11/MUG-B-1.6 - rag_hybrid_bm25_weight: 0.1s + rag_hybrid_bm25_weight: 0.1 rag_query_decomposition_llm_name: gpt-4o-mini rag_query_decomposition_num_queries: 6 rag_fusion_mode: reciprocal_rerank From efd687393a4e89b09d10cff73c91a46149b923d3 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 16 Jun 2025 10:43:04 -0500 Subject: [PATCH 06/14] Add ray cache to indexes --- syftr/flows.py | 4 ++-- syftr/retrievers/storage.py | 12 ++++++++++++ 2 files changed, 14 insertions(+), 2 deletions(-) diff --git a/syftr/flows.py b/syftr/flows.py index 1880ebe8..24b1417e 100644 --- a/syftr/flows.py +++ b/syftr/flows.py @@ -246,9 +246,9 @@ def cached_retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: start_time = time.perf_counter() with get_retrieval_cache_key(query, self.retriever_cache_fingerprint) as key: if (retrieval_result := get_retrieval_cache(key)) is not None: - logger.info(f"Retriever cache hit: {query}") + logger.debug(f"Retriever cache hit: {query}") else: - logger.info(f"Retriever cache miss: {query}") + logger.debug(f"Retriever cache miss: {query}") qb = QueryBundle(query) retrieval_result = self.query_engine.retrieve(qb) put_retrieval_cache(key, retrieval_result) diff --git a/syftr/retrievers/storage.py b/syftr/retrievers/storage.py index afe89c00..f7da842e 100644 --- a/syftr/retrievers/storage.py +++ b/syftr/retrievers/storage.py @@ -14,6 +14,7 @@ from syftr.amazon import delete_file_from_s3, get_file_from_s3 from syftr.configuration import cfg from syftr.logger import logger +from syftr.ray.utils import ray_cache_get, ray_cache_put from syftr.studies import StudyConfig from syftr.utils.locks import distributed_lock @@ -88,6 +89,12 @@ def put_cache(cache_key, index, local_only: bool = False) -> None: cache.set(cache_key, serialized_obj) logger.info(f"Done storing index to {cache.directory}") + try: + logger.info(f"Storing {cache_key} to Ray cache") + ray_cache_put(cache_key, serialized_obj) + except Exception as e: + logger.warning(f"Skipping Ray cache put due to error: {e}") + if not local_only and cfg.storage.s3_cache_enabled: try: import boto3 @@ -117,6 +124,11 @@ def get_cached(cache_key: str) -> Optional[Any]: logger.info(f"Loaded pre-built index from {cache.directory}") return index + data = ray_cache_get(cache_key) + if data is not None: + logger.info(f"Loading {cache_key} from Ray cache") + return cloudpickle.loads(decompress(data)) + if cfg.storage.s3_cache_enabled: if (data := get_file_from_s3(s3_cache_key)) is not None: logger.info(f"Loading pre-built index from S3: {s3_cache_key}") From b324524abffac254aacd723b190dc7abfa7755e2 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 16 Jun 2025 11:20:10 -0500 Subject: [PATCH 07/14] Add option to enable/disable retriever caching --- syftr/studies.py | 3 +++ syftr/tuner/qa_tuner.py | 13 +++++++++++-- 2 files changed, 14 insertions(+), 2 deletions(-) diff --git a/syftr/studies.py b/syftr/studies.py index c5e7c480..88a3ac5e 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -1622,6 +1622,9 @@ class StudyConfig(BaseSettings): toy_mode: bool = Field( default=False, description="Whether to run in toy mode (with smaller dataset)." ) + cache_query_responses: bool = Field( + default=True, description="Whether to cache queries and retrieval responses." + ) model_config = SettingsConfigDict( extra="forbid", # Forbids unknown fields diff --git a/syftr/tuner/qa_tuner.py b/syftr/tuner/qa_tuner.py index ff7812f5..76cebe08 100644 --- a/syftr/tuner/qa_tuner.py +++ b/syftr/tuner/qa_tuner.py @@ -128,6 +128,11 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: enforce_full_evaluation = params.get("enforce_full_evaluation", False) if study_config.is_retriever_study: + retriever_cache_fingerprint = ( + get_retriever_fingerprint(study_config, params) + if study_config.cache_query_responses + else None + ) hyde_llm = ( get_llm(params["hyde_llm_name"]) if params.get("hyde_enabled") else None ) @@ -140,7 +145,7 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: additional_context_num_nodes=params.get("additional_context_num_nodes", 0), params=params, enforce_full_evaluation=enforce_full_evaluation, - retriever_cache_fingerprint=get_retriever_fingerprint(study_config, params), + retriever_cache_fingerprint=retriever_cache_fingerprint, ) get_qa_examples = None @@ -174,7 +179,11 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: enforce_full_evaluation=enforce_full_evaluation, ) else: - retriever_cache_fingerprint = get_retriever_fingerprint(study_config, params) + retriever_cache_fingerprint = ( + get_retriever_fingerprint(study_config, params) + if study_config.cache_query_responses + else None + ) hyde_llm = reranker_llm = reranker_top_k = None if params.get("hyde_enabled"): hyde_llm = get_llm(params["hyde_llm_name"]) From e2b174c1d4c8d4cb64f769f2bc52f91775a5cfa6 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Fri, 20 Jun 2025 16:33:26 -0300 Subject: [PATCH 08/14] Major refactor --- syftr/evaluation.py | 14 ++- syftr/flows.py | 161 ++++++++++++++------------ syftr/retrievers/cached_retriever.py | 10 +- syftr/scripts/build_retriever_flow.py | 1 + syftr/studies.py | 4 +- syftr/tuner/qa_tuner.py | 19 ++- 6 files changed, 115 insertions(+), 94 deletions(-) diff --git a/syftr/evaluation.py b/syftr/evaluation.py index a3206033..fd211472 100644 --- a/syftr/evaluation.py +++ b/syftr/evaluation.py @@ -371,15 +371,17 @@ async def aretrieve_pair( flow: RetrieverFlow, rate_limiter: AsyncLimiter, raise_on_exception: bool | None = EVAL__RAISE_ON_EXCEPTION, -) -> T.Tuple[T.List[NodeWithScore] | None, float, Exception | None]: +) -> T.Tuple[ + T.List[NodeWithScore] | None, float, T.List[LLMCallData], Exception | None +]: """Get flow's retrieved documents from an Q&A pair asynchronously.""" - result, run_time, exception = await exception_catcher( + nodes, duration, call_data, exception = await exception_catcher( func=flow.aretrieve, return_values_on_exception=(None, np.nan), raise_on_exception=raise_on_exception, query=qa_pair.question, ) - return result, run_time, exception + return nodes, duration, call_data, exception async def _aeval_retriever_pair( @@ -392,8 +394,8 @@ async def _aeval_retriever_pair( """Evaluate retrieval performance on a single Q&A item.""" if not qa_pair.gold_evidence: raise ValueError("QAPair gold_evidence is empty: %s", qa_pair) - retrieval_results, run_time, retrieval_exception = await aretrieve_pair( - qa_pair, flow, rate_limiter + retrieval_results, run_time, call_data, retrieval_exception = await aretrieve_pair( + qa_pair, flow, rate_limiter, raise_on_exception ) if retrieval_results: retrieved_contexts: T.List[str] = [ @@ -412,7 +414,7 @@ async def _aeval_retriever_pair( run_time=run_time, generation_exception=retrieval_exception, evaluation_exception=None, - llm_call_data=[], + llm_call_data=call_data, retriever_recall=result.score, retriever_context_length=retrieved_contexts_length, passing=result.passing, diff --git a/syftr/flows.py b/syftr/flows.py index 24b1417e..544199e3 100644 --- a/syftr/flows.py +++ b/syftr/flows.py @@ -130,14 +130,16 @@ def generate( ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: invocation_id = uuid4().hex self._llm_call_data[invocation_id] = [] - response, duration = self._generate(query, invocation_id) + response, duration, retrieval_call_data = self._generate(query, invocation_id) call_data = self._llm_call_data.pop(invocation_id) + if retrieval_call_data: + call_data.extend(retrieval_call_data) return response, duration, call_data @dispatcher.span def _generate( self, query: str, invocation_id: str - ) -> T.Tuple[CompletionResponse, float]: + ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: assert self.response_synthesizer_llm is not None, ( "Response synthesizer LLM is not set. Cannot generate." ) @@ -145,21 +147,25 @@ def _generate( prompt = self.get_prompt(query) response: CompletionResponse = self.response_synthesizer_llm.complete(prompt) duration = time.perf_counter() - start_time - return response, duration + return response, duration, [] async def agenerate( self, query: str ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: invocation_id = uuid4().hex self._llm_call_data[invocation_id] = [] - response, duration = await self._agenerate(query, invocation_id) + response, duration, retrieval_call_data = await self._agenerate( + query, invocation_id + ) call_data = self._llm_call_data.pop(invocation_id) + if retrieval_call_data: + call_data.extend(retrieval_call_data) return response, duration, call_data @dispatcher.span async def _agenerate( self, query: str, invocation_id: str - ) -> T.Tuple[CompletionResponse, float]: + ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: assert self.response_synthesizer_llm is not None, ( "Response synthesizer LLM is not set. Cannot generate." ) @@ -169,7 +175,7 @@ async def _agenerate( prompt ) duration = time.perf_counter() - start_time - return response, duration + return response, duration, [] @dataclass(kw_only=True) @@ -228,75 +234,88 @@ def _generate(self, query: str, *args, **kwargs): async def _agenerate(self, query: str, *args, **kwargs): raise NotImplementedError("RetrieverFlow does not support generation.") - @dispatcher.span - def retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: + def _check_cache( + self, query: str + ) -> T.Optional[T.Tuple[T.List[NodeWithScore], float, T.List[LLMCallData]]]: + """Check the cache for the result of a query. + + Returns the retrieval information if present, otherwise, None. + """ if self.retriever_cache_fingerprint: - return self.cached_retrieve(query) - start_time = time.perf_counter() - qb = QueryBundle(query) - retrieval_result = self.query_engine.retrieve(qb) - duration = time.perf_counter() - start_time - return retrieval_result, duration + with get_retrieval_cache_key( + query, self.retriever_cache_fingerprint + ) as key: + if (retrieval_result := get_retrieval_cache(key)) is not None: + logger.info(f"Retriever cache hit: {query}") + nodes, duration, call_data = retrieval_result + return nodes, duration, call_data + else: + logger.info(f"Retriever cache miss: {query}") + return None + + def _put_cache( + self, + query: str, + nodes: T.List[NodeWithScore], + duration: float, + call_data: T.List[LLMCallData], + ) -> None: + """Store values in the retrieval cache.""" + if not self.retriever_cache_fingerprint: + return + with get_retrieval_cache_key(query, self.retriever_cache_fingerprint) as key: + put_retrieval_cache(key, (nodes, duration, call_data)) + + def retrieve( + self, query: str + ) -> T.Tuple[T.List[NodeWithScore], float, T.List[LLMCallData]]: + if (retrieval_result := self._check_cache(query)) is not None: + return retrieval_result + invocation_id = uuid4().hex + self._llm_call_data[invocation_id] = [] + nodes, duration = self._retrieve(query, invocation_id) + call_data = self._llm_call_data.pop(invocation_id) + self._put_cache(query, nodes, duration, call_data) + return nodes, duration, call_data @dispatcher.span - def cached_retrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: - assert self.retriever_cache_fingerprint, ( - "Retriever fingerprint must be set to use cache." - ) + def _retrieve( + self, query: str, invocation_id: str + ) -> T.Tuple[T.List[NodeWithScore], float]: start_time = time.perf_counter() - with get_retrieval_cache_key(query, self.retriever_cache_fingerprint) as key: - if (retrieval_result := get_retrieval_cache(key)) is not None: - logger.debug(f"Retriever cache hit: {query}") - else: - logger.debug(f"Retriever cache miss: {query}") - qb = QueryBundle(query) - retrieval_result = self.query_engine.retrieve(qb) - put_retrieval_cache(key, retrieval_result) + qb = QueryBundle(query) + nodes = self.query_engine.retrieve(qb) duration = time.perf_counter() - start_time - return retrieval_result, duration + return nodes, duration + + async def aretrieve( + self, query: str + ) -> T.Tuple[T.List[NodeWithScore], float, T.List[LLMCallData]]: + if (retrieval_result := self._check_cache(query)) is not None: + return retrieval_result + invocation_id = uuid4().hex + self._llm_call_data[invocation_id] = [] + nodes, duration = await self._aretrieve(query, invocation_id) + call_data = self._llm_call_data.pop(invocation_id) + self._put_cache(query, nodes, duration, call_data) + return nodes, duration, call_data @dispatcher.span - async def aretrieve(self, query: str) -> T.Tuple[T.List[NodeWithScore], float]: - if self.retriever_cache_fingerprint: - return await self.cached_aretrieve(query) + async def _aretrieve( + self, query: str, invocation_id: str + ) -> T.Tuple[T.List[NodeWithScore], float]: start_time = time.perf_counter() qb = QueryBundle(query) if isinstance(self.query_engine, TransformQueryEngine): # TransformQueryEngine does not have aretrieve method - retrieval_result = self.query_engine.retrieve(qb) + nodes = self.query_engine.retrieve(qb) else: assert hasattr(self.query_engine, "aretrieve"), ( f"{self.query_engine} does not have 'aretrieve' method" ) - retrieval_result = await self.query_engine.aretrieve(qb) + nodes = await self.query_engine.aretrieve(qb) duration = time.perf_counter() - start_time - return retrieval_result, duration - - @dispatcher.span - async def cached_aretrieve( - self, query: str - ) -> T.Tuple[T.List[NodeWithScore], float]: - assert self.retriever_cache_fingerprint, ( - "Retriever fingerprint must be set to use cache." - ) - start_time = time.perf_counter() - with get_retrieval_cache_key(query, self.retriever_cache_fingerprint) as key: - if (retrieval_result := get_retrieval_cache(key)) is not None: - logger.info(f"Retriever cache hit: {query}") - else: - logger.info(f"Retriever cache miss: {query}") - qb = QueryBundle(query) - if isinstance(self.query_engine, TransformQueryEngine): - # TransformQueryEngine does not have aretrieve method - retrieval_result = self.query_engine.retrieve(qb) - else: - assert hasattr(self.query_engine, "aretrieve"), ( - f"{self.query_engine} does not have 'aretrieve' method" - ) - retrieval_result = await self.query_engine.aretrieve(qb) - put_retrieval_cache(key, retrieval_result) - duration = time.perf_counter() - start_time - return retrieval_result, duration + return nodes, duration @dataclass(kw_only=True) @@ -338,9 +357,9 @@ def get_prompt(self, query) -> str: @dispatcher.span def _generate( self, query: str, invocation_id: str - ) -> T.Tuple[CompletionResponse, float]: + ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: + nodes, retrieval_duration, retrieval_call_data = self.retrieve(query) start_time = time.perf_counter() - nodes, _ = self.retrieve(query) response = self.query_engine.synthesize( query_bundle=QueryBundle(query), nodes=nodes, @@ -355,15 +374,15 @@ def _generate( **(response.metadata or {}), # type: ignore }, ) - duration = time.perf_counter() - start_time - return completion_response, duration + duration = time.perf_counter() - start_time + retrieval_duration + return completion_response, duration, retrieval_call_data @dispatcher.span async def _agenerate( self, query: str, invocation_id: str - ) -> T.Tuple[CompletionResponse, float]: + ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: + nodes, retrieval_duration, retrieval_call_data = await self.aretrieve(query) start_time = time.perf_counter() - nodes, _ = await self.aretrieve(query) response = await self.query_engine.asynthesize( query_bundle=QueryBundle(query), nodes=nodes, @@ -381,8 +400,8 @@ async def _agenerate( **(response.metadata or {}), }, ) - duration = time.perf_counter() - start_time - return completion_response, duration + duration = time.perf_counter() - start_time + retrieval_duration + return completion_response, duration, retrieval_call_data @dataclass(kw_only=True) @@ -488,7 +507,7 @@ def _generate( self, query: str, invocation_id: str, - ) -> T.Tuple[CompletionResponse, float]: + ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: start_time = time.perf_counter() response: AgentChatResponse = self.agent.chat(query) try: @@ -500,14 +519,14 @@ def _generate( logger.error("Incorrect response from an agent: %s", response) raise duration = time.perf_counter() - start_time - return completion_response, duration + return completion_response, duration, [] @dispatcher.span async def _agenerate( self, query: str, invocation_id: str, - ) -> T.Tuple[CompletionResponse, float]: + ) -> T.Tuple[CompletionResponse, float, T.List[LLMCallData]]: start_time = time.perf_counter() response: AgentChatResponse = await self.agent.achat(query) try: @@ -518,7 +537,7 @@ async def _agenerate( logger.error("Incorrect response from an agent: %s", response) raise duration = time.perf_counter() - start_time - return completion_response, duration + return completion_response, duration, [] @dataclass(kw_only=True) diff --git a/syftr/retrievers/cached_retriever.py b/syftr/retrievers/cached_retriever.py index 01c19f67..10b4ffa5 100644 --- a/syftr/retrievers/cached_retriever.py +++ b/syftr/retrievers/cached_retriever.py @@ -1,8 +1,8 @@ import hashlib import io import json +import typing as T from contextlib import contextmanager -from typing import Any, Dict, Optional import cloudpickle import diskcache @@ -21,7 +21,7 @@ @contextmanager -def get_retrieval_cache_key(question: str, retriever_params_dict: Dict[str, Any]): +def get_retrieval_cache_key(question: str, retriever_params_dict: T.Dict[str, T.Any]): """ Build a cache key from question text and retriever params. """ @@ -35,7 +35,7 @@ def get_retrieval_cache_key(question: str, retriever_params_dict: Dict[str, Any] def get_retriever_fingerprint( study_config: StudyConfig, params: ParamDict -) -> Dict[str, Any]: +) -> T.Dict[str, T.Any]: param_names = [ "hyde_enabled", "additional_context_enabled", @@ -78,7 +78,7 @@ def local_retrieval_cache(): yield cache -def put_retrieval_cache(cache_key: str, obj: Any, local_only: bool = False): +def put_retrieval_cache(cache_key: str, obj: T.Any, local_only: bool = False): """ Mirror to both diskcache & S3 under “retrieval_cache/{cache_key}.pkl”. """ @@ -112,7 +112,7 @@ def put_retrieval_cache(cache_key: str, obj: Any, local_only: bool = False): logger.info("Done storing object to S3") -def get_retrieval_cache(cache_key: str) -> Optional[Any]: +def get_retrieval_cache(cache_key: str) -> T.Optional[T.Any]: """ First check diskcache, then fall back to S3. Returns the retrieved list (NodeWithScore) or None if missing. diff --git a/syftr/scripts/build_retriever_flow.py b/syftr/scripts/build_retriever_flow.py index b1bf8118..ee9e4494 100644 --- a/syftr/scripts/build_retriever_flow.py +++ b/syftr/scripts/build_retriever_flow.py @@ -55,6 +55,7 @@ def main( retriever, docstore = build_dummy_retriever(study_config) flow = RetrieverFlow( response_synthesizer_llm=get_llm("gpt-4o-mini"), + template="default", retriever=retriever, docstore=docstore, hyde_llm=None, diff --git a/syftr/studies.py b/syftr/studies.py index 88a3ac5e..c7af44ee 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -1622,8 +1622,8 @@ class StudyConfig(BaseSettings): toy_mode: bool = Field( default=False, description="Whether to run in toy mode (with smaller dataset)." ) - cache_query_responses: bool = Field( - default=True, description="Whether to cache queries and retrieval responses." + retriever_cache_enabled: bool = Field( + default=False, description="Enable caching of queries and retrieval responses." ) model_config = SettingsConfigDict( diff --git a/syftr/tuner/qa_tuner.py b/syftr/tuner/qa_tuner.py index 76cebe08..7d24a78a 100644 --- a/syftr/tuner/qa_tuner.py +++ b/syftr/tuner/qa_tuner.py @@ -126,19 +126,23 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: response_synthesizer_llm = get_llm(params["response_synthesizer_llm"]) enforce_full_evaluation = params.get("enforce_full_evaluation", False) + if study_config.retriever_cache_enabled: + retriever_cache_fingerprint = get_retriever_fingerprint(study_config, params) + logger.info( + f"Retriever caching enabled with fingerprint {retriever_cache_fingerprint}" + ) + else: + retriever_cache_fingerprint = None + logger.info("Retriever caching disabled") if study_config.is_retriever_study: - retriever_cache_fingerprint = ( - get_retriever_fingerprint(study_config, params) - if study_config.cache_query_responses - else None - ) hyde_llm = ( get_llm(params["hyde_llm_name"]) if params.get("hyde_enabled") else None ) retriever, docstore = build_rag_retriever(study_config, params) return RetrieverFlow( response_synthesizer_llm=response_synthesizer_llm, + template=params.get("template_name", "default"), retriever=retriever, docstore=docstore, hyde_llm=hyde_llm, @@ -179,11 +183,6 @@ def build_flow(params: T.Dict, study_config: StudyConfig) -> Flow: enforce_full_evaluation=enforce_full_evaluation, ) else: - retriever_cache_fingerprint = ( - get_retriever_fingerprint(study_config, params) - if study_config.cache_query_responses - else None - ) hyde_llm = reranker_llm = reranker_top_k = None if params.get("hyde_enabled"): hyde_llm = get_llm(params["hyde_llm_name"]) From 1244940ce142e954fbddb5d354be9efe1de647fc Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 23 Jun 2025 12:07:27 -0300 Subject: [PATCH 09/14] param_override --- syftr/optimization.py | 2 +- syftr/studies.py | 7 +++---- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/syftr/optimization.py b/syftr/optimization.py index 614e90b4..e28f0156 100644 --- a/syftr/optimization.py +++ b/syftr/optimization.py @@ -307,7 +307,7 @@ def run(self) -> optuna.Study: block.name, defaults, ) - study_config.search_space.custom_defaults.update(defaults) + study_config.search_space.param_override.update(defaults) future.result() # Wait for seeding to finish diff --git a/syftr/studies.py b/syftr/studies.py index c7af44ee..42a8abc7 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -945,9 +945,9 @@ class SearchSpace(BaseModel): default_factory=LATSRagAgent, description="Configuration for the LATS RAG agent.", ) - custom_defaults: ParamDict = Field( + param_override: ParamDict = Field( default_factory=dict, - description="Override default parameters for the search space.", + description="Override parameters for the search space.", ) def defaults(self) -> ParamDict: @@ -1118,8 +1118,7 @@ def sample(self, trial: Trial, parameters: T.List[str] = PARAMETERS) -> ParamDic if params["few_shot_enabled"]: params.update(**self.few_shot_retriever.defaults()) - # Use custom defaults to override parameters - params.update(self.custom_defaults) + params.update(self.param_override) return params From 7f6e126cf6908b1db76375a948b5e85b4effd738 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 23 Jun 2025 12:22:50 -0300 Subject: [PATCH 10/14] Remove no-retriever studies. --- studies/no-retriever-financebench.yaml | 83 -------------------------- studies/no-retriever-hotpotqa.yaml | 80 ------------------------- studies/no-retriever-multihoprag.yaml | 73 ---------------------- 3 files changed, 236 deletions(-) delete mode 100644 studies/no-retriever-financebench.yaml delete mode 100644 studies/no-retriever-hotpotqa.yaml delete mode 100644 studies/no-retriever-multihoprag.yaml diff --git a/studies/no-retriever-financebench.yaml b/studies/no-retriever-financebench.yaml deleted file mode 100644 index 8d4120fb..00000000 --- a/studies/no-retriever-financebench.yaml +++ /dev/null @@ -1,83 +0,0 @@ -name: "no-retriever-financebench" -dataset: - dataset_dir: partitioned - description: Financial dataset that contains everything about finance, including - real-world financial documents, SEC filings, earning reports, call transcripts, - and much more. It has all the financial live data, historical data, just about - everything about finance, for instance, definitions and explanations of financial - term, insights on company revenues, mergers, founders, or stock performance, details - on financial laws, compliance, or government policies, information required to - evaluate finance risk, and information about banking operations, credit systems, - or loan structures. - examples_data_path: examples - grounding_data_path: aryn_html - load_examples_timeout_s: 3600 - load_grounding_data_timeout_s: 3600 - partition_map: - holdout: holdout - sample: pepsi - test: test - train: train - path_root: octo-syftr-benchmarking-data - storage_options: - cache_check: 60 - cache_storage: benchmarking/data/cache - check_files: true - expiry_time: 2592000 - protocol: filecache - same_names: true - target_protocol: s3 - storage_partitions: - - sample - - train - - test - - holdout - xname: financebench_hf -reuse_study: false -search_space: - custom_defaults: - rag_mode: rag - hyde_enabled: true - additional_context_enabled: true - rag_method: dense - rag_query_decomposition_enabled: true - rag_top_k: 19 - rag_embedding_model: BAAI/bge-small-en-v1.5 - rag_query_decomposition_llm_name: gpt-4o-mini - rag_query_decomposition_num_queries: 6 - rag_fusion_mode: simple - splitter_method: sentence - splitter_chunk_exp: 12 - splitter_chunk_overlap_frac: 0 - hyde_llm_name: gpt-4o-mini - additional_context_num_nodes: 16 -optimization: - baselines: [] - use_individual_baselines: false - use_agent_baselines: false - use_variations_of_baselines: false - shuffle_baselines: false - blocks: - - components: - - few_shot_retriever - - reranker - - sub_question_rag - - critique_rag_agent - - lats_rag_agent - - react_rag_agent - - response_synthesizer_llm - - template_name - name: no-retriever - num_trials: 500 - cpus_per_trial: 5 - embedding_device: cuda - # use_hf_embedding_models: true - gpus_per_trial: 0.2 - max_concurrent_trials: 10 - num_eval_samples: 100 - num_eval_batch: 100 - num_retries_unique_params: 10 - num_trials: 100 - use_cost_pruner: false - use_pareto_pruner: false - use_runtime_pruner: false diff --git a/studies/no-retriever-hotpotqa.yaml b/studies/no-retriever-hotpotqa.yaml deleted file mode 100644 index b77dadc5..00000000 --- a/studies/no-retriever-hotpotqa.yaml +++ /dev/null @@ -1,80 +0,0 @@ -name: "no-retriever-hotpotqa" -dataset: - dataset_dir: partitioned - description: This dataset is a vast collection of all kind of information that you - can find on Wikipedia. It can be used, for instance, to retrieve straightforward - facts from one or more documents, compare two entities based on shared attributes, - identify relationships, roles, or attributes of entities, reason about dates, - timelines, or chronological order, determine geographical relationships or locations, - explain causes or sequences of events or processes, synthesize facts from multiple - documents to infer answers, and validate or refute premises in the context of - the question. - examples_data_path: examples - grounding_data_path: grounding_data - load_examples_timeout_s: 3600 - load_grounding_data_timeout_s: 3600 - partition_map: - holdout: holdout - sample: sample - test: test - train: train - path_root: octo-syftr-benchmarking-data - storage_options: - cache_check: 60 - cache_storage: benchmarking/data/cache - check_files: true - expiry_time: 2592000 - protocol: filecache - same_names: true - target_protocol: s3 - storage_partitions: - - sample - - train - - test - - holdout - subset: train_hard - xname: hotpotqa_hf -reuse_study: false -optimization: - baselines: [] - use_individual_baselines: false - use_agent_baselines: false - use_variations_of_baselines: false - shuffle_baselines: false - blocks: - - components: - - few_shot_retriever - - reranker - - rag_mode - - sub_question_rag - - critique_rag_agent - - lats_rag_agent - - react_rag_agent - - response_synthesizer_llm - - template_name - name: no-retriever - num_trials: 500 - custom_defaults: - hyde_enabled: false - additional_context_enabled: true - rag_method: hybrid - rag_query_decomposition_enabled: false - rag_top_k: 19 - rag_embedding_model: avsolatorio/GIST-large-Embedding-v0 - rag_hybrid_bm25_weight: 0.6 - rag_query_decomposition_num_queries: 6 - rag_fusion_mode: dist_based_score - splitter_method: html - splitter_chunk_exp: 10 - splitter_chunk_overlap_frac: 0.25 - hyde_llm_name: llama-33-70B - additional_context_num_nodes: 7 - cpus_per_trial: 5 - # embedding_device: cuda - # use_hf_embedding_models: true - gpus_per_trial: 0.2 - max_concurrent_trials: 50 - num_eval_samples: 100 - num_eval_batch: 10 - num_retries_unique_params: 10 - num_trials: 500 diff --git a/studies/no-retriever-multihoprag.yaml b/studies/no-retriever-multihoprag.yaml deleted file mode 100644 index 1239f887..00000000 --- a/studies/no-retriever-multihoprag.yaml +++ /dev/null @@ -1,73 +0,0 @@ -name: "no-retriever-multihoprag" -dataset: - dataset_dir: partitioned - description: multihoprag dataset - examples_data_path: examples - grounding_data_path: grounding_data - load_examples_timeout_s: 3600 - load_grounding_data_timeout_s: 3600 - partition_map: - holdout: holdout - sample: sample - test: test - train: train - path_root: octo-syftr-benchmarking-data - storage_options: - cache_check: 60 - cache_storage: benchmarking/data/cache - check_files: true - expiry_time: 2592000 - protocol: filecache - same_names: true - target_protocol: s3 - storage_partitions: - - sample - - train - - test - - holdout - xname: multihoprag_hf -reuse_study: true -toy_mode: false -search_space: - custom_defaults: - rag_mode: rag - hyde_enabled: false - additional_context_enabled: false - rag_method: hybrid - rag_query_decomposition_enabled: true - rag_top_k: 19 - rag_embedding_model: Labib11/MUG-B-1.6 - rag_hybrid_bm25_weight: 0.1 - rag_query_decomposition_llm_name: gpt-4o-mini - rag_query_decomposition_num_queries: 6 - rag_fusion_mode: reciprocal_rerank - splitter_method: sentence - splitter_chunk_exp: 12 - splitter_chunk_overlap_frac: 0 -optimization: - baselines: [] - use_individual_baselines: false - use_agent_baselines: false - use_variations_of_baselines: false - shuffle_baselines: false - blocks: - - components: - - few_shot_retriever - - reranker - - sub_question_rag - - critique_rag_agent - - lats_rag_agent - - react_rag_agent - - response_synthesizer_llm - - template_name - name: no-retriever - num_trials: 500 - cpus_per_trial: 5 - embedding_device: cuda - # use_hf_embedding_models: true - gpus_per_trial: 0.2 - max_concurrent_trials: 50 - num_eval_samples: 200 - num_eval_batch: 20 - num_retries_unique_params: 10 - num_trials: 500 From a18885c506959c7e3f0aa2796cd53c59e053a698 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Mon, 23 Jun 2025 13:00:24 -0300 Subject: [PATCH 11/14] ruff --- syftr/ray/utils.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/syftr/ray/utils.py b/syftr/ray/utils.py index bbe1b14d..9f797e2d 100644 --- a/syftr/ray/utils.py +++ b/syftr/ray/utils.py @@ -1,7 +1,8 @@ +import functools import getpass import os -import functools from typing import Any, Final + import ray from syftr.configuration import cfg From 3c1bbdddcbb767645fdac773e87166ac3ff50f93 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Wed, 25 Jun 2025 16:54:18 -0300 Subject: [PATCH 12/14] Further updates --- syftr/custom_metrics.py | 262 +++++++++++++++++++++++++++ syftr/evaluation.py | 130 +------------ syftr/retrievers/cached_retriever.py | 5 +- syftr/retrievers/storage.py | 7 +- syftr/studies.py | 2 +- 5 files changed, 273 insertions(+), 133 deletions(-) diff --git a/syftr/custom_metrics.py b/syftr/custom_metrics.py index f2d926dd..3048e9ff 100644 --- a/syftr/custom_metrics.py +++ b/syftr/custom_metrics.py @@ -1,10 +1,16 @@ +import asyncio +import math import typing as T import numpy as np +from llama_index.core.evaluation import BaseEvaluator +from llama_index.core.evaluation.base import EvaluationResult from llama_index.core.evaluation.retrieval.metrics_base import ( BaseRetrievalMetric, RetrievalMetricResult, ) +from llama_index.core.prompts.mixin import PromptDictType +from rapidfuzz.fuzz import partial_ratio from rouge_score import rouge_scorer @@ -87,3 +93,259 @@ def lognormal_confidence(values: T.List[float], zscore: float) -> float: if len(values) == 0: return np.nan return zscore * float(np.std(values, ddof=1)) / np.sqrt(len(values)) + + +class ExactMatchEvaluator(BaseEvaluator): + """ + Evaluator that calculates exact match by comparing reference contexts + with retrieved contexts. + """ + + async def aevaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + """ + Evaluate exact match by computing the proportion of reference contexts + that are present in the retrieved contexts. + """ + reference = kwargs.get("reference") + + if not reference: + raise ValueError("Reference contexts are empty.") + if not contexts: + raise ValueError("Retrieved contexts are empty.") + + matched = sum(any(ref in context for context in contexts) for ref in reference) + recall = matched / len(reference) if reference else 0.0 + return EvaluationResult( + passing=recall > 0, + score=recall, + ) + + def evaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + """ + Synchronous version of the evaluation method for compatibility with base class. + """ + loop = asyncio.get_event_loop() + return loop.run_until_complete( + self.aevaluate(query, response, contexts, **kwargs) + ) + + def _get_prompts(self) -> PromptDictType: + """Get prompts.""" + return {} + + def _update_prompts(self, prompts_dict: PromptDictType) -> None: + """Update prompts.""" + pass + + +class FuzzyRecallEvaluator(BaseEvaluator): + """ + Evaluator that calculates fuzzy recall by comparing reference contexts + with retrieved contexts using partial_ratio from rapidfuzz. + """ + + def __init__(self, threshold: float = 90.0): + self.threshold = threshold + + async def fuzzy_match_async(self, ref: str, doc: str) -> bool: + return await asyncio.to_thread(partial_ratio, ref, doc) >= self.threshold + + async def fuzzy_contains_async(self, ref: str, docs: T.Sequence[str]) -> bool: + tasks = [self.fuzzy_match_async(ref, doc) for doc in docs] + for coro in asyncio.as_completed(tasks): + if await coro: + return True + return False + + async def aevaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + """ + Evaluate fuzzy recall by computing the proportion of reference contexts + that have a fuzzy match in the retrieved contexts. + """ + reference = kwargs.get("reference") + + if not reference: + raise ValueError("Reference contexts are empty.") + if not contexts: + raise ValueError("Retrieved contexts are empty.") + + tasks = [self.fuzzy_contains_async(ref, contexts) for ref in reference] + results = await asyncio.gather(*tasks) + matched = sum(results) + recall = matched / len(reference) if reference else 0.0 + return EvaluationResult( + passing=recall > 0, + score=recall, + ) + + def evaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + """ + Synchronous version of the evaluation method for compatibility with base class. + """ + loop = asyncio.get_event_loop() + return loop.run_until_complete( + self.aevaluate(query, response, contexts, **kwargs) + ) + + def _get_prompts(self) -> PromptDictType: + """Get prompts.""" + return {} + + def _update_prompts(self, prompts_dict: PromptDictType) -> None: + """Update prompts.""" + pass + + +class MRREvaluator(BaseEvaluator): + """ + Evaluator that calculates Mean Reciprocal Rank (MRR) for a single query by + finding the first matching reference in the retrieved contexts. + """ + + async def aevaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + """ + Evaluate reciprocal rank of the first relevant document. + Assumes `reference` is a list of correct answers and `contexts` is + a list of retrieved documents ordered by relevance. + """ + reference = kwargs.get("reference") + + if not reference: + raise ValueError("Reference contexts are empty.") + if not contexts: + raise ValueError("Retrieved contexts are empty.") + + reciprocal_rank = 0.0 + for i, context in enumerate(contexts): + if any(ref in context for ref in reference): + reciprocal_rank = 1.0 / (i + 1) + break + + return EvaluationResult( + passing=reciprocal_rank > 0, + score=reciprocal_rank, + ) + + def evaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + loop = asyncio.get_event_loop() + return loop.run_until_complete( + self.aevaluate(query, response, contexts, **kwargs) + ) + + def _get_prompts(self) -> PromptDictType: + return {} + + def _update_prompts(self, prompts_dict: PromptDictType) -> None: + pass + + +class NDCGEvaluator(BaseEvaluator): + """ + Evaluator that calculates Normalized Discounted Cumulative Gain (NDCG) + based on relevance of retrieved contexts. + """ + + def __init__(self, k: int = 10): + self.k = k + + def _dcg(self, relevance_scores: T.List[float]) -> float: + return sum(rel / math.log2(i + 2) for i, rel in enumerate(relevance_scores)) + + def _get_relevance( + self, context: str, reference: T.Union[T.Sequence[str], T.Mapping[str, float]] + ) -> float: + if isinstance(reference, dict): + # Graded relevance + return max( + (score for ref, score in reference.items() if ref in context), + default=0.0, + ) + else: + # Binary relevance + return 1.0 if any(ref in context for ref in reference) else 0.0 + + async def aevaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + reference = kwargs.get("reference") + + if not reference: + raise ValueError("Reference contexts are empty.") + if not contexts: + raise ValueError("Retrieved contexts are empty.") + + top_k_contexts = contexts[: self.k] + relevance_scores = [self._get_relevance(c, reference) for c in top_k_contexts] + dcg = self._dcg(relevance_scores) + + # Ideal DCG: sorted relevance scores + if isinstance(reference, dict): + ideal_scores = sorted(reference.values(), reverse=True)[: self.k] + else: + ideal_scores = [1.0] * min(len(reference), self.k) + + idcg = self._dcg(ideal_scores) + ndcg = dcg / idcg if idcg > 0 else 0.0 + + return EvaluationResult( + passing=ndcg > 0, + score=ndcg, + ) + + def evaluate( + self, + query: T.Optional[str] = None, + response: T.Optional[str] = None, + contexts: T.Optional[T.Sequence[str]] = None, + **kwargs: T.Any, + ) -> EvaluationResult: + loop = asyncio.get_event_loop() + return loop.run_until_complete( + self.aevaluate(query, response, contexts, **kwargs) + ) + + def _get_prompts(self) -> PromptDictType: + return {} + + def _update_prompts(self, prompts_dict: PromptDictType) -> None: + pass diff --git a/syftr/evaluation.py b/syftr/evaluation.py index fd211472..22edfaad 100644 --- a/syftr/evaluation.py +++ b/syftr/evaluation.py @@ -22,10 +22,8 @@ from llama_index.core.evaluation import BaseEvaluator, CorrectnessEvaluator from llama_index.core.evaluation.base import EvaluationResult from llama_index.core.llms import CompletionResponse -from llama_index.core.prompts.mixin import PromptDictType from llama_index.core.schema import MetadataMode, NodeWithScore from optuna import TrialPruned -from rapidfuzz.fuzz import partial_ratio from tenacity import ( before_sleep_log, retry, @@ -38,6 +36,7 @@ from syftr import core, custom_metrics from syftr.configuration import EVAL__RAISE_ON_EXCEPTION +from syftr.custom_metrics import ExactMatchEvaluator from syftr.flows import Flow, RetrieverFlow from syftr.helpers import get_exception_report from syftr.instrumentation.tokens import LLMCallData @@ -74,131 +73,6 @@ ) -class ExactMatchEvaluator(BaseEvaluator): - """ - Evaluator that calculates exact match by comparing reference contexts - with retrieved contexts. - """ - - async def aevaluate( - self, - query: T.Optional[str] = None, - response: T.Optional[str] = None, - contexts: T.Optional[T.Sequence[str]] = None, - **kwargs: T.Any, - ) -> EvaluationResult: - """ - Evaluate exact match by computing the proportion of reference contexts - that are present in the retrieved contexts. - """ - reference = kwargs.get("reference") - - if not reference: - raise ValueError("Reference contexts are empty.") - if not contexts: - raise ValueError("Retrieved contexts are empty.") - - matched = sum(any(ref in context for context in contexts) for ref in reference) - recall = matched / len(reference) if reference else 0.0 - return EvaluationResult( - passing=recall > 0, - score=recall, - ) - - def evaluate( - self, - query: T.Optional[str] = None, - response: T.Optional[str] = None, - contexts: T.Optional[T.Sequence[str]] = None, - **kwargs: T.Any, - ) -> EvaluationResult: - """ - Synchronous version of the evaluation method for compatibility with base class. - """ - loop = asyncio.get_event_loop() - return loop.run_until_complete( - self.aevaluate(query, response, contexts, **kwargs) - ) - - def _get_prompts(self) -> PromptDictType: - """Get prompts.""" - return {} - - def _update_prompts(self, prompts_dict: PromptDictType) -> None: - """Update prompts.""" - pass - - -class FuzzyRecallEvaluator(BaseEvaluator): - """ - Evaluator that calculates fuzzy recall by comparing reference contexts - with retrieved contexts using partial_ratio from rapidfuzz. - """ - - def __init__(self, threshold: float = 90.0): - self.threshold = threshold - - async def fuzzy_match_async(self, ref: str, doc: str) -> bool: - return await asyncio.to_thread(partial_ratio, ref, doc) >= self.threshold - - async def fuzzy_contains_async(self, ref: str, docs: T.Sequence[str]) -> bool: - tasks = [self.fuzzy_match_async(ref, doc) for doc in docs] - for coro in asyncio.as_completed(tasks): - if await coro: - return True - return False - - async def aevaluate( - self, - query: T.Optional[str] = None, - response: T.Optional[str] = None, - contexts: T.Optional[T.Sequence[str]] = None, - **kwargs: T.Any, - ) -> EvaluationResult: - """ - Evaluate fuzzy recall by computing the proportion of reference contexts - that have a fuzzy match in the retrieved contexts. - """ - reference = kwargs.get("reference") - - if not reference: - raise ValueError("Reference contexts are empty.") - if not contexts: - raise ValueError("Retrieved contexts are empty.") - - tasks = [self.fuzzy_contains_async(ref, contexts) for ref in reference] - results = await asyncio.gather(*tasks) - matched = sum(results) - recall = matched / len(reference) if reference else 0.0 - return EvaluationResult( - passing=recall > 0, - score=recall, - ) - - def evaluate( - self, - query: T.Optional[str] = None, - response: T.Optional[str] = None, - contexts: T.Optional[T.Sequence[str]] = None, - **kwargs: T.Any, - ) -> EvaluationResult: - """ - Synchronous version of the evaluation method for compatibility with base class. - """ - loop = asyncio.get_event_loop() - return loop.run_until_complete( - self.aevaluate(query, response, contexts, **kwargs) - ) - - def _get_prompts(self) -> PromptDictType: - """Get prompts.""" - return {} - - def _update_prompts(self, prompts_dict: PromptDictType) -> None: - """Update prompts.""" - pass - - class SyftrEvaluationResult(EvaluationResult): class Config: arbitrary_types_allowed = True @@ -830,7 +704,7 @@ def eval_dataset( items=filtered_dataset, flow=flow, study_config=study_config, - evaluators=[FuzzyRecallEvaluator()], + evaluators=[ExactMatchEvaluator()], raise_on_exception=study_config.evaluation.raise_on_exception, pruner=pruner, timeout_pruner=timeout, diff --git a/syftr/retrievers/cached_retriever.py b/syftr/retrievers/cached_retriever.py index 10b4ffa5..3fa190c5 100644 --- a/syftr/retrievers/cached_retriever.py +++ b/syftr/retrievers/cached_retriever.py @@ -46,6 +46,7 @@ def get_retriever_fingerprint( "rag_query_decomposition_llm_name", "rag_query_decomposition_num_queries", "rag_fusion_mode", + "rag_hybrid_bm25_weight", "splitter_method", "splitter_chunk_exp", "splitter_chunk_overlap_frac", @@ -92,8 +93,8 @@ def put_retrieval_cache(cache_key: str, obj: T.Any, local_only: bool = False): try: logger.info(f"Storing {cache_key} to Ray cache") ray_cache_put(cache_key, serialized) - except Exception as e: - logger.warning(f"Skipping Ray cache put due to error: {e}") + except Exception: + pass # S3 mirror if not local_only and cfg.storage.s3_cache_enabled: diff --git a/syftr/retrievers/storage.py b/syftr/retrievers/storage.py index f7da842e..a6d3be59 100644 --- a/syftr/retrievers/storage.py +++ b/syftr/retrievers/storage.py @@ -115,6 +115,10 @@ def put_cache(cache_key, index, local_only: bool = False) -> None: def get_cached(cache_key: str) -> Optional[Any]: s3_cache_key = f"index_cache/{cache_key}.pkl" + try: + ray_data = ray_cache_get(cache_key) + except Exception: + ray_data = None try: with local_cache() as cache: @@ -124,8 +128,7 @@ def get_cached(cache_key: str) -> Optional[Any]: logger.info(f"Loaded pre-built index from {cache.directory}") return index - data = ray_cache_get(cache_key) - if data is not None: + if ray_data is not None: logger.info(f"Loading {cache_key} from Ray cache") return cloudpickle.loads(decompress(data)) diff --git a/syftr/studies.py b/syftr/studies.py index 42a8abc7..6f43827a 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -1622,7 +1622,7 @@ class StudyConfig(BaseSettings): default=False, description="Whether to run in toy mode (with smaller dataset)." ) retriever_cache_enabled: bool = Field( - default=False, description="Enable caching of queries and retrieval responses." + default=True, description="Enable caching of queries and retrieval responses." ) model_config = SettingsConfigDict( From 3db8110eb40ba70e903111bbf32a8cbd9b9bfe66 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Thu, 26 Jun 2025 10:24:52 -0300 Subject: [PATCH 13/14] Bugfix --- studies/retriever-only-financebench.yaml | 11 ++++++++--- syftr/retrievers/storage.py | 2 +- syftr/scripts/clear_ray_cache.py | 11 +++++++++++ 3 files changed, 20 insertions(+), 4 deletions(-) create mode 100644 syftr/scripts/clear_ray_cache.py diff --git a/studies/retriever-only-financebench.yaml b/studies/retriever-only-financebench.yaml index 0650c8c7..c30bb8f8 100644 --- a/studies/retriever-only-financebench.yaml +++ b/studies/retriever-only-financebench.yaml @@ -1,4 +1,4 @@ -name: "retriever-only-financebench" +name: "retriever-only-financebench-v2" dataset: dataset_dir: partitioned description: Financial dataset that contains everything about finance, including @@ -104,11 +104,16 @@ search_space: num_queries_max: 20 num_queries_min: 2 num_queries_step: 2 + top_k: + kmax: 128 + kmin: 2 + log: true + step: 1 splitter: chunk_overlap_frac_max: 0.75 chunk_overlap_frac_min: 0.0 chunk_overlap_frac_step: 0.25 - chunk_max_exp: 12 + chunk_max_exp: 13 chunk_min_exp: 9 methods: - html @@ -121,7 +126,7 @@ optimization: embedding_device: cuda # use_hf_embedding_models: true gpus_per_trial: 0.2 - max_concurrent_trials: 50 + max_concurrent_trials: 40 num_eval_samples: 100 num_eval_batch: 10 num_retries_unique_params: 10 diff --git a/syftr/retrievers/storage.py b/syftr/retrievers/storage.py index a6d3be59..97f05d8a 100644 --- a/syftr/retrievers/storage.py +++ b/syftr/retrievers/storage.py @@ -130,7 +130,7 @@ def get_cached(cache_key: str) -> Optional[Any]: if ray_data is not None: logger.info(f"Loading {cache_key} from Ray cache") - return cloudpickle.loads(decompress(data)) + return cloudpickle.loads(decompress(ray_data)) if cfg.storage.s3_cache_enabled: if (data := get_file_from_s3(s3_cache_key)) is not None: diff --git a/syftr/scripts/clear_ray_cache.py b/syftr/scripts/clear_ray_cache.py new file mode 100644 index 00000000..0ce1d726 --- /dev/null +++ b/syftr/scripts/clear_ray_cache.py @@ -0,0 +1,11 @@ +from syftr.ray.utils import ray_cache_restart, ray_init + + +def main(): + ray_init(force_remote=True) + ray_cache_restart() + print("Ray cache cleared successfully.") + + +if __name__ == "__main__": + main() From acb2946080bc201f0636678ad0e32ca0be788be6 Mon Sep 17 00:00:00 2001 From: Matthew Hausknecht Date: Thu, 26 Jun 2025 12:41:50 -0300 Subject: [PATCH 14/14] Change top-k defaults --- syftr/studies.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/syftr/studies.py b/syftr/studies.py index 6f43827a..b644b5ba 100644 --- a/syftr/studies.py +++ b/syftr/studies.py @@ -457,7 +457,7 @@ def get_cardinality(self) -> int: class Retriever(BaseModel, SearchSpaceMixin): top_k: TopK = Field( - default_factory=TopK, + default_factory=lambda: TopK(kmax=64, log=True), description="Configuration for the number of items to retrieve.", ) methods: T.List[str] = Field( @@ -606,7 +606,7 @@ class Reranker(BaseModel, SearchSpaceMixin): """ top_k: TopK = Field( - default_factory=lambda: TopK(kmax=128, log=True), + default_factory=lambda: TopK(kmax=32, log=True), description="Configuration for the number of items to rerank.", ) llms: T.List[str] = Field( @@ -904,7 +904,7 @@ class SearchSpace(BaseModel): description="Configuration for the few-shot retriever.", ) rag_retriever: Retriever = Field( - default_factory=lambda: Retriever(top_k=TopK(kmax=128, log=True)), + default_factory=Retriever, description="Configuration for the RAG retriever.", ) splitter: Splitter = Field( @@ -1179,7 +1179,7 @@ class RetrieverSearchSpace(BaseModel): description="LLMs used for response synthesis.", ) rag_retriever: Retriever = Field( - default_factory=lambda: Retriever(top_k=TopK(kmax=128, log=True)), + default_factory=Retriever, description="Configuration for the RAG retriever.", ) splitter: Splitter = Field(