""" Clause retrieval module. Builds a BM25 + embedding index over a clause corpus. Retrieves relevant precedent clauses for a drafting query. """ import json import pickle from typing import List, Dict, Tuple, Optional import numpy as np try: from rank_bm25 import BM25Okapi except ImportError: BM25Okapi = None try: from sentence_transformers import SentenceTransformer, util except ImportError: SentenceTransformer = None class ClauseRetriever: def __init__( self, embedding_model_name: str = "sentence-transformers/all-MiniLM-L6-v2", use_bm25: bool = True, use_embeddings: bool = True, ): self.use_bm25 = use_bm25 and BM25Okapi is not None self.use_embeddings = use_embeddings and SentenceTransformer is not None self.embedding_model_name = embedding_model_name self.bm25 = None self.corpus: List[Dict] = [] self.tokenized_corpus: List[List[str]] = [] self.embeddings: Optional[np.ndarray] = None self.embedding_model = None if self.use_embeddings: self.embedding_model = SentenceTransformer(embedding_model_name) def _tokenize(self, text: str) -> List[str]: return text.lower().split() def add_clauses(self, clauses: List[Dict[str, str]]): """ clauses: list of dicts with keys 'clause_text', 'clause_type', 'source', etc. """ self.corpus.extend(clauses) if self.use_bm25: self.tokenized_corpus = [self._tokenize(c["clause_text"]) for c in self.corpus] self.bm25 = BM25Okapi(self.tokenized_corpus) if self.use_embeddings and self.embedding_model is not None: texts = [c["clause_text"] for c in self.corpus] self.embeddings = self.embedding_model.encode( texts, show_progress_bar=True, convert_to_numpy=True ) def retrieve( self, query: str, clause_type: Optional[str] = None, top_k: int = 5, bm25_weight: float = 0.3, embedding_weight: float = 0.7, ) -> List[Dict]: if not self.corpus: return [] scores = np.zeros(len(self.corpus)) if self.use_bm25 and self.bm25 is not None: tokenized_query = self._tokenize(query) bm25_scores = np.array(self.bm25.get_scores(tokenized_query)) if bm25_scores.max() > 0: bm25_scores = bm25_scores / bm25_scores.max() scores += bm25_weight * bm25_scores if self.use_embeddings and self.embedding_model is not None and self.embeddings is not None: query_emb = self.embedding_model.encode(query, convert_to_numpy=True) sims = util.cos_sim(query_emb, self.embeddings)[0].cpu().numpy() scores += embedding_weight * sims # Filter by clause_type if requested indices = list(range(len(self.corpus))) if clause_type: indices = [i for i in indices if self.corpus[i].get("clause_type") == clause_type] ranked = sorted(indices, key=lambda i: scores[i], reverse=True)[:top_k] results = [] for i in ranked: item = dict(self.corpus[i]) item["score"] = float(scores[i]) results.append(item) return results def save(self, path_prefix: str): meta = { "corpus": self.corpus, "embedding_model_name": self.embedding_model_name, "use_bm25": self.use_bm25, "use_embeddings": self.use_embeddings, } with open(path_prefix + "_meta.json", "w") as f: json.dump(meta, f) if self.embeddings is not None: np.save(path_prefix + "_embeddings.npy", self.embeddings) if self.bm25 is not None: with open(path_prefix + "_bm25.pkl", "wb") as f: pickle.dump(self.bm25, f) def load(self, path_prefix: str): with open(path_prefix + "_meta.json", "r") as f: meta = json.load(f) self.corpus = meta["corpus"] self.embedding_model_name = meta["embedding_model_name"] self.use_bm25 = meta["use_bm25"] self.use_embeddings = meta["use_embeddings"] if self.use_bm25: with open(path_prefix + "_bm25.pkl", "rb") as f: self.bm25 = pickle.load(f) self.tokenized_corpus = [self._tokenize(c["clause_text"]) for c in self.corpus] if self.use_embeddings: self.embeddings = np.load(path_prefix + "_embeddings.npy") if self.embedding_model is None: self.embedding_model = SentenceTransformer(self.embedding_model_name) def build_retriever_from_hf_datasets( clause_dataset_name: str = "asapworks/Contract_Clause_SampleDataset", contract_dataset_name: str = "albertvillanova/legal_contracts", max_contracts: int = 500, max_clauses_per_contract: int = 20, ) -> ClauseRetriever: from datasets import load_dataset retriever = ClauseRetriever() # Load labeled clause dataset try: ds = load_dataset(clause_dataset_name, split="train") for row in ds: retriever.add_clauses([{ "clause_text": row["clause_text"], "clause_type": row.get("clause_type", "unknown"), "source": row.get("file", clause_dataset_name), }]) except Exception as e: print(f"Warning: could not load {clause_dataset_name}: {e}") # Load raw contracts and chunk for retrieval corpus try: ds = load_dataset(contract_dataset_name, split="train", streaming=True) count = 0 for row in ds: text = row["text"] # Simple paragraph chunking paragraphs = [p.strip() for p in text.split("\n\n") if len(p.strip()) > 100] for para in paragraphs[:max_clauses_per_contract]: retriever.add_clauses([{ "clause_text": para, "clause_type": "unknown", "source": contract_dataset_name, }]) count += 1 if count >= max_contracts: break except Exception as e: print(f"Warning: could not load {contract_dataset_name}: {e}") return retriever