"""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, )