Spaces:
Sleeping
Sleeping
| """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 | |
| # --------------------------------------------------------------------------- | |
| class PruningConfig: | |
| max_input_tokens: int = 1500 | |
| segment_tokens: int = 384 | |
| mmr_lambda: float = 0.7 | |
| 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, | |
| ) | |