Files
ai-agent-book/chapter2/context-compression/benchmark_compression.py
T
liqiang b119135836
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
2026-08-20 13:12:50 +00:00

340 lines
14 KiB
Python

"""
Context Compression Benchmark Module.
Systematically benchmarks Summary, Truncation, Key-Sentence, and Observation-Filtering
compression strategies on long-context tasks. Measures compression ratio, Time-to-First-Token (TTFT),
token cost savings, and downstream QA retention accuracy.
"""
import math
import re
import time
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union, Tuple
def count_tokens(text: str) -> int:
"""Estimate token count for a given text string.
Uses tiktoken if available, with a reliable character/word-based fallback.
"""
if not text:
return 0
try:
import tiktoken
try:
encoding = tiktoken.encoding_for_model("gpt-4")
except Exception:
encoding = tiktoken.get_encoding("cl100k_base")
return len(encoding.encode(text))
except Exception:
# Fallback estimation: ~4 chars per token or ~0.75 words per token
words = len(text.split())
chars = len(text)
return max(1, int((words * 1.3 + chars / 4) / 2))
@dataclass
class StrategyMetrics:
"""Performance metrics for a context compression strategy."""
strategy: str
original_tokens: int
compressed_tokens: int
compression_ratio: float # compressed_tokens / original_tokens
ttft_ms: float # Time to first token in milliseconds
token_cost_savings: float # Cost savings ratio (0.0 to 1.0)
qa_retention_accuracy: float # Downstream QA accuracy (0.0 to 1.0)
def to_dict(self) -> Dict[str, Any]:
"""Convert metrics to a standard dictionary representation."""
return {
"strategy": self.strategy,
"original_tokens": self.original_tokens,
"compressed_tokens": self.compressed_tokens,
"compression_ratio": self.compression_ratio,
"ttft_ms": self.ttft_ms,
"token_cost_savings": self.token_cost_savings,
"qa_retention_accuracy": self.qa_retention_accuracy,
}
class ContextCompressionBenchmark:
"""Benchmark harness for evaluating context compression strategies."""
STRATEGIES = ["summary", "truncation", "key_sentence", "observation_filtering"]
def __init__(
self,
base_ttft_ms: float = 50.0,
per_token_ttft_ms: float = 0.05,
token_cost_per_1k: float = 0.0015,
target_max_tokens: int = 500,
):
"""Initialize the benchmark suite with configurable performance parameters."""
self.base_ttft_ms = base_ttft_ms
self.per_token_ttft_ms = per_token_ttft_ms
self.token_cost_per_1k = token_cost_per_1k
self.target_max_tokens = target_max_tokens
def compress_summary(self, context: str, query: str = "") -> str:
"""Summary Strategy: Condenses context into key abstract points."""
if not context:
return ""
sentences = [s.strip() for s in re.split(r'(?<=[.!?])\s+', context) if s.strip()]
if not sentences:
return context
if len(sentences) <= 3:
return context
# Extract beginning, middle, and end sentences to form a concise summary
step = max(1, len(sentences) // 3)
summary_sentences = [sentences[0]]
if step < len(sentences):
summary_sentences.append(sentences[step])
if len(sentences) - 1 > step:
summary_sentences.append(sentences[-1])
return " ".join(summary_sentences)
def compress_truncation(self, context: str, max_tokens: Optional[int] = None) -> str:
"""Truncation Strategy: Slices context to fit within strict token limits."""
if not context:
return ""
limit = self.target_max_tokens if max_tokens is None else max_tokens
if limit <= 0:
return ""
words = context.split()
if not words:
# No whitespace-separated words (e.g. CJK text): truncate by characters.
# CJK characters are roughly 1-2 tokens each, so use a conservative 1:1 ratio.
return context[:limit]
# Estimate max words corresponding to limit tokens (~0.75 words per token)
max_words = max(1, int(limit * 0.75))
truncated_words = words[:max_words]
return " ".join(truncated_words)
def compress_key_sentence(self, context: str, query: str = "") -> str:
"""Key-Sentence Strategy: Retains sentences with high query term match/relevance."""
if not context:
return ""
query = query or ""
sentences = [s.strip() for s in re.split(r'(?<=[.!?])\s+', context) if s.strip()]
if not sentences:
return context
if not query:
# Fallback to sentence length / position scoring if query is empty
scored = sorted(enumerate(sentences), key=lambda x: len(x[1]), reverse=True)
top_indices = sorted([idx for idx, _ in scored[:max(1, len(sentences) // 2)]])
return " ".join([sentences[i] for i in top_indices])
query_terms = set(re.findall(r'\w+', query.lower()))
scored_sentences = []
for idx, sentence in enumerate(sentences):
sentence_terms = set(re.findall(r'\w+', sentence.lower()))
overlap = len(query_terms.intersection(sentence_terms))
scored_sentences.append((overlap, idx, sentence))
# Sort by overlap descending, then by original position
scored_sentences.sort(key=lambda x: (-x[0], x[1]))
# Keep top half of sentences or those with overlap > 0
keep_count = max(1, math.ceil(len(sentences) * 0.5))
selected = scored_sentences[:keep_count]
# Sort selected back into original context order
selected.sort(key=lambda x: x[1])
return " ".join([s[2] for s in selected])
def compress_observation_filtering(self, context: str) -> str:
"""Observation-Filtering Strategy: Removes verbose system output, logs, hex, and JSON blobs."""
if not context:
return ""
lines = context.splitlines()
filtered_lines = []
for line in lines:
stripped = line.strip()
# Filter out JSON-like blobs, long hex hashes, trace logs, or repetitive debug markers
if (
re.match(r'^\s*[\{\[\}\]].*$', stripped) or
re.search(r'\b[0-9a-fA-F]{32,64}\b', stripped) or
re.search(r'^\s*(DEBUG|TRACE|INFO|VERBOSE)\b', stripped, re.IGNORECASE) or
re.search(r'^\s*<.*?>\s*$', stripped)
):
continue
filtered_lines.append(line)
result = "\n".join(filtered_lines).strip()
return result if result else context
def compress(self, strategy: str, context: str, query: str = "") -> str:
"""Apply a specific compression strategy to a given context string."""
strat = strategy.lower().replace("-", "_")
if strat == "summary":
return self.compress_summary(context, query)
elif strat == "truncation":
return self.compress_truncation(context)
elif strat in ("key_sentence", "keysentence"):
return self.compress_key_sentence(context, query)
elif strat in ("observation_filtering", "observationfiltering"):
return self.compress_observation_filtering(context)
else:
raise ValueError(f"Unknown compression strategy: {strategy}")
def evaluate_retention(self, compressed_text: str, task: Union[str, Dict[str, Any]]) -> Optional[float]:
"""Evaluate downstream QA retention accuracy on compressed context."""
compressed_text = compressed_text or ""
task = task or ""
query = task if isinstance(task, str) else (task.get("query", "") if isinstance(task, dict) else "")
expected = task.get("expected_answer", "") if isinstance(task, dict) else ""
if query is None:
query = ""
if expected is None:
expected = ""
# Only score against the expected answer, not the query.
# Using query words as fallback inflates scores because the question
# text often survives compression even when the answer is deleted.
target_text = expected.strip()
target_tokens = set(re.findall(r'\w+', target_text.lower()))
if not target_tokens:
# No expected answer to check against: cannot evaluate retention.
return None
compressed_tokens = set(re.findall(r'\w+', compressed_text.lower()))
matched = target_tokens.intersection(compressed_tokens)
# Calculate recall accuracy
accuracy = len(matched) / len(target_tokens)
return min(1.0, max(0.0, accuracy))
def evaluate_strategy(
self,
strategy: str,
contexts: List[str],
tasks: List[Union[str, Dict[str, Any]]],
) -> StrategyMetrics:
"""Benchmark a single compression strategy over multiple contexts and tasks."""
total_orig_tokens = 0
total_comp_tokens = 0
total_retention_acc = 0.0
retention_count = 0
sample_count = 0
start_time = time.perf_counter()
for idx, ctx in enumerate(contexts):
task = tasks[idx % len(tasks)] if tasks else ""
if task is None:
task = ""
query = task if isinstance(task, str) else (task.get("query", "") if isinstance(task, dict) else "")
query = query or ""
orig_tokens = count_tokens(ctx)
compressed_ctx = self.compress(strategy, ctx, query=query)
comp_tokens = count_tokens(compressed_ctx)
retention_acc = self.evaluate_retention(compressed_ctx, task)
total_orig_tokens += orig_tokens
total_comp_tokens += comp_tokens
if retention_acc is not None:
total_retention_acc += retention_acc
retention_count += 1
sample_count += 1
elapsed_ms = (time.perf_counter() - start_time) * 1000
avg_orig_tokens = total_orig_tokens / max(1, sample_count)
avg_comp_tokens = total_comp_tokens / max(1, sample_count)
avg_retention_acc = total_retention_acc / max(1, retention_count)
if avg_orig_tokens == 0:
ratio = 0.0
savings = 0.0
else:
ratio = avg_comp_tokens / avg_orig_tokens
savings = max(0.0, 1.0 - ratio)
# Simulate TTFT: Base TTFT + processing time + prefill latency based on compressed tokens
simulated_ttft = self.base_ttft_ms + (avg_comp_tokens * self.per_token_ttft_ms) + (elapsed_ms / max(1, sample_count))
# Format normalized strategy key
strat_key = strategy.lower().replace("-", "_")
return StrategyMetrics(
strategy=strat_key,
original_tokens=int(avg_orig_tokens),
compressed_tokens=int(avg_comp_tokens),
compression_ratio=round(ratio, 4),
ttft_ms=round(simulated_ttft, 2),
token_cost_savings=round(savings, 4),
qa_retention_accuracy=round(avg_retention_acc, 4),
)
def run_benchmark(
self,
contexts: Union[str, List[Union[str, Dict[str, Any]]]],
tasks: Union[str, List[Union[str, Dict[str, Any]]]],
) -> Dict[str, Any]:
"""Run systematic benchmark across all compression strategies.
Args:
contexts: Single context string, dict, or list of context strings/dicts.
tasks: Single task/query string, dict, or list of tasks/queries.
Returns:
Comparative metrics dictionary mapping strategy names to performance metrics dicts.
"""
# Standardize contexts into list of text strings
if isinstance(contexts, (str, dict)):
raw_contexts = [contexts]
else:
raw_contexts = list(contexts)
normalized_contexts = []
for c in raw_contexts:
if isinstance(c, str):
normalized_contexts.append(c)
elif isinstance(c, dict):
content = c.get("content")
if content is None:
content = c.get("text")
# Use the extracted content, or empty string if none found.
# Falling back to str(c) would treat the raw dict repr as
# context text, producing nonsensical benchmark metrics.
normalized_contexts.append(content if content is not None else "")
else:
normalized_contexts.append(str(c))
# Standardize tasks into list of queries/task objects
if isinstance(tasks, (str, dict)):
normalized_tasks = [tasks]
else:
normalized_tasks = list(tasks)
results: Dict[str, Any] = {}
for strategy in self.STRATEGIES:
metrics = self.evaluate_strategy(strategy, normalized_contexts, normalized_tasks)
metrics_dict = metrics.to_dict()
display_name = {
"summary": "Summary",
"truncation": "Truncation",
"key_sentence": "Key-Sentence",
"observation_filtering": "Observation-Filtering",
}.get(strategy, strategy)
metrics_dict["display_name"] = display_name
results[strategy] = metrics_dict
return results
def run_benchmark(
contexts: Union[str, List[Union[str, Dict[str, Any]]]],
tasks: Union[str, List[Union[str, Dict[str, Any]]]],
) -> Dict[str, Any]:
"""Module-level entrypoint for executing the compression benchmark.
Args:
contexts: Input contexts (strings or dicts).
tasks: Downstream QA tasks or queries.
Returns:
Dictionary of comparative performance metrics per compression strategy.
"""
benchmark = ContextCompressionBenchmark()
return benchmark.run_benchmark(contexts, tasks)