From 6fae7021a6973f2dbe6beefe15cb67f8ea71b648 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 13 Apr 2026 10:32:03 -0700 Subject: [PATCH] improve compression quality: line-based truncation, multi-message budget, 70% default target - Switch truncate_message from word-based to line-based splitting to preserve code structure (function boundaries, indentation) - Allow multiple messages to be truncated instead of burning entire budget on one overflow message - Raise default compression target from 50% to 70% of trigger for better quality/cost tradeoff - Add --compression-target CLI arg to SWE-bench eval harness - Move tests to canonical locations (tests/test_litellm/, scripts/) - Add docs page and sidebar entries for compress() Eval results (5 problems, Opus, trigger=10k): Hunk overlap delta improved from -0.417 to -0.221 Content similarity now matches baseline (+0.006) Cost savings: 72% Co-Authored-By: Claude Opus 4.6 --- .../docs/completion/prompt_compression.md | 79 +++++ docs/my-website/package-lock.json | 7 + docs/my-website/sidebars.js | 6 + litellm/compression/compress.py | 37 ++- litellm/compression/message_stubbing.py | 30 +- .../compression/scoring/embedding_scorer.py | 9 +- {tests => scripts}/eval_compression.py | 6 +- tests/eval_swe_bench.py | 300 ++++++++++++++++-- tests/{ => test_litellm}/test_compression.py | 66 ++++ 9 files changed, 482 insertions(+), 58 deletions(-) create mode 100644 docs/my-website/docs/completion/prompt_compression.md rename {tests => scripts}/eval_compression.py (99%) rename tests/{ => test_litellm}/test_compression.py (82%) diff --git a/docs/my-website/docs/completion/prompt_compression.md b/docs/my-website/docs/completion/prompt_compression.md new file mode 100644 index 00000000000..0ea679fb6b7 --- /dev/null +++ b/docs/my-website/docs/completion/prompt_compression.md @@ -0,0 +1,79 @@ +# Prompt Compression (`compress()`) + +Use `litellm.compress()` to shrink long conversation history before calling `completion()`. + +The function keeps high-relevance and recent context, replaces low-relevance content with lightweight stubs, and returns a retrieval tool so the model can request full content only when needed. + +## Quickstart + +```python +import litellm + +messages = [ + {"role": "system", "content": "You are a coding assistant."}, + {"role": "user", "content": "# auth.py\n" + "def authenticate():\n pass\n" * 2000}, + {"role": "user", "content": "# utils.py\n" + "def helper():\n pass\n" * 2000}, + {"role": "user", "content": "Fix the bug in auth.py"}, +] + +compressed = litellm.compress( + messages=messages, + model="gpt-4o", + compression_trigger=1000, + compression_target=500, +) + +response = litellm.completion( + model="gpt-4o", + messages=compressed["messages"], + tools=compressed["tools"], +) +``` + +## What It Returns + +`compress()` returns a dictionary with: + +- `messages`: compressed conversation messages +- `original_tokens`: token count before compression +- `compressed_tokens`: token count after compression +- `compression_ratio`: fraction of tokens removed +- `cache`: key-value mapping of stub key -> original full content +- `tools`: retrieval tool definition (`litellm_content_retrieve`) for on-demand restoration + +## Parameters + +- `messages` (`List[dict]`, required): input conversation messages +- `model` (`str`, required): model name used for token counting +- `compression_trigger` (`int`, default `200000`): compress only if input token count exceeds this +- `compression_target` (`Optional[int]`, default `compression_trigger // 2`): desired post-compression token budget +- `embedding_model` (`Optional[str]`): if set, combines BM25 + embedding relevance scoring +- `embedding_model_params` (`Optional[dict]`): additional kwargs passed to `litellm.embedding()` +- `compression_cache` (`Optional[DualCache]`): optional cache used by embedding scoring + +## Behavior Notes + +- Messages below `compression_trigger` are passed through unchanged. +- System messages, the last user message, and the last assistant message are always preserved. +- If a relevant message does not fully fit the remaining budget, `compress()` may keep a truncated version of it. +- Compressed-out content is never lost; it is stored in `cache` and addressable by `litellm_content_retrieve`. + +## Handling Retrieval Tool Calls + +If the model calls `litellm_content_retrieve`, look up the requested key in `compressed["cache"]` and return that value as tool output. + +```python +import json + +tool_call = response.choices[0].message.tool_calls[0] +args = json.loads(tool_call.function.arguments) +full_content = compressed["cache"][args["key"]] +``` + +## Evaluate Compression Quality + +You can benchmark baseline vs compressed behavior with: + +```bash +python scripts/eval_compression.py --model gpt-4o --problems 5 +``` diff --git a/docs/my-website/package-lock.json b/docs/my-website/package-lock.json index 56684b737de..d14ca96cf5b 100644 --- a/docs/my-website/package-lock.json +++ b/docs/my-website/package-lock.json @@ -20403,6 +20403,13 @@ "url": "https://opencollective.com/webpack" } }, + "node_modules/search-insights": { + "version": "2.17.3", + "resolved": "https://registry.npmjs.org/search-insights/-/search-insights-2.17.3.tgz", + "integrity": "sha512-RQPdCYTa8A68uM2jwxoY842xDhvx3E5LFL1LxvxCNMev4o5mLuokczhzjAgGwUZBAmOKZknArSxLKmXtIi2AxQ==", + "license": "MIT", + "peer": true + }, "node_modules/section-matter": { "version": "1.0.0", "resolved": "https://registry.npmjs.org/section-matter/-/section-matter-1.0.0.tgz", diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index b2ac8433911..46e392037a6 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -254,6 +254,11 @@ const sidebars = { id: "image_generation", label: "image_generation()", }, + { + type: "doc", + id: "completion/prompt_compression", + label: "compress()", + }, { type: "doc", id: "audio_transcription", @@ -1280,6 +1285,7 @@ const learnSidebar = { items: [ "completion/prefix", "completion/predict_outputs", + "completion/prompt_compression", "completion/message_trimming", "completion/prompt_caching", "completion/prompt_formatting", diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index 8a79f7274ca..718bc1c45c3 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -3,10 +3,14 @@ Main compress() function — orchestrates BM25/embedding scoring, message stubbi and retrieval tool injection. """ -from typing import Dict, List, Optional, Set +from typing import Any, Dict, List, Optional, Set from litellm.caching.dual_cache import DualCache -from litellm.compression.message_stubbing import extract_key, stub_message, truncate_message +from litellm.compression.message_stubbing import ( + extract_key, + stub_message, + truncate_message, +) from litellm.compression.retrieval_tool import build_retrieval_tool from litellm.compression.scoring.bm25 import bm25_score_messages from litellm.litellm_core_utils.token_counter import token_counter @@ -88,6 +92,7 @@ def compress( compression_trigger: int = 200_000, compression_target: Optional[int] = None, embedding_model: Optional[str] = None, + embedding_model_params: Optional[Dict[str, Any]] = None, compression_cache: Optional[DualCache] = None, ) -> CompressedResult: """ @@ -107,6 +112,8 @@ def compress( Defaults to ``compression_trigger // 2``. embedding_model: If provided, use BM25 + embeddings for scoring. If ``None``, BM25 only. + embedding_model_params: Optional kwargs forwarded to + ``litellm.embedding()`` when ``embedding_model`` is set. compression_cache: Passed through to ``litellm.embedding()`` for cross-turn caching of embedding vectors. @@ -115,7 +122,7 @@ def compress( counts, a cache of original content, and the retrieval tool definition. """ if compression_target is None: - compression_target = compression_trigger // 2 + compression_target = compression_trigger * 7 // 10 original_tokens = token_counter(model=model, messages=messages) @@ -142,7 +149,11 @@ def compress( ) emb_scores = embedding_score_messages( - query, messages, model=embedding_model, cache=compression_cache + query, + messages, + model=embedding_model, + cache=compression_cache, + embedding_model_params=embedding_model_params, ) combined_scores = _combine_scores(bm25_scores, emb_scores, bm25_weight=0.4) else: @@ -170,10 +181,10 @@ def compress( # For each candidate (ranked by relevance): # - If it fits entirely → keep it as-is. # - If it doesn't fit but there's meaningful remaining budget → truncate it - # to fill that budget (so the LLM always has real content to work with). + # to fill as much of the budget as possible. # - Otherwise → stub it (pointer only, content goes to cache). - # We only truncate one message (the highest-scoring one that overflows) so - # the budget is consumed and the rest are stubbed cleanly. + # Multiple messages may be truncated so we preserve partial content from + # several high-scoring messages rather than fully stubbing all but one. truncated_overrides: Dict[int, dict] = {} # idx -> truncated message dict for idx in ranked_indices: @@ -183,17 +194,23 @@ def compress( msg_tokens = token_counter(model=model, text=msg_content) remaining = compression_target - current_tokens + if remaining <= 0: + break # budget exhausted + if current_tokens + msg_tokens <= compression_target: # Fits entirely kept_indices.add(idx) current_tokens += msg_tokens - elif remaining >= 100 and idx not in truncated_overrides: + elif remaining >= 100: # Too large to fit whole, but we have budget — truncate it. - # Only do this once (the highest-scoring overflow message). truncated = truncate_message(messages[idx], remaining) + truncated_tokens = token_counter( + model=model, + text=truncated.get("content", "") or "", + ) truncated_overrides[idx] = truncated kept_indices.add(idx) - current_tokens = compression_target # budget consumed + current_tokens += truncated_tokens # Build compressed messages and cache compressed_messages: List[dict] = [] diff --git a/litellm/compression/message_stubbing.py b/litellm/compression/message_stubbing.py index 6ca0a8399e0..2330f1bbc9e 100644 --- a/litellm/compression/message_stubbing.py +++ b/litellm/compression/message_stubbing.py @@ -80,7 +80,11 @@ def stub_message(message: dict, key: str) -> dict: def truncate_message(message: dict, max_tokens: int) -> dict: """ Truncate a message's content to approximately max_tokens by keeping - the first 70% and last 30% of words with a separator in between. + the first 70% and last 30% of lines with a separator in between. + + Uses line-based splitting to preserve code structure (function + boundaries, indentation) rather than word-based splitting which + mangles code. Used when a message is too large to fit entirely in the budget but too relevant to fully stub out. @@ -91,18 +95,26 @@ def truncate_message(message: dict, max_tokens: int) -> dict: p.get("text", "") if isinstance(p, dict) else str(p) for p in content ) - # Rough conversion: 1 token ≈ 0.75 words - target_words = max(1, int(max_tokens * 0.75)) - words = content.split() + # Rough conversion: 1 token ≈ 3 characters + target_chars = max(100, max_tokens * 3) - if len(words) <= target_words: + if len(content) <= target_chars: return {**message, "content": content} - first_count = (target_words * 2) // 3 - last_count = target_words - first_count + lines = content.split("\n") + + # Estimate target line count from character budget + avg_line_len = max(1, len(content) // max(1, len(lines))) + target_lines = max(2, target_chars // avg_line_len) + + if len(lines) <= target_lines: + return {**message, "content": content} + + first_count = (target_lines * 7) // 10 + last_count = target_lines - first_count truncated = ( - " ".join(words[:first_count]) + "\n".join(lines[:first_count]) + "\n...[truncated for context window]...\n" - + " ".join(words[-last_count:]) + + "\n".join(lines[-last_count:]) ) return {**message, "content": truncated} diff --git a/litellm/compression/scoring/embedding_scorer.py b/litellm/compression/scoring/embedding_scorer.py index 2f5f3e66f67..f3558ae8f5c 100644 --- a/litellm/compression/scoring/embedding_scorer.py +++ b/litellm/compression/scoring/embedding_scorer.py @@ -5,7 +5,7 @@ Computes cosine similarity between the query embedding and each message embeddin """ import math -from typing import List, Optional +from typing import Any, Dict, List, Optional from litellm.caching.dual_cache import DualCache @@ -49,6 +49,7 @@ def embedding_score_messages( messages: List[dict], model: str, cache: Optional[DualCache] = None, + embedding_model_params: Optional[Dict[str, Any]] = None, ) -> List[float]: """ Score each message's semantic similarity to the query using embeddings. @@ -58,6 +59,8 @@ def embedding_score_messages( messages: List of message dicts with "content" fields. model: The embedding model to use (e.g., "text-embedding-3-small"). cache: Optional DualCache for cross-turn embedding caching. + embedding_model_params: Optional additional kwargs forwarded to + ``litellm.embedding()``. Returns: List of float scores (cosine similarity), one per message. @@ -71,11 +74,13 @@ def embedding_score_messages( # Filter out empty texts — replace with a placeholder to maintain indexing processed_texts = [t if t.strip() else "empty" for t in texts] - kwargs = { + kwargs: Dict[str, Any] = { "model": model, "input": processed_texts, "caching": cache is not None, } + if embedding_model_params: + kwargs = {**kwargs, **embedding_model_params} response = litellm.embedding(**kwargs) diff --git a/tests/eval_compression.py b/scripts/eval_compression.py similarity index 99% rename from tests/eval_compression.py rename to scripts/eval_compression.py index ac9853b7fc4..d7d90dacc2e 100644 --- a/tests/eval_compression.py +++ b/scripts/eval_compression.py @@ -4,9 +4,9 @@ Prompt Compression Evaluation Harness Compare model performance on coding tasks with and without prompt compression. Usage: - python tests/eval_compression.py --model gpt-4o --problems 5 - python tests/eval_compression.py --model claude-sonnet-4-20250514 --problems 12 --runs 3 - python tests/eval_compression.py --model gpt-4o-mini --padding-factor 50 + python scripts/eval_compression.py --model gpt-4o --problems 5 + python scripts/eval_compression.py --model claude-sonnet-4-20250514 --problems 12 --runs 3 + python scripts/eval_compression.py --model gpt-4o-mini --padding-factor 50 The harness runs each problem in two modes: 1. **baseline** — raw prompt sent directly to the model. diff --git a/tests/eval_swe_bench.py b/tests/eval_swe_bench.py index 275748b8dfa..5cc9958584e 100644 --- a/tests/eval_swe_bench.py +++ b/tests/eval_swe_bench.py @@ -49,7 +49,10 @@ SYSTEM_MSG = ( "You are an expert software engineer resolving GitHub issues. " "You will be given an issue description and relevant source files. " "Produce a minimal unified diff patch that fixes the issue. " - "Output ONLY the patch starting with `diff --git`, no explanation." + "Your response must contain ONLY the patch in unified diff format. " + "Start with `diff --git a/path b/path`, then `---`, `+++`, and " + "`@@` hunks. Do NOT include any explanation, commentary, or markdown " + "fences — just the raw diff text." ) @@ -72,19 +75,34 @@ def _load_via_datasets(n: int, split: str) -> list[dict]: def _load_via_api(n: int, split: str) -> list[dict]: - """Fallback: fetch rows directly from the HuggingFace dataset API (no deps).""" + """Fallback: fetch rows directly from the HuggingFace dataset API (no deps). + + The API returns at most 100 rows per request, so we paginate. + """ import json import urllib.request - url = ( - "https://datasets-server.huggingface.co/rows" - "?dataset=princeton-nlp/SWE-bench_Lite_bm25_27K" - f"&config=default&split={split}&offset=0&length={n}" - ) - req = urllib.request.Request(url, headers={"User-Agent": "litellm-eval"}) - with urllib.request.urlopen(req, timeout=60) as resp: - data = json.loads(resp.read().decode()) - return [row["row"] for row in data["rows"]] + # 0 means "all" — SWE-bench Lite has 300 test instances + target = n if n > 0 else 300 + page_size = 100 + all_rows: list[dict] = [] + + for offset in range(0, target, page_size): + length = min(page_size, target - offset) + url = ( + "https://datasets-server.huggingface.co/rows" + "?dataset=princeton-nlp/SWE-bench_Lite_bm25_27K" + f"&config=default&split={split}&offset={offset}&length={length}" + ) + req = urllib.request.Request(url, headers={"User-Agent": "litellm-eval"}) + with urllib.request.urlopen(req, timeout=60) as resp: + data = json.loads(resp.read().decode()) + rows = [row["row"] for row in data["rows"]] + all_rows.extend(rows) + if len(rows) < length: + break # no more data + + return all_rows def load_problems(n: int = 10, split: str = "test") -> list[dict]: @@ -154,8 +172,16 @@ def build_messages(instance: dict) -> list[dict]: def parse_patch_files(patch: str) -> set[str]: - """Extract modified file paths from a unified diff.""" - return set(re.findall(r"^diff --git a/(.*) b/", patch, re.MULTILINE)) + """Extract modified file paths from a unified diff. + + Tries `diff --git a/path b/path` first, then falls back to + `--- a/path` lines for diffs that omit the git header. + """ + files = set(re.findall(r"^diff --git a/(.*?) b/", patch, re.MULTILINE)) + if not files: + # Fallback: extract from --- a/path lines + files = set(re.findall(r"^--- a/(.+)", patch, re.MULTILINE)) + return files def extract_patch(text: str) -> str: @@ -182,19 +208,81 @@ def is_valid_diff(patch: str) -> bool: # --------------------------------------------------------------------------- +def _parse_hunk_line_ranges(patch: str) -> dict[str, list[tuple[int, int]]]: + """Parse a unified diff into {filepath: [(start, end), ...]} for modified line ranges.""" + current_file = None + ranges: dict[str, list[tuple[int, int]]] = {} + for line in patch.split("\n"): + m = re.match(r"^diff --git a/(.*?) b/", line) + if m: + current_file = m.group(1) + if current_file not in ranges: + ranges[current_file] = [] + continue + if not current_file: + m2 = re.match(r"^--- a/(.+)", line) + if m2: + current_file = m2.group(1) + if current_file not in ranges: + ranges[current_file] = [] + continue + m3 = re.match(r"^@@ -(\d+)(?:,(\d+))? \+(\d+)(?:,(\d+))? @@", line) + if m3 and current_file: + start = int(m3.group(1)) + length = int(m3.group(2) or "1") + ranges[current_file].append((start, start + length)) + return ranges + + +def _extract_changed_lines(patch: str) -> set[str]: + """Extract the actual added/removed lines (stripped) from a diff.""" + lines = set() + for line in patch.split("\n"): + if line.startswith(("+", "-")) and not line.startswith(("+++", "---")): + stripped = line[1:].strip() + if stripped: + lines.add(stripped) + return lines + + +def _line_range_overlap( + ranges_a: dict[str, list[tuple[int, int]]], + ranges_b: dict[str, list[tuple[int, int]]], +) -> float: + """Compute fraction of gold hunk line ranges that overlap with generated ranges.""" + shared_files = set(ranges_a.keys()) & set(ranges_b.keys()) + if not shared_files: + return 0.0 + + total_gold_lines = 0 + overlapping_lines = 0 + + for f in shared_files: + for g_start, g_end in ranges_a[f]: + gold_set = set(range(g_start, g_end)) + total_gold_lines += len(gold_set) + for c_start, c_end in ranges_b[f]: + overlapping_lines += len(gold_set & set(range(c_start, c_end))) + + if total_gold_lines == 0: + return 0.0 + return min(overlapping_lines / total_gold_lines, 1.0) + + def proxy_eval(generated_text: str, instance: dict) -> dict: """ Evaluate a generated patch without running the test suite. Returns: - has_diff: bool — model produced a valid unified diff - file_overlap: float — fraction of gold files present in patch - exact_file_match: bool — generated patch touches exactly the right files - gold_files: list[str] - generated_files: list[str] + has_diff: bool — model produced a valid unified diff + file_overlap: float — fraction of gold files present in patch + exact_file_match: bool — generated patch touches exactly the right files + hunk_overlap: float — fraction of gold line ranges covered by generated hunks + content_similarity: float — Jaccard similarity of changed lines (added/removed) """ generated_patch = extract_patch(generated_text) - gold_files = parse_patch_files(instance["patch"]) + gold_patch = instance["patch"] + gold_files = parse_patch_files(gold_patch) generated_files = parse_patch_files(generated_patch) has_diff = is_valid_diff(generated_patch) @@ -204,10 +292,25 @@ def proxy_eval(generated_text: str, instance: dict) -> dict: ) exact_file_match = (gold_files == generated_files) and bool(gold_files) + # Hunk-level: do they modify the same line ranges? + gold_ranges = _parse_hunk_line_ranges(gold_patch) + gen_ranges = _parse_hunk_line_ranges(generated_patch) + hunk_overlap = _line_range_overlap(gold_ranges, gen_ranges) + + # Content-level: Jaccard similarity of the actual changed lines + gold_lines = _extract_changed_lines(gold_patch) + gen_lines = _extract_changed_lines(generated_patch) + if gold_lines or gen_lines: + content_similarity = len(gold_lines & gen_lines) / len(gold_lines | gen_lines) + else: + content_similarity = 0.0 + return { "has_diff": has_diff, "file_overlap": round(file_overlap, 3), "exact_file_match": exact_file_match, + "hunk_overlap": round(hunk_overlap, 3), + "content_similarity": round(content_similarity, 3), "gold_files": sorted(gold_files), "generated_files": sorted(generated_files), } @@ -225,10 +328,13 @@ class SWERunResult: has_diff: bool file_overlap: float exact_file_match: bool + hunk_overlap: float + content_similarity: float prompt_tokens: int completion_tokens: int total_tokens: int latency_ms: float + cost_usd: float = 0.0 compression_ratio: float = 0.0 error: str = "" @@ -238,39 +344,114 @@ class SWERunResult: # --------------------------------------------------------------------------- +def _run_with_retrieval_loop( + model: str, + messages: list[dict], + tools: list[dict], + cache: dict[str, str], + max_retrievals: int = 5, +) -> tuple[str, object, float, float]: + """ + Call the model, and if it invokes litellm_content_retrieve, fulfill + the tool call from the cache and re-call until the model produces a + final text response (or we hit max_retrievals). + + Returns (generated_text, final_usage, total_latency_ms, total_cost). + """ + total_latency = 0.0 + total_cost = 0.0 + total_usage = None + kwargs: dict = { + "model": model, + "messages": list(messages), + "temperature": 0.0, + "max_tokens": 4096, + } + if tools: + kwargs["tools"] = tools + + for _ in range(max_retrievals + 1): + t0 = time.time() + resp = litellm.completion(**kwargs) + total_latency += (time.time() - t0) * 1000 + total_cost += resp._hidden_params.get("response_cost", 0) or 0 + total_usage = resp.usage + + choice = resp.choices[0] + + # If the model produced tool calls, fulfill them and loop + tool_calls = getattr(choice.message, "tool_calls", None) + if tool_calls: + # Append the assistant message with tool calls + kwargs["messages"].append(choice.message.model_dump()) + + for tc in tool_calls: + if tc.function.name == "litellm_content_retrieve": + import json as _json + + args = _json.loads(tc.function.arguments) + key = args.get("key", "") + content = cache.get(key, f"[key {key!r} not found in cache]") + kwargs["messages"].append( + { + "role": "tool", + "tool_call_id": tc.id, + "content": content, + } + ) + else: + kwargs["messages"].append( + { + "role": "tool", + "tool_call_id": tc.id, + "content": "[unknown tool]", + } + ) + continue + + # No tool calls — model produced a final text response + return choice.message.content or "", total_usage, total_latency, total_cost + + # Exhausted retries — return whatever we have + return resp.choices[0].message.content or "", total_usage, total_latency, total_cost + + def eval_instance( instance: dict, model: str, use_compression: bool, compression_trigger: int, + compression_target: Optional[int] = None, embedding_model: Optional[str] = None, ) -> SWERunResult: mode = "compressed" if use_compression else "baseline" messages = build_messages(instance) compression_ratio = 0.0 + tools: list[dict] = [] + cache: dict[str, str] = {} if use_compression: - result = litellm_compress( - messages=messages, - model=model, - compression_trigger=compression_trigger, - embedding_model=embedding_model, - ) + compress_kwargs: dict = { + "messages": messages, + "model": model, + "compression_trigger": compression_trigger, + "embedding_model": embedding_model, + } + if compression_target is not None: + compress_kwargs["compression_target"] = compression_target + result = litellm_compress(**compress_kwargs) messages = result["messages"] + tools = result["tools"] + cache = result["cache"] compression_ratio = result["compression_ratio"] try: - t0 = time.time() - resp = litellm.completion( + generated_text, usage, latency_ms, cost = _run_with_retrieval_loop( model=model, messages=messages, - temperature=0.0, - max_tokens=4096, + tools=tools, + cache=cache, ) - latency_ms = (time.time() - t0) * 1000 - - generated_text = resp.choices[0].message.content or "" - usage = resp.usage ev = proxy_eval(generated_text, instance) return SWERunResult( @@ -279,10 +460,13 @@ def eval_instance( has_diff=ev["has_diff"], file_overlap=ev["file_overlap"], exact_file_match=ev["exact_file_match"], + hunk_overlap=ev["hunk_overlap"], + content_similarity=ev["content_similarity"], prompt_tokens=usage.prompt_tokens, completion_tokens=usage.completion_tokens, total_tokens=usage.total_tokens, latency_ms=latency_ms, + cost_usd=cost, compression_ratio=compression_ratio, ) except Exception as e: @@ -292,6 +476,8 @@ def eval_instance( has_diff=False, file_overlap=0.0, exact_file_match=False, + hunk_overlap=0.0, + content_similarity=0.0, prompt_tokens=0, completion_tokens=0, total_tokens=0, @@ -321,12 +507,18 @@ def aggregate(results: list[SWERunResult]) -> dict: "exact_file_match_rate": round( sum(r.exact_file_match for r in results) / len(results) * 100, 1 ), + "avg_hunk_overlap": round(statistics.mean(r.hunk_overlap for r in results), 3), + "avg_content_similarity": round( + statistics.mean(r.content_similarity for r in results), 3 + ), "avg_prompt_tokens": round(statistics.mean(r.prompt_tokens for r in results)), "avg_total_tokens": round(statistics.mean(r.total_tokens for r in results)), "avg_latency_ms": round(statistics.mean(r.latency_ms for r in results), 1), "avg_compression_ratio": round( statistics.mean(r.compression_ratio for r in results), 4 ), + "total_cost_usd": round(sum(r.cost_usd for r in results), 6), + "avg_cost_usd": round(statistics.mean(r.cost_usd for r in results), 6), } @@ -339,6 +531,7 @@ def run_benchmark( model: str, num_problems: int = 10, compression_trigger: int = 10_000, + compression_target: Optional[int] = None, embedding_model: Optional[str] = None, ) -> dict: """ @@ -359,7 +552,13 @@ def run_benchmark( print(f"{'=' * 60}") print(f"Model: {model}") print(f"Problems: {len(problems)}") + effective_target = ( + compression_target + if compression_target is not None + else compression_trigger * 7 // 10 + ) print(f"Compression trigger: {compression_trigger} tokens") + print(f"Compression target: {effective_target} tokens") print(f"Embedding model: {embedding_model or 'None (BM25 only)'}") print(f"{'=' * 60}\n") @@ -377,6 +576,7 @@ def run_benchmark( model, use_compression=False, compression_trigger=compression_trigger, + compression_target=compression_target, ) baseline_results.append(r_base) if r_base.error: @@ -385,7 +585,8 @@ def run_benchmark( print( f"{'✓' if r_base.has_diff else '✗'} diff " f"file_overlap={r_base.file_overlap:.2f} " - f"{r_base.prompt_tokens} tok" + f"{r_base.prompt_tokens} tok " + f"${r_base.cost_usd:.4f}" ) print(f" compressed ...", end=" ", flush=True) @@ -394,6 +595,7 @@ def run_benchmark( model, use_compression=True, compression_trigger=compression_trigger, + compression_target=compression_target, embedding_model=embedding_model, ) compressed_results.append(r_comp) @@ -404,6 +606,7 @@ def run_benchmark( f"{'✓' if r_comp.has_diff else '✗'} diff " f"file_overlap={r_comp.file_overlap:.2f} " f"{r_comp.prompt_tokens} tok " + f"${r_comp.cost_usd:.4f} " f"(ratio: {r_comp.compression_ratio:.2%})" ) @@ -417,15 +620,23 @@ def run_benchmark( print(f" Has-diff rate: {base_agg['has_diff_rate']}%") print(f" Avg file overlap: {base_agg['avg_file_overlap']:.3f}") print(f" Exact file match: {base_agg['exact_file_match_rate']}%") + print(f" Avg hunk overlap: {base_agg['avg_hunk_overlap']:.3f}") + print(f" Avg content sim: {base_agg['avg_content_similarity']:.3f}") print(f" Avg prompt tokens: {base_agg['avg_prompt_tokens']}") print(f" Avg latency: {base_agg['avg_latency_ms']}ms") + print(f" Total cost: ${base_agg['total_cost_usd']:.4f}") + print(f" Avg cost/problem: ${base_agg['avg_cost_usd']:.6f}") print(f"\n Compressed:") print(f" Has-diff rate: {comp_agg['has_diff_rate']}%") print(f" Avg file overlap: {comp_agg['avg_file_overlap']:.3f}") print(f" Exact file match: {comp_agg['exact_file_match_rate']}%") + print(f" Avg hunk overlap: {comp_agg['avg_hunk_overlap']:.3f}") + print(f" Avg content sim: {comp_agg['avg_content_similarity']:.3f}") print(f" Avg prompt tokens: {comp_agg['avg_prompt_tokens']}") print(f" Avg latency: {comp_agg['avg_latency_ms']}ms") + print(f" Total cost: ${comp_agg['total_cost_usd']:.4f}") + print(f" Avg cost/problem: ${comp_agg['avg_cost_usd']:.6f}") print(f" Avg compression: {comp_agg['avg_compression_ratio']:.2%}") token_savings = base_agg["avg_prompt_tokens"] - comp_agg["avg_prompt_tokens"] @@ -448,6 +659,19 @@ def run_benchmark( print( f" Exact match delta: {comp_agg['exact_file_match_rate'] - base_agg['exact_file_match_rate']:+.1f}%" ) + print( + f" Hunk overlap delta: {comp_agg['avg_hunk_overlap'] - base_agg['avg_hunk_overlap']:+.3f}" + ) + print( + f" Content sim delta: {comp_agg['avg_content_similarity'] - base_agg['avg_content_similarity']:+.3f}" + ) + cost_savings = base_agg["total_cost_usd"] - comp_agg["total_cost_usd"] + cost_pct = ( + round(cost_savings / base_agg["total_cost_usd"] * 100, 1) + if base_agg["total_cost_usd"] + else 0 + ) + print(f" Cost savings: ${cost_savings:.4f} ({cost_pct}%)") ts = time.strftime("%Y-%m-%d_%H-%M-%S") report_path = f"eval_swe_bench_report_{ts}.json" @@ -491,6 +715,13 @@ if __name__ == "__main__": help="Token threshold to activate compression (default: 10000). " "The bm25_27K dataset has ~27k tokens of context per problem.", ) + parser.add_argument( + "--compression-target", + type=int, + default=None, + help="Target token count after compression (default: 70%% of trigger). " + "Higher values preserve more context at the cost of less compression.", + ) parser.add_argument( "--embedding-model", type=str, @@ -503,5 +734,6 @@ if __name__ == "__main__": model=args.model, num_problems=args.problems, compression_trigger=args.compression_trigger, + compression_target=args.compression_target, embedding_model=args.embedding_model, ) diff --git a/tests/test_compression.py b/tests/test_litellm/test_compression.py similarity index 82% rename from tests/test_compression.py rename to tests/test_litellm/test_compression.py index ade89f7a329..13dda0cbcbc 100644 --- a/tests/test_compression.py +++ b/tests/test_litellm/test_compression.py @@ -8,6 +8,7 @@ import pytest import litellm from litellm.compression.scoring.bm25 import bm25_score_messages +from litellm.compression.scoring.embedding_scorer import embedding_score_messages from litellm.compression.content_detection import detect_content_type from litellm.compression.message_stubbing import extract_key, stub_message from litellm.compression.retrieval_tool import build_retrieval_tool @@ -246,6 +247,71 @@ def test_compress_default_target(): assert result["compressed_tokens"] <= result["original_tokens"] +def test_compress_forwards_embedding_model_params(monkeypatch): + captured = {} + + def fake_embedding_score_messages( + query, messages, model, cache=None, embedding_model_params=None + ): + captured["query"] = query + captured["model"] = model + captured["embedding_model_params"] = embedding_model_params + return [0.0] * len(messages) + + monkeypatch.setattr( + "litellm.compression.scoring.embedding_scorer.embedding_score_messages", + fake_embedding_score_messages, + ) + + result = litellm.compress( + messages=[ + {"role": "user", "content": "Authentication code " * 2000}, + {"role": "user", "content": "Fix auth"}, + ], + model="gpt-4o", + compression_trigger=1000, + embedding_model="text-embedding-3-small", + embedding_model_params={"api_base": "https://example-embeddings.test"}, + ) + + assert result["compressed_tokens"] <= result["original_tokens"] + assert captured["model"] == "text-embedding-3-small" + assert captured["embedding_model_params"] == { + "api_base": "https://example-embeddings.test" + } + + +def test_embedding_scorer_forwards_embedding_model_params(monkeypatch): + captured = {} + + class _MockResponse: + data = [ + {"embedding": [1.0, 0.0]}, + {"embedding": [1.0, 0.0]}, + {"embedding": [0.0, 1.0]}, + ] + + def fake_embedding(**kwargs): + captured.update(kwargs) + return _MockResponse() + + monkeypatch.setattr(litellm, "embedding", fake_embedding) + + scores = embedding_score_messages( + query="auth", + messages=[ + {"role": "user", "content": "auth code"}, + {"role": "user", "content": "cooking recipe"}, + ], + model="text-embedding-3-small", + embedding_model_params={"api_base": "https://example-embeddings.test"}, + ) + + assert len(scores) == 2 + assert captured["model"] == "text-embedding-3-small" + assert captured["api_base"] == "https://example-embeddings.test" + + # --------------------------------------------------------------------------- # Embedding scorer — integration test (skipped without API key) # ---------------------------------------------------------------------------