74 lines
3.4 KiB
Python
74 lines
3.4 KiB
Python
"""Поиск ниши по индексу: сперва синонимы (бесплатно), затем эмбеддинг-косинус.
|
||
Лесенка: niche (нашли нишу) → domain (нашли только домен) → miss (совсем новое)."""
|
||
from dataclasses import dataclass
|
||
import math
|
||
|
||
# пороги подобраны на живом прогоне 17.07.2026 (text-embedding-3-small, 4 карточки):
|
||
# смежные ниши давали косинус 0.48–0.58, домен ~0.34, мусор 0.24–0.27 (спека §9).
|
||
# ⚠️ Провизорные — перепроверить/перетюнить, когда карточек станет заметно больше.
|
||
THR_NICHE = 0.45
|
||
THR_DOMAIN = 0.33
|
||
|
||
|
||
@dataclass
|
||
class Result:
|
||
kind: str # "niche" | "domain" | "miss"
|
||
entry: dict | None # запись индекса или None
|
||
score: float = 0.0
|
||
|
||
|
||
def cosine(a, b) -> float:
|
||
dot = sum(x * y for x, y in zip(a, b))
|
||
na = math.sqrt(sum(x * x for x in a))
|
||
nb = math.sqrt(sum(y * y for y in b))
|
||
if na == 0 or nb == 0:
|
||
return 0.0
|
||
return dot / (na * nb)
|
||
|
||
|
||
def _norm(s: str) -> str:
|
||
return s.strip().lower()
|
||
|
||
|
||
def _synonym_hit(spoken: str, index):
|
||
q = _norm(spoken)
|
||
if not q: # пустой/пробельный запрос — не совпадение (было: "" in n = всегда True)
|
||
return None
|
||
for e in index:
|
||
names = [_norm(e["nisha"])] + [_norm(s) for s in e.get("sinonimy", [])]
|
||
if q in names: # ТОЛЬКО точное совпадение после нормализации (без подстроки)
|
||
return e
|
||
return None
|
||
|
||
|
||
def find_niche(spoken, index, embed_fn, thr_niche=THR_NICHE, thr_domain=THR_DOMAIN) -> Result:
|
||
# 1) дешёвый слой: точное/подстрочное совпадение имени или синонима
|
||
hit = _synonym_hit(spoken, index)
|
||
if hit is not None:
|
||
kind = "domain" if hit.get("roditel") in (None, "", "—") else "niche"
|
||
return Result(kind=kind, entry=hit, score=1.0)
|
||
if embed_fn is None or not index:
|
||
return Result(kind="miss", entry=None)
|
||
# 2) смысловой слой: ближайший по косинусу
|
||
qv = embed_fn(spoken)
|
||
best, best_score = None, -1.0
|
||
for e in index:
|
||
sc = cosine(qv, e["vector"])
|
||
if sc > best_score:
|
||
best, best_score = e, sc
|
||
if best is None:
|
||
return Result(kind="miss", entry=None)
|
||
is_domain = best.get("roditel") in (None, "", "—")
|
||
if not is_domain and best_score >= thr_niche:
|
||
return Result(kind="niche", entry=best, score=best_score)
|
||
if best_score >= thr_domain:
|
||
# ближе к домену: если лучший — ниша, поднимаемся к её родителю-домену
|
||
if is_domain:
|
||
return Result(kind="domain", entry=best, score=best_score)
|
||
parent_name = best.get("roditel")
|
||
parent = next((e for e in index if _norm(e["nisha"]) == _norm(parent_name)), None)
|
||
if parent is not None:
|
||
return Result(kind="domain", entry=parent, score=best_score)
|
||
return Result(kind="miss", entry=None, score=best_score) # родителя нет в индексе → НЕ выдаём лист за домен
|
||
return Result(kind="miss", entry=None, score=best_score)
|