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/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/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..85f00da
--- /dev/null
+++ b/coolprompt/optimizer/sapo/__init__.py
@@ -0,0 +1,5 @@
+"""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..d7cb9f7
--- /dev/null
+++ b/coolprompt/optimizer/sapo/prompt_templates.py
@@ -0,0 +1,62 @@
+"""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..2e4688e
--- /dev/null
+++ b/coolprompt/optimizer/sapo/sapo.py
@@ -0,0 +1,393 @@
+"""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/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
diff --git a/test/coolprompt/optimizer/test_sapo.py b/test/coolprompt/optimizer/test_sapo.py
new file mode 100644
index 0000000..ca131cc
--- /dev/null
+++ b/test/coolprompt/optimizer/test_sapo.py
@@ -0,0 +1,91 @@
+"""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):