Enhance chat memory scoring and token handling by adding expanded query tokens and refining ranking weights.
This commit is contained in:
@@ -316,6 +316,22 @@ def _query_tokens(prompt: str, limit: int = 16) -> List[str]:
|
||||
return tokens
|
||||
|
||||
|
||||
def _expanded_query_tokens(tokens: Sequence[str], limit: int = 24) -> List[str]:
|
||||
expanded: List[str] = []
|
||||
seen = set()
|
||||
for token in tokens:
|
||||
values = [token, *_QUERY_EXPANSIONS.get(token, [])]
|
||||
for value in values:
|
||||
cleaned = str(value or "").casefold().strip()
|
||||
if not cleaned or cleaned in seen:
|
||||
continue
|
||||
seen.add(cleaned)
|
||||
expanded.append(cleaned)
|
||||
if len(expanded) >= limit:
|
||||
return expanded
|
||||
return expanded
|
||||
|
||||
|
||||
def _fts_match_query(tokens: Sequence[str]) -> str:
|
||||
safe_terms = []
|
||||
for token in tokens:
|
||||
@@ -325,6 +341,14 @@ def _fts_match_query(tokens: Sequence[str]) -> str:
|
||||
return " OR ".join(safe_terms)
|
||||
|
||||
|
||||
def _minimum_memory_score(token_count: int) -> float:
|
||||
if token_count <= 1:
|
||||
return 0.85
|
||||
if token_count == 2:
|
||||
return 0.9
|
||||
return 0.65
|
||||
|
||||
|
||||
def _row_mapping(row: Any) -> Dict[str, Any]:
|
||||
if isinstance(row, dict):
|
||||
return row
|
||||
@@ -334,7 +358,13 @@ def _row_mapping(row: Any) -> Dict[str, Any]:
|
||||
return dict(row)
|
||||
|
||||
|
||||
def _score_hit(row: Dict[str, Any], tokens: Sequence[str], fts_rank: Optional[float] = None) -> float:
|
||||
def _score_hit(
|
||||
row: Dict[str, Any],
|
||||
tokens: Sequence[str],
|
||||
fts_rank: Optional[float] = None,
|
||||
*,
|
||||
query_token_count: Optional[int] = None,
|
||||
) -> float:
|
||||
title = str(row.get("session_name") or "").casefold()
|
||||
user_text = str(row.get("user_content") or "").casefold()
|
||||
assistant_text = str(row.get("assistant_content") or "").casefold()
|
||||
@@ -343,19 +373,20 @@ def _score_hit(row: Dict[str, Any], tokens: Sequence[str], fts_rank: Optional[fl
|
||||
hits = [token for token in tokens if token in combined]
|
||||
if not hits:
|
||||
return 0.0
|
||||
if len(tokens) >= 5 and len(hits) < 2:
|
||||
core_count = max(1, int(query_token_count or len(tokens) or 1))
|
||||
if core_count >= 2 and len(hits) < 2:
|
||||
return 0.0
|
||||
score = len(hits) / max(1, min(len(tokens), 8))
|
||||
score = len(hits) / max(1, min(core_count, 4))
|
||||
score += 0.18 * sum(1 for token in hits if token in title)
|
||||
score += 0.10 * sum(1 for token in hits if token in user_text)
|
||||
if fts_rank is not None:
|
||||
try:
|
||||
score += min(0.25, 1.0 / (1.0 + abs(float(fts_rank))))
|
||||
score += min(0.08, 1.0 / (1.0 + abs(float(fts_rank))))
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
assistant_row = int(row.get("assistant_message_row_id") or 0)
|
||||
score += min(0.12, assistant_row / 1_000_000)
|
||||
score += min(0.04, assistant_row / 2_000_000)
|
||||
except Exception:
|
||||
pass
|
||||
return score
|
||||
@@ -367,6 +398,7 @@ def _load_fts_candidates(
|
||||
*,
|
||||
exclude_session_id: Optional[str],
|
||||
limit: int,
|
||||
query_token_count: int,
|
||||
) -> List[Dict[str, Any]]:
|
||||
match = _fts_match_query(tokens)
|
||||
if not match:
|
||||
@@ -396,7 +428,7 @@ def _load_fts_candidates(
|
||||
candidates: List[Dict[str, Any]] = []
|
||||
for row in rows:
|
||||
item = _row_mapping(row)
|
||||
item["_score"] = _score_hit(item, tokens, item.get("fts_rank"))
|
||||
item["_score"] = _score_hit(item, tokens, item.get("fts_rank"), query_token_count=query_token_count)
|
||||
if item["_score"] > 0:
|
||||
candidates.append(item)
|
||||
return candidates
|
||||
@@ -408,6 +440,7 @@ def _load_table_candidates(
|
||||
*,
|
||||
exclude_session_id: Optional[str],
|
||||
limit: int,
|
||||
query_token_count: int,
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
rows = db.execute(text(
|
||||
@@ -430,7 +463,7 @@ def _load_table_candidates(
|
||||
candidates: List[Dict[str, Any]] = []
|
||||
for row in rows:
|
||||
item = _row_mapping(row)
|
||||
item["_score"] = _score_hit(item, tokens)
|
||||
item["_score"] = _score_hit(item, tokens, query_token_count=query_token_count)
|
||||
if item["_score"] > 0:
|
||||
candidates.append(item)
|
||||
return candidates
|
||||
@@ -442,6 +475,7 @@ def _load_live_candidates(
|
||||
*,
|
||||
exclude_session_id: Optional[str],
|
||||
limit: int,
|
||||
query_token_count: int,
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
rows = db.execute(text(
|
||||
@@ -492,7 +526,7 @@ def _load_live_candidates(
|
||||
"attachment_summary": attachments,
|
||||
"created_at": str(item.get("created_at") or ""),
|
||||
}
|
||||
turn["_score"] = _score_hit(turn, tokens)
|
||||
turn["_score"] = _score_hit(turn, tokens, query_token_count=query_token_count)
|
||||
if turn["_score"] > 0:
|
||||
turns.append(turn)
|
||||
return turns
|
||||
|
||||
Reference in New Issue
Block a user