From 18e7b7e14c10008e762840de5f8ab0564dd9600d Mon Sep 17 00:00:00 2001 From: Nikita Kulin Date: Sun, 27 Sep 2026 18:57:35 +0300 Subject: [PATCH 1/4] SAPO --- README.md | 22 +- coolprompt/meta_selector/selector.py | 4 +- .../method_evaluation/method_evaluation.py | 4 +- coolprompt/optimizer/sapo/__init__.py | 6 + coolprompt/optimizer/sapo/prompt_templates.py | 63 +++ coolprompt/optimizer/sapo/sapo.py | 386 ++++++++++++++++++ coolprompt/utils/var_validation.py | 2 + docs/API.md | 5 +- test/coolprompt/optimizer/test_sapo.py | 89 ++++ .../test_smoke_optimizer_interface.py | 6 + test/coolprompt/test_prompt_tuner_auto.py | 13 +- 11 files changed, 585 insertions(+), 15 deletions(-) create mode 100644 coolprompt/optimizer/sapo/__init__.py create mode 100644 coolprompt/optimizer/sapo/prompt_templates.py create mode 100644 coolprompt/optimizer/sapo/sapo.py create mode 100644 test/coolprompt/optimizer/test_sapo.py diff --git a/README.md b/README.md index f89eb9f..159bfb7 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,7 @@ CoolPrompt is a framework for automatic prompt creation and optimization. - RE-GPS - RIDER - BRAVE + - SAPO - PromptCompressor - *(legacy/deprecated)*: ReflectivePrompt, DistillPrompt - **LLM-Agnostic Choice:** work with your custom llm (from open-sourced to proprietary) using [supported Langchain LLMs](https://python.langchain.com/docs/integrations/llms/) @@ -70,6 +71,7 @@ Compared metrics: | `regps` | Required | High | Very High | High | | `rider` | Required | Very High | Very High | Very High | | `brave` | Required | High | Very High | Budget-controlled | +| `sapo` | Required | High | Very High | High | | `compress` | None | Low | Medium | Low | | `reflective` | Required | High | High | High | | `distill` | Required | High | High | High | @@ -144,11 +146,25 @@ final_prompt = prompt_tuner.run( ) ``` -The bundled metadata contains SAPO, RIDER, and HyPER results. CoolPrompt does -not yet implement SAPO, so an SAPO recommendation is transparently executed -with `hyper`; the fallback reason is recorded in `meta_selection`. Pass +The bundled metadata contains SAPO, RIDER, and HyPER results. A SAPO +recommendation is executed directly with the built-in segment-based optimizer. Pass `meta_classifier_path="/path/to/metadata.csv"` to use a custom CSV. Without `dataset` and `target`, `method="auto"` uses `hyper_light`. + +Run SAPO directly when you want contrastive, segment-level prompt refinement: + +```python +final_prompt = prompt_tuner.run( + "Summarize the article concisely.", + task="generation", + dataset=["Article one", "Article two"], + target=["Summary one", "Summary two"], + method="sapo", + validation_size=0.5, + n_iterations=3, + n_candidates=4, +) +``` ## Examples diff --git a/coolprompt/meta_selector/selector.py b/coolprompt/meta_selector/selector.py index ade0b6b..f3a0717 100644 --- a/coolprompt/meta_selector/selector.py +++ b/coolprompt/meta_selector/selector.py @@ -57,7 +57,7 @@ def to_dict(self) -> dict[str, Any]: class APOMetaSelector: - """Choose RIDER or HyPER from quality, cost, and runtime metadata. + """Choose SAPO, RIDER, or HyPER from quality, cost, and runtime metadata. The selector keeps method choice independent from the caller's LLM: model recommendations are returned for diagnostics only and never replace it. @@ -371,5 +371,5 @@ def _map_method(method: str) -> tuple[str, str | None]: if method == "HyPER": return "hyper", None if method == "SAPO": - return "hyper", "SAPO is not implemented in CoolPrompt" + return "sapo", None return "hyper", f"Unsupported recommendation '{method}'; defaulted to HyPER." diff --git a/coolprompt/method_evaluation/method_evaluation.py b/coolprompt/method_evaluation/method_evaluation.py index 6c803ad..b0882bc 100644 --- a/coolprompt/method_evaluation/method_evaluation.py +++ b/coolprompt/method_evaluation/method_evaluation.py @@ -12,6 +12,7 @@ from coolprompt.optimizer.reflective_prompt import ReflectiveMethod from coolprompt.optimizer.regps import ReGPSMethod from coolprompt.optimizer.rider import RIDERGenesisMethod +from coolprompt.optimizer.sapo import SAPOMethod _BENCHMARK_IMPL: dict[str, AutoPromptingMethod] = { "hyper_light": HyPERLightMethod, @@ -23,6 +24,7 @@ "regps": ReGPSMethod, "rider": RIDERGenesisMethod, "brave": BRAVEMethod, + "sapo": SAPOMethod, } @@ -39,7 +41,7 @@ def evaluate_method( Args: method: One of ``hyper_light``, ``hyper``, ``reflective`` / ``reflectiveprompt``, - ``distill``, ``compress``, ``regps``, ``rider``, ``brave`` + ``distill``, ``compress``, ``regps``, ``rider``, ``brave``, ``sapo`` (same names as in ``PromptTuner`` / ``validate_method`` where applicable). model: LangChain language model used for optimization and evaluation. diff --git a/coolprompt/optimizer/sapo/__init__.py b/coolprompt/optimizer/sapo/__init__.py new file mode 100644 index 0000000..972f7b7 --- /dev/null +++ b/coolprompt/optimizer/sapo/__init__.py @@ -0,0 +1,6 @@ +"""Public SAPO optimizer exports.""" + +from .sapo import SAPOMethod, SAPOOptimizer + +__all__ = ["SAPOMethod", "SAPOOptimizer"] + diff --git a/coolprompt/optimizer/sapo/prompt_templates.py b/coolprompt/optimizer/sapo/prompt_templates.py new file mode 100644 index 0000000..6e51a38 --- /dev/null +++ b/coolprompt/optimizer/sapo/prompt_templates.py @@ -0,0 +1,63 @@ +"""Meta-prompts used by the segment-based SAPO optimizer.""" + +SEGMENTATION_TEMPLATE = """You are an expert prompt engineer. +Decompose the prompt below into four segments. Return JSON only with string +fields: role, context, tasks, output_format. Use an empty string when a segment +is absent. Do not invent text that is not present in the prompt. + +Prompt: +\"\"\" +{prompt} +\"\"\" +""" + +WEAKNESS_ANALYSIS_TEMPLATE = """You are an expert prompt engineer. +Analyze the prompt using the strongest and weakest examples below. Return JSON +only with these fields: +- weak_segments: a list containing only role, context, tasks, output_format +- strong_segments: a list containing only role, context, tasks, output_format +- recommendations: an object mapping weak segment names to concise actions + +Prompt: +\"\"\" +{prompt} +\"\"\" + +Current segments: +Role: {role} +Context: {context} +Tasks: {tasks} +Output format: {output_format} + +Best examples: +{best_examples} + +Worst examples: +{worst_examples} +""" + +CANDIDATE_GENERATION_TEMPLATE = """You are an expert prompt engineer. +Generate {n_candidates} diverse, standalone improved versions of the current +prompt. Modify the weak segments according to the recommendations and preserve +the strong segments. Do not copy dataset examples into a prompt. Do not add an +input placeholder: CoolPrompt appends each input separately at runtime. + +Return JSON only as {{"prompts": ["candidate 1", "candidate 2"]}}. + +Current prompt: +\"\"\" +{current_prompt} +\"\"\" + +Segments: +Role: {role} +Context: {context} +Tasks: {tasks} +Output format: {output_format} + +Weak segments: {weak_segments} +Strong segments: {strong_segments} +Recommendations: +{recommendations} +""" + diff --git a/coolprompt/optimizer/sapo/sapo.py b/coolprompt/optimizer/sapo/sapo.py new file mode 100644 index 0000000..8a94dc9 --- /dev/null +++ b/coolprompt/optimizer/sapo/sapo.py @@ -0,0 +1,386 @@ +"""Segment-based Automatic Prompt Optimization (SAPO) for CoolPrompt.""" + +from __future__ import annotations + +import time +from typing import Any, override + +from langchain_core.language_models.base import BaseLanguageModel +from pydantic import BaseModel, Field + +from coolprompt.evaluator import Evaluator +from coolprompt.optimizer.autoprompting_method import ( + AutoPromptingMethod, + BenchmarkContext, + TelemetryCallback, +) +from coolprompt.utils.logging_config import logger +from coolprompt.utils.parsing import extract_json + +from .prompt_templates import ( + CANDIDATE_GENERATION_TEMPLATE, + SEGMENTATION_TEMPLATE, + WEAKNESS_ANALYSIS_TEMPLATE, +) + +_SEGMENTS = ("role", "context", "tasks", "output_format") + + +class PromptSegments(BaseModel): + """A prompt split into the four segments used by SAPO.""" + + role: str = "" + context: str = "" + tasks: str = "" + output_format: str = "" + + +class WeaknessAnalysis(BaseModel): + """Segment-level diagnosis produced from contrastive examples.""" + + weak_segments: list[str] = Field(default_factory=list) + strong_segments: list[str] = Field(default_factory=list) + recommendations: dict[str, str] = Field(default_factory=dict) + + +class CandidatePrompts(BaseModel): + """Structured response for candidate prompt generation.""" + + prompts: list[str] = Field(default_factory=list) + + +class SAPOOptimizer: + """Optimize a prompt by diagnosing and rewriting its logical segments.""" + + def __init__( + self, + model: BaseLanguageModel, + evaluator: Evaluator, + *, + n_iterations: int = 5, + n_candidates: int = 5, + early_stopping_rounds: int = 3, + examples_per_side: int = 5, + max_request_retries: int = 3, + retry_delay_seconds: float = 1.0, + telemetry_callback: TelemetryCallback | None = None, + ) -> None: + if n_iterations < 0: + raise ValueError("n_iterations must be non-negative") + if n_candidates < 1: + raise ValueError("n_candidates must be at least 1") + if early_stopping_rounds < 1: + raise ValueError("early_stopping_rounds must be at least 1") + self.model = model + self.evaluator = evaluator + self.n_iterations = n_iterations + self.n_candidates = n_candidates + self.early_stopping_rounds = early_stopping_rounds + self.examples_per_side = max(1, examples_per_side) + self.max_request_retries = max(1, max_request_retries) + self.retry_delay_seconds = max(0.0, retry_delay_seconds) + self.telemetry_callback = telemetry_callback + self.history: list[dict[str, Any]] = [] + + @staticmethod + def _response_text(response: Any) -> str: + content = getattr(response, "content", response) + if isinstance(content, list): + return "".join( + str(item.get("text", "")) if isinstance(item, dict) else str(item) + for item in content + ) + return str(content) + + def _retry(self, operation: str, function): + last_error: Exception | None = None + for attempt in range(1, self.max_request_retries + 1): + try: + return function() + except Exception as exc: # noqa: BLE001 - provider exceptions vary + last_error = exc + if attempt < self.max_request_retries: + logger.warning( + "SAPO %s failed (%s/%s): %s", + operation, + attempt, + self.max_request_retries, + exc, + ) + time.sleep(self.retry_delay_seconds) + raise RuntimeError( + f"SAPO {operation} failed after {self.max_request_retries} attempts" + ) from last_error + + def _invoke_structured(self, prompt: str, schema: type[BaseModel]) -> BaseModel: + """Use native structured output when available, with JSON fallback.""" + + def invoke() -> BaseModel: + if hasattr(self.model, "with_structured_output"): + try: + result = self.model.with_structured_output(schema).invoke(prompt) + if isinstance(result, schema): + return result + return schema.model_validate(result) + except (AttributeError, NotImplementedError, TypeError, ValueError): + pass + parsed = extract_json(self._response_text(self.model.invoke(prompt))) + if parsed is None: + raise ValueError(f"Model did not return JSON for {schema.__name__}") + return schema.model_validate(parsed) + + return self._retry(schema.__name__, invoke) + + def _evaluate( + self, prompt: str, dataset: list[str], targets: list[Any] + ) -> tuple[float, list[float], list[str]]: + result = self.evaluator.evaluate( + prompt=prompt, + dataset=dataset, + targets=targets, + return_detailed=True, + ) + return ( + float(result.aggregate_score), + [float(score) for score in result.score_per_task], + list(result.raw_outputs), + ) + + def _extract_segments(self, prompt: str) -> PromptSegments: + return self._invoke_structured( + SEGMENTATION_TEMPLATE.format(prompt=prompt), PromptSegments + ) + + def _ranked_examples( + self, + dataset: list[str], + targets: list[Any], + scores: list[float], + outputs: list[str], + ) -> tuple[str, str]: + indices = sorted(range(len(scores)), key=scores.__getitem__, reverse=True) + count = min(self.examples_per_side, len(indices)) + + def render(selected: list[int]) -> str: + return "\n\n".join( + f"Input: {dataset[index]}\nReference: {targets[index]}\n" + f"Model response: {outputs[index]}\nScore: {scores[index]:.4f}" + for index in selected + ) + + return render(indices[:count]), render(indices[-count:]) + + @staticmethod + def _sanitize_analysis(analysis: WeaknessAnalysis) -> WeaknessAnalysis: + weak = list(dict.fromkeys(x for x in analysis.weak_segments if x in _SEGMENTS)) + strong = list( + dict.fromkeys( + x for x in analysis.strong_segments if x in _SEGMENTS and x not in weak + ) + ) + recommendations = { + key: str(value).strip() + for key, value in analysis.recommendations.items() + if key in weak and str(value).strip() + } + if not weak: + weak = ["tasks"] + for segment in weak: + recommendations.setdefault(segment, f"Make the {segment} clearer and more specific.") + return WeaknessAnalysis( + weak_segments=weak, + strong_segments=strong, + recommendations=recommendations, + ) + + def _analyze( + self, + prompt: str, + segments: PromptSegments, + dataset: list[str], + targets: list[Any], + scores: list[float], + outputs: list[str], + ) -> WeaknessAnalysis: + best, worst = self._ranked_examples(dataset, targets, scores, outputs) + analysis = self._invoke_structured( + WEAKNESS_ANALYSIS_TEMPLATE.format( + prompt=prompt, + best_examples=best, + worst_examples=worst, + **segments.model_dump(), + ), + WeaknessAnalysis, + ) + return self._sanitize_analysis(analysis) + + def _generate_candidates( + self, + prompt: str, + segments: PromptSegments, + analysis: WeaknessAnalysis, + ) -> list[str]: + recommendations = "\n".join( + f"- {key}: {value}" for key, value in analysis.recommendations.items() + ) + result = self._invoke_structured( + CANDIDATE_GENERATION_TEMPLATE.format( + n_candidates=self.n_candidates, + current_prompt=prompt, + weak_segments=", ".join(analysis.weak_segments), + strong_segments=", ".join(analysis.strong_segments) or "none", + recommendations=recommendations, + **segments.model_dump(), + ), + CandidatePrompts, + ) + candidates: list[str] = [] + seen = {prompt.strip().casefold()} + for candidate in result.prompts: + candidate = candidate.strip() + key = candidate.casefold() + if candidate and key not in seen: + candidates.append(candidate) + seen.add(key) + if not candidates: + logger.warning("SAPO produced no distinct candidates; stopping") + return candidates[: self.n_candidates] + + def optimize( + self, + initial_prompt: str, + dataset_split: tuple[list[str], list[str], list[Any], list[Any]], + ) -> str: + train_data, val_data, train_targets, val_targets = map(list, dataset_split) + if not train_data or not train_targets: + raise ValueError("SAPO requires a non-empty training dataset and targets") + if not val_data: + val_data, val_targets = train_data, train_targets + + current_prompt = initial_prompt.strip() + train_score, train_scores, train_outputs = self._evaluate( + current_prompt, train_data, train_targets + ) + best_score, _, _ = self._evaluate(current_prompt, val_data, val_targets) + best_prompt = current_prompt + self.history = [ + {"iteration": 0, "prompt": current_prompt, "train_score": train_score, + "val_score": best_score, "improved": False} + ] + if self.telemetry_callback: + self.telemetry_callback(0, best_score, best_prompt) + + no_improvement = 0 + for iteration in range(1, self.n_iterations + 1): + segments = self._extract_segments(current_prompt) + analysis = self._analyze( + current_prompt, + segments, + train_data, + train_targets, + train_scores, + train_outputs, + ) + candidates = self._generate_candidates(current_prompt, segments, analysis) + if not candidates: + break + + evaluated = [] + for candidate in candidates: + candidate_train = self._evaluate(candidate, train_data, train_targets) + candidate_val = self._evaluate(candidate, val_data, val_targets)[0] + evaluated.append((candidate_val, candidate, candidate_train)) + candidate_score, candidate_prompt, candidate_train = max( + evaluated, key=lambda item: item[0] + ) + improved = candidate_score > best_score + if improved: + best_score = candidate_score + best_prompt = candidate_prompt + current_prompt = candidate_prompt + train_score, train_scores, train_outputs = candidate_train + no_improvement = 0 + else: + no_improvement += 1 + + self.history.append( + { + "iteration": iteration, + "prompt": current_prompt, + "train_score": train_score, + "val_score": best_score, + "segments": segments.model_dump(), + "analysis": analysis.model_dump(), + "candidates": candidates, + "candidate_scores": [item[0] for item in evaluated], + "improved": improved, + } + ) + logger.info( + "SAPO iteration %s: best validation score %.4f%s", + iteration, + best_score, + " (improved)" if improved else "", + ) + if self.telemetry_callback: + self.telemetry_callback(iteration, best_score, best_prompt) + if no_improvement >= self.early_stopping_rounds: + break + return best_prompt + + +class SAPOMethod(AutoPromptingMethod): + """CoolPrompt adapter for Segment-based Automatic Prompt Optimization.""" + + def __init__(self) -> None: + self.last_optimizer: SAPOOptimizer | None = None + + @override + def optimize( + self, + model, + initial_prompt, + dataset_split, + evaluator, + problem_description, + **kwargs, + ) -> str: + del problem_description + telemetry_callback = kwargs.pop("telemetry_callback", None) + self.last_optimizer = SAPOOptimizer( + model=model, + evaluator=evaluator, + n_iterations=kwargs.pop("n_iterations", 5), + n_candidates=kwargs.pop("n_candidates", 5), + early_stopping_rounds=kwargs.pop( + "early_stopping_rounds", kwargs.pop("patience", 3) + ), + examples_per_side=kwargs.pop("examples_per_side", 5), + max_request_retries=kwargs.pop("max_request_retries", 3), + retry_delay_seconds=kwargs.pop("retry_delay_seconds", 1.0), + telemetry_callback=telemetry_callback, + ) + if kwargs: + logger.debug("Ignoring unsupported SAPO options: %s", sorted(kwargs)) + return self.last_optimizer.optimize(initial_prompt, dataset_split) + + @override + def run_configured_benchmark(self, ctx: BenchmarkContext, start_prompt: str) -> str: + config = dict(ctx.config.get("method", {})) + return self.optimize( + model=ctx.model, + initial_prompt=start_prompt, + dataset_split=ctx.dataset_split, + evaluator=ctx.evaluator, + problem_description=ctx.config.get("problem_description"), + **config, + ) + + @override + def is_data_driven(self) -> bool: + return True + + @property + @override + def name(self) -> str: + return "sapo" diff --git a/coolprompt/utils/var_validation.py b/coolprompt/utils/var_validation.py index e159a15..225fe33 100644 --- a/coolprompt/utils/var_validation.py +++ b/coolprompt/utils/var_validation.py @@ -12,6 +12,7 @@ from coolprompt.optimizer.reflective_prompt import ReflectiveMethod from coolprompt.optimizer.regps import ReGPSMethod from coolprompt.optimizer.rider import RIDERGenesisMethod +from coolprompt.optimizer.sapo import SAPOMethod from coolprompt.utils.enums import PD_Method, Task from coolprompt.utils.logging_config import logger @@ -24,6 +25,7 @@ "compress": CompressorMethod, "rider": RIDERGenesisMethod, "brave": BRAVEMethod, + "sapo": SAPOMethod, } diff --git a/docs/API.md b/docs/API.md index 897c6c7..6fc97f7 100644 --- a/docs/API.md +++ b/docs/API.md @@ -18,7 +18,7 @@ retrieve similar datasets, then ranks comparable configurations by quality, cost, and runtime. The result is stored in `PromptTuner.meta_selection` and included in telemetry exports. `meta_classifier_path` overrides the bundled CSV; `meta_dataset_name` is an optional retrieval hint. SAPO recommendations -fall back to `hyper`, because SAPO is not yet implemented in CoolPrompt. +select the built-in `sapo` optimizer. --- ## `evaluator/` @@ -63,6 +63,7 @@ Method names accepted by `PromptTuner.run(method=...)`: - `regps` - RE-GPS optimizer. - `rider` - RIDER optimizer. Documentation - `brave` - BRAVE budget-aware evolutionary optimizer. Documentation +- `sapo` - segment-based optimizer using contrastive best/worst examples. - `compress` - PromptCompressor. Documentation - `reflective` - legacy ReflectivePrompt. Documentation - `distill` - legacy DistillPrompt. Documentation @@ -73,7 +74,7 @@ Custom methods should implement `AutoPromptingMethod` and can be passed to `Prom ## `method_evaluation/` Benchmark interface for comparing autoprompting methods on dataset/config-based experiments. -`evaluate_method(...)` supports the built-in method names `hyper_light`, `hyper`, `reflective`, `reflectiveprompt`, `distill`, `compress`, `regps`, `rider`, and `brave`. +`evaluate_method(...)` supports the built-in method names `hyper_light`, `hyper`, `reflective`, `reflectiveprompt`, `distill`, `compress`, `regps`, `rider`, `brave`, and `sapo`. --- ## `data_generator/` and `task_detector/` diff --git a/test/coolprompt/optimizer/test_sapo.py b/test/coolprompt/optimizer/test_sapo.py new file mode 100644 index 0000000..808a042 --- /dev/null +++ b/test/coolprompt/optimizer/test_sapo.py @@ -0,0 +1,89 @@ +"""Unit and integration-style tests for the CoolPrompt SAPO adapter.""" + +from types import SimpleNamespace + +import pytest + +from coolprompt.meta_selector import APOMetaSelector +from coolprompt.optimizer.sapo import SAPOMethod, SAPOOptimizer + + +class _JSONModel: + def __init__(self): + self.prompts = [] + + def invoke(self, prompt): + self.prompts.append(prompt) + if "Decompose the prompt" in prompt: + return '{"role":"","context":"","tasks":"Summarize",' \ + '"output_format":""}' + if "Analyze the prompt" in prompt: + return '{"weak_segments":["tasks"],"strong_segments":[],' \ + '"recommendations":{"tasks":"Make the instruction clear"}}' + if "Generate 2 diverse" in prompt: + return '{"prompts":["Write a clear one-sentence summary.",' \ + '"Write a concise summary."]}' + raise AssertionError(f"Unexpected prompt: {prompt}") + + +class _Evaluator: + def __init__(self): + self.prompts = [] + + def evaluate(self, prompt, dataset, targets, return_detailed=False): + assert return_detailed is True + assert len(dataset) == len(targets) + self.prompts.append(prompt) + score = 0.9 if "clear" in prompt.lower() else 0.2 + return SimpleNamespace( + aggregate_score=score, + score_per_task=[score] * len(dataset), + raw_outputs=["summary"] * len(dataset), + ) + + +def test_sapo_optimizes_with_coolprompt_model_and_evaluator(): + telemetry = [] + optimizer = SAPOOptimizer( + model=_JSONModel(), + evaluator=_Evaluator(), + n_iterations=2, + n_candidates=2, + early_stopping_rounds=1, + retry_delay_seconds=0, + telemetry_callback=lambda *args: telemetry.append(args), + ) + + result = optimizer.optimize( + "Summarize the article.", + (["train"], ["validation"], ["train summary"], ["val summary"]), + ) + + assert result == "Write a clear one-sentence summary." + assert optimizer.history[1]["improved"] is True + assert optimizer.history[-1]["val_score"] == pytest.approx(0.9) + assert telemetry[0] == (0, 0.2, "Summarize the article.") + assert telemetry[-1][1:] == (0.9, result) + + +def test_sapo_method_adapter_and_empty_training_validation(): + method = SAPOMethod() + with pytest.raises(ValueError, match="non-empty training"): + method.optimize( + model=_JSONModel(), + initial_prompt="Summarize.", + dataset_split=([], [], [], []), + evaluator=_Evaluator(), + problem_description=None, + n_iterations=0, + ) + + +def test_meta_selector_maps_sapo_without_fallback(): + assert APOMetaSelector._map_method("SAPO") == ("sapo", None) + + +def test_sapo_rejects_invalid_search_configuration(): + with pytest.raises(ValueError, match="n_candidates"): + SAPOOptimizer(model=_JSONModel(), evaluator=_Evaluator(), n_candidates=0) + diff --git a/test/coolprompt/optimizer/test_smoke_optimizer_interface.py b/test/coolprompt/optimizer/test_smoke_optimizer_interface.py index c5b91eb..82ad9b1 100644 --- a/test/coolprompt/optimizer/test_smoke_optimizer_interface.py +++ b/test/coolprompt/optimizer/test_smoke_optimizer_interface.py @@ -6,6 +6,7 @@ from coolprompt.optimizer.brave import BRAVEMethod from coolprompt.optimizer.hyper.meta_prompt import HyPERLightMethod from coolprompt.optimizer.rider import RIDERGenesisMethod +from coolprompt.optimizer.sapo import SAPOMethod from coolprompt.utils.var_validation import _METHOD_BY_NAME, validate_method @@ -18,6 +19,9 @@ def test_autoprompting_module_exports(): assert issubclass(BRAVEMethod, AutoPromptingMethod) assert BRAVEMethod().name == "brave" assert BRAVEMethod().is_data_driven() is True + assert issubclass(SAPOMethod, AutoPromptingMethod) + assert SAPOMethod().name == "sapo" + assert SAPOMethod().is_data_driven() is True def test_validate_method_string_class_and_instance_equivalent(): @@ -60,6 +64,7 @@ def test_method_by_name_covers_expected_keys(): "compress", "rider", "brave", + "sapo", } @@ -70,6 +75,7 @@ def test_method_evaluation_entrypoint(): assert hasattr(HyPERLightMethod(), "run") assert "rider" in me._BENCHMARK_IMPL assert "brave" in me._BENCHMARK_IMPL + assert "sapo" in me._BENCHMARK_IMPL def test_prompt_tuner_importable(): diff --git a/test/coolprompt/test_prompt_tuner_auto.py b/test/coolprompt/test_prompt_tuner_auto.py index 9fdd325..2aeb3fb 100644 --- a/test/coolprompt/test_prompt_tuner_auto.py +++ b/test/coolprompt/test_prompt_tuner_auto.py @@ -78,6 +78,7 @@ def no_dataset_result(cls): **validation._METHOD_BY_NAME, "rider": lambda: _Method("rider"), "hyper": lambda: _Method("hyper"), + "sapo": lambda: _Method("sapo"), "hyper_light": lambda: _Method("hyper_light", data_driven=False), }, ) @@ -114,14 +115,14 @@ def test_auto_uses_selected_rider_and_exposes_selection(monkeypatch): assert tuner.meta_selection is selection -def test_auto_applies_sapo_fallback_to_hyper(monkeypatch): +def test_auto_runs_selected_sapo(monkeypatch): selection = MetaSelectionResult( profile={"domain": "news"}, similar_datasets=["xsum"], candidates=[], recommended_method="SAPO", - selected_method="hyper", - fallback_reason="SAPO is not implemented in CoolPrompt", + selected_method="sapo", + fallback_reason=None, recommended_model="openai/gpt-4o-mini", recommended_split="150/100/300", recommended_metric="BERTScore F1", @@ -140,10 +141,8 @@ def test_auto_applies_sapo_fallback_to_hyper(monkeypatch): enable_telemetry=False, ) - assert result == "optimized with hyper" - assert ( - tuner.meta_selection.fallback_reason == "SAPO is not implemented in CoolPrompt" - ) + assert result == "optimized with sapo" + assert tuner.meta_selection.fallback_reason is None def test_auto_without_data_uses_hyper_light(monkeypatch): From 5a4a062d4a89c2af8b33dda409d11a88b511e654 Mon Sep 17 00:00:00 2001 From: Nikita Kulin Date: Fri, 2 Oct 2026 15:04:08 +0300 Subject: [PATCH 2/4] black moment --- coolprompt/meta_selector/data/__init__.py | 1 - coolprompt/optimizer/sapo/__init__.py | 1 - coolprompt/optimizer/sapo/prompt_templates.py | 1 - coolprompt/optimizer/sapo/sapo.py | 13 ++++++++++--- 4 files changed, 10 insertions(+), 6 deletions(-) delete mode 100644 coolprompt/meta_selector/data/__init__.py diff --git a/coolprompt/meta_selector/data/__init__.py b/coolprompt/meta_selector/data/__init__.py deleted file mode 100644 index cf5d113..0000000 --- a/coolprompt/meta_selector/data/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Bundled metadata used by :mod:`coolprompt.meta_selector`.""" diff --git a/coolprompt/optimizer/sapo/__init__.py b/coolprompt/optimizer/sapo/__init__.py index 972f7b7..85f00da 100644 --- a/coolprompt/optimizer/sapo/__init__.py +++ b/coolprompt/optimizer/sapo/__init__.py @@ -3,4 +3,3 @@ from .sapo import SAPOMethod, SAPOOptimizer __all__ = ["SAPOMethod", "SAPOOptimizer"] - diff --git a/coolprompt/optimizer/sapo/prompt_templates.py b/coolprompt/optimizer/sapo/prompt_templates.py index 6e51a38..d7cb9f7 100644 --- a/coolprompt/optimizer/sapo/prompt_templates.py +++ b/coolprompt/optimizer/sapo/prompt_templates.py @@ -60,4 +60,3 @@ Recommendations: {recommendations} """ - diff --git a/coolprompt/optimizer/sapo/sapo.py b/coolprompt/optimizer/sapo/sapo.py index 8a94dc9..2e4688e 100644 --- a/coolprompt/optimizer/sapo/sapo.py +++ b/coolprompt/optimizer/sapo/sapo.py @@ -186,7 +186,9 @@ def _sanitize_analysis(analysis: WeaknessAnalysis) -> WeaknessAnalysis: if not weak: weak = ["tasks"] for segment in weak: - recommendations.setdefault(segment, f"Make the {segment} clearer and more specific.") + recommendations.setdefault( + segment, f"Make the {segment} clearer and more specific." + ) return WeaknessAnalysis( weak_segments=weak, strong_segments=strong, @@ -264,8 +266,13 @@ def optimize( best_score, _, _ = self._evaluate(current_prompt, val_data, val_targets) best_prompt = current_prompt self.history = [ - {"iteration": 0, "prompt": current_prompt, "train_score": train_score, - "val_score": best_score, "improved": False} + { + "iteration": 0, + "prompt": current_prompt, + "train_score": train_score, + "val_score": best_score, + "improved": False, + } ] if self.telemetry_callback: self.telemetry_callback(0, best_score, best_prompt) From 7faf958cdd05d241ea94c2d4c762d9721e8b266f Mon Sep 17 00:00:00 2001 From: Nikita Kulin Date: Fri, 2 Oct 2026 15:07:09 +0300 Subject: [PATCH 3/4] black --- test/coolprompt/optimizer/test_sapo.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/test/coolprompt/optimizer/test_sapo.py b/test/coolprompt/optimizer/test_sapo.py index 808a042..ca131cc 100644 --- a/test/coolprompt/optimizer/test_sapo.py +++ b/test/coolprompt/optimizer/test_sapo.py @@ -15,14 +15,17 @@ def __init__(self): def invoke(self, prompt): self.prompts.append(prompt) if "Decompose the prompt" in prompt: - return '{"role":"","context":"","tasks":"Summarize",' \ - '"output_format":""}' + return '{"role":"","context":"","tasks":"Summarize",' '"output_format":""}' if "Analyze the prompt" in prompt: - return '{"weak_segments":["tasks"],"strong_segments":[],' \ + return ( + '{"weak_segments":["tasks"],"strong_segments":[],' '"recommendations":{"tasks":"Make the instruction clear"}}' + ) if "Generate 2 diverse" in prompt: - return '{"prompts":["Write a clear one-sentence summary.",' \ + return ( + '{"prompts":["Write a clear one-sentence summary.",' '"Write a concise summary."]}' + ) raise AssertionError(f"Unexpected prompt: {prompt}") @@ -86,4 +89,3 @@ def test_meta_selector_maps_sapo_without_fallback(): def test_sapo_rejects_invalid_search_configuration(): with pytest.raises(ValueError, match="n_candidates"): SAPOOptimizer(model=_JSONModel(), evaluator=_Evaluator(), n_candidates=0) - From 662fc8b0f72e371c9fdf7c2b16c60f96e5d0ecdd Mon Sep 17 00:00:00 2001 From: Nikita Kulin Date: Fri, 2 Oct 2026 15:27:19 +0300 Subject: [PATCH 4/4] fix --- test/coolprompt/meta_selector/test_selector.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/test/coolprompt/meta_selector/test_selector.py b/test/coolprompt/meta_selector/test_selector.py index df7bde2..fc38758 100644 --- a/test/coolprompt/meta_selector/test_selector.py +++ b/test/coolprompt/meta_selector/test_selector.py @@ -68,8 +68,8 @@ def test_selector_ranks_quality_cost_time_and_maps_sapo(tmp_path): result = selector.select(model, "Summarize", "generation") assert result.recommended_method == "SAPO" - assert result.selected_method == "hyper" - assert result.fallback_reason == "SAPO is not implemented in CoolPrompt" + assert result.selected_method == "sapo" + assert result.fallback_reason is None assert result.candidates[0].final_score > result.candidates[-1].final_score