import math from collections import Counter def context_coherence_score( candidate: dict, context_terms: list[str], smoothing: float = 1.0, ) -> float: """ Pointwise mutual information between context_terms and candidate description tokens. Higher score = more context terms co-occur with this entity in training data. """ if not context_terms: return 0.5 # neutral when no context available description_tokens = set(candidate.get("description_tokens", [])) cooccurrence_counts = candidate.get("cooccurrence_counts", {}) total_docs = candidate.get("total_docs_in_corpus", 1) pmi_sum = 0.0 for term in context_terms: p_term = cooccurrence_counts.get(term, 0) / total_docs p_entity = candidate.get("doc_count", 0) / total_docs p_joint = cooccurrence_counts.get(f"joint:{term}", 0) / total_docs if p_joint > 0 and p_term > 0 and p_entity > 0: pmi = math.log2(p_joint / (p_term * p_entity) + smoothing) pmi_sum += max(0.0, pmi) return min(1.0, pmi_sum / (len(context_terms) + smoothing))