|
- from collections import Counter
- import logging
- from math import sqrt
- from collections.abc import Sequence
-
- logger = logging.getLogger(__name__)
-
-
- def cosine_similarity(a: Sequence[float], b: Sequence[float]) -> float:
- logger.debug("Computing cosine similarity between numerical vectors of length %d and %d", len(a), len(b))
- if len(a) != len(b):
- logger.error("Vector dimension mismatch: vector 'a' length (%d) != vector 'b' length (%d)", len(a), len(b))
- raise ValueError("Vectors must have the same dimension")
-
- dot = 0.0
- norm_a_sq = 0.0
- norm_b_sq = 0.0
-
- for x, y in zip(a, b):
- dot += x * y
- norm_a_sq += x * x
- norm_b_sq += y * y
-
- denominator = sqrt(norm_a_sq * norm_b_sq)
- if denominator == 0.0:
- logger.debug("Zero denominator encountered in cosine_similarity (norm_a_sq=%f, norm_b_sq=%f). Returning 0.0", norm_a_sq, norm_b_sq)
- return 0.0
-
- similarity = dot / denominator
- logger.debug("Calculated vector cosine similarity: dot=%f, denominator=%f, similarity=%f", dot, denominator, similarity)
- return similarity
-
-
- def cosine_lists(a: list[str], b: list[str], *, casefold: bool = True) -> float:
- logger.debug("Computing token cosine similarity for list_a=%s and list_b=%s (casefold=%s)", a, b, casefold)
-
- def tokens(xs: list[str]) -> Counter[str]:
- return Counter(x.casefold() if casefold else x for x in xs)
-
- ca, cb = tokens(a), tokens(b)
- if not ca or not cb:
- logger.debug("Empty token set detected (count_a=%d, count_b=%d). Cosine similarity is 0.0", len(ca), len(cb))
- return 0.0
-
- common_tokens = ca.keys() & cb.keys()
- dot = sum(ca[t] * cb[t] for t in common_tokens)
- norm_a = sqrt(sum(v * v for v in ca.values()))
- norm_b = sqrt(sum(v * v for v in cb.values()))
-
- if norm_a == 0.0 or norm_b == 0.0:
- logger.debug("Zero norm detected (norm_a=%f, norm_b=%f). Cosine similarity is 0.0", norm_a, norm_b)
- return 0.0
-
- similarity = dot / (norm_a * norm_b)
- logger.debug("Token similarity calculation: common_tokens=%s, dot=%f, norm_a=%f, norm_b=%f -> similarity=%.4f",
- list(common_tokens), dot, norm_a, norm_b, similarity)
- return similarity
|