patent-sum / backend /pruning.py
ishma03
Initial deployment
1cc3c69
Raw
History Blame Contribute Delete
4.89 kB
"""Pruning logic — segmentation + MMR selection, calls embedding server for vectors."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Callable, List, Tuple
import numpy as np
from nltk.tokenize import sent_tokenize
from sklearn.metrics.pairwise import cosine_similarity
from segmentation import PatentSegment, segment_by_structure
# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
@dataclass
class PruningConfig:
max_input_tokens: int = 1500
segment_tokens: int = 384
mmr_lambda: float = 0.7
@dataclass
class PruningResult:
pruned_text: str
original_tokens: int
pruned_tokens: int
num_segments: int
segments_selected: int
sections_used: List[str]
compression_ratio: float
was_pruned: bool
def estimate_tokens(text: str) -> int:
"""Rough token estimate (words * 1.3). No tokenizer dependency."""
return int(len(text.split()) * 1.3)
# ---------------------------------------------------------------------------
# MMR selection
# ---------------------------------------------------------------------------
def mmr_select_ordered(
segments: List[PatentSegment],
embeddings: np.ndarray,
scores: np.ndarray,
max_tokens: int,
mmr_lambda: float = 0.7,
) -> Tuple[List[int], int]:
selected_ids: list[int] = []
remaining = list(range(len(segments)))
total_tokens = 0
while remaining:
mmr_scores: list[tuple[int, float]] = []
for i in remaining:
if not selected_ids:
mmr = float(scores[i])
else:
redundancy = float(np.max(
cosine_similarity(
embeddings[i].reshape(1, -1),
embeddings[selected_ids],
)[0]
))
mmr = mmr_lambda * scores[i] - (1 - mmr_lambda) * redundancy
mmr_scores.append((i, mmr))
best_idx = max(mmr_scores, key=lambda x: x[1])[0]
seg_tokens = estimate_tokens(segments[best_idx].text)
if total_tokens + seg_tokens > max_tokens:
remaining.remove(best_idx)
continue
selected_ids.append(best_idx)
total_tokens += seg_tokens
remaining.remove(best_idx)
selected_ids.sort(key=lambda i: segments[i].original_position)
return selected_ids, total_tokens
# ---------------------------------------------------------------------------
# Main pruning (async — calls embedding server)
# ---------------------------------------------------------------------------
async def prune_document(
text: str,
config: PruningConfig | None = None,
embed_fn: Callable | None = None,
) -> PruningResult:
if config is None:
config = PruningConfig()
original_tokens = estimate_tokens(text)
if original_tokens <= config.max_input_tokens:
return PruningResult(
pruned_text=text,
original_tokens=original_tokens,
pruned_tokens=original_tokens,
num_segments=1,
segments_selected=1,
sections_used=[],
compression_ratio=1.0,
was_pruned=False,
)
segments = segment_by_structure(
text,
max_tokens=config.segment_tokens,
token_counter=estimate_tokens,
)
if not segments:
return PruningResult(
pruned_text="",
original_tokens=original_tokens,
pruned_tokens=0,
num_segments=0,
segments_selected=0,
sections_used=[],
compression_ratio=0.0,
was_pruned=True,
)
# Get embeddings from embedding server
segment_texts = [seg.text for seg in segments]
embeddings_list = await embed_fn(segment_texts)
seg_embeddings = np.array(embeddings_list)
# Centrality scoring
doc_centroid = np.mean(seg_embeddings, axis=0)
scores = cosine_similarity(
seg_embeddings, doc_centroid.reshape(1, -1)
).flatten()
# MMR selection
selected_ids, pruned_tokens = mmr_select_ordered(
segments=segments,
embeddings=seg_embeddings,
scores=scores,
max_tokens=config.max_input_tokens,
mmr_lambda=config.mmr_lambda,
)
pruned_text = "\n\n".join(segments[i].text for i in selected_ids)
sections_used = [segments[i].section_type for i in selected_ids]
return PruningResult(
pruned_text=pruned_text,
original_tokens=original_tokens,
pruned_tokens=pruned_tokens,
num_segments=len(segments),
segments_selected=len(selected_ids),
sections_used=sections_used,
compression_ratio=pruned_tokens / max(original_tokens, 1),
was_pruned=True,
)