/
SupremeSoviet
/
llm-memory
Обзор
Документация
Войти
/
SupremeSoviet
/
llm-memory
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
eval/eval_longmemeval_retrieval.py
378 строк
15 KB
Vladimir Gubin
feat: add local memory dialogue system
25 май 2026, 01:05
25 май 2026, 01:05
0f5c6d4
Код
Авторство
О чём код?
"""Evaluate LongMemEval retrieval with simple hashing or HF E5 embeddings.""" from __future__ import annotations import argparse import hashlib import json import math from dataclasses import dataclass from pathlib import Path from typing import Any, Protocol TOP_K_LIMITS = (1, 3, 5) class Embedder(Protocol): def embed(self, text: str, input_type: str) -> list[float]: ... def embed_many(self, texts: list[str], input_type: str) -> list[list[float]]: ... @dataclass(frozen=True) class Document: identifier: str text: str score: float = 0.0 def parse_arguments() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--granularity", choices=["session", "turn"], default="session") parser.add_argument("--embedder", choices=["hashing", "hf-e5"], default="hashing") parser.add_argument("--model-name", default="intfloat/multilingual-e5-base") parser.add_argument("--limit", type=int, default=None) parser.add_argument( "--audit-output", type=Path, default=None, help="Optional JSONL path for per-query metric audit rows.", ) return parser.parse_args() def main() -> int: arguments = parse_arguments() instances = json.loads(arguments.input.read_text(encoding="utf-8")) if arguments.limit is not None: instances = instances[: arguments.limit] embedder: Embedder if arguments.embedder == "hf-e5": embedder = HFE5Embedder(arguments.model_name) else: embedder = HashingEmbedder() report, audit_rows = evaluate_instances( instances=instances, embedder=embedder, granularity=arguments.granularity, input_path=arguments.input, embedder_name=arguments.embedder, model_name=arguments.model_name if arguments.embedder == "hf-e5" else None, ) audit_output = arguments.audit_output or default_audit_output(arguments.output) report["metric_definitions"] = metric_definitions() report["audit_output"] = str(audit_output) validate_metric_invariants(report) arguments.output.parent.mkdir(parents=True, exist_ok=True) arguments.output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") write_jsonl(audit_output, audit_rows) print(json.dumps(report, ensure_ascii=False, indent=2)) return 0 def evaluate_instances( *, instances: list[dict[str, Any]], embedder: Embedder, granularity: str, input_path: Path, embedder_name: str, model_name: str | None, ) -> tuple[dict[str, Any], list[dict[str, Any]]]: precision_totals = {limit: 0.0 for limit in TOP_K_LIMITS} recall_totals = {limit: 0.0 for limit in TOP_K_LIMITS} f1_totals = {limit: 0.0 for limit in TOP_K_LIMITS} hit_rate_totals = {limit: 0.0 for limit in TOP_K_LIMITS} reciprocal_rank_total = 0.0 evaluated = 0 skipped_abstention = 0 skipped_no_ground_truth = 0 question_type_counts: dict[str, int] = {} audit_rows: list[dict[str, Any]] = [] for instance in instances: if is_abstention(instance): skipped_abstention += 1 continue documents, relevant_ids = build_documents(instance, granularity=granularity) if not relevant_ids: skipped_no_ground_truth += 1 continue retrieved = retrieve( query=str(instance["question"]), documents=documents, embedder=embedder, top_k=5, ) retrieved_ids = [document.identifier for document in retrieved] query_metrics = calculate_query_metrics(relevant_ids, retrieved_ids) for limit in TOP_K_LIMITS: precision_totals[limit] += query_metrics[f"precision_at_{limit}"] recall_totals[limit] += query_metrics[f"recall_at_{limit}"] f1_totals[limit] += query_metrics[f"f1_at_{limit}"] hit_rate_totals[limit] += query_metrics[f"hit_rate_at_{limit}"] reciprocal_rank_total += query_metrics["reciprocal_rank_at_5"] evaluated += 1 question_type = str(instance.get("question_type", "unknown")) question_type_counts[question_type] = question_type_counts.get(question_type, 0) + 1 audit_rows.append( { "query_id": str(instance.get("question_id", evaluated)), "question_type": question_type, "granularity": granularity, "relevant_count": query_metrics["relevant_count"], "relevant_ids": relevant_ids, "retrieved_ids": retrieved_ids, **query_metrics, } ) if evaluated % 25 == 0: print(f"[retrieval:{embedder_name}:{granularity}] evaluated {evaluated}", flush=True) if evaluated == 0: raise ValueError("No LongMemEval instances with retrieval ground truth were evaluated.") report: dict[str, Any] = { "input_path": str(input_path), "granularity": granularity, "embedder": embedder_name, "model_name": model_name, "evaluated_query_count": evaluated, "skipped_abstention_count": skipped_abstention, "skipped_no_ground_truth_count": skipped_no_ground_truth, "question_type_counts": question_type_counts, "mrr_at_5": reciprocal_rank_total / evaluated, "mean_reciprocal_rank": reciprocal_rank_total / evaluated, } for limit in TOP_K_LIMITS: report[f"precision_at_{limit}"] = precision_totals[limit] / evaluated report[f"recall_at_{limit}"] = recall_totals[limit] / evaluated report[f"f1_at_{limit}"] = f1_totals[limit] / evaluated report[f"hit_rate_at_{limit}"] = hit_rate_totals[limit] / evaluated return report, audit_rows def build_documents(instance: dict[str, Any], *, granularity: str) -> tuple[list[Document], list[str]]: session_ids = [str(identifier) for identifier in instance.get("haystack_session_ids", [])] sessions = instance.get("haystack_sessions", []) if granularity == "session": documents = [ Document(identifier=session_id, text=format_session(session)) for session_id, session in zip(session_ids, sessions) ] relevant_ids = [str(identifier) for identifier in instance.get("answer_session_ids", [])] return documents, relevant_ids documents: list[Document] = [] relevant_ids: list[str] = [] for session_id, session in zip(session_ids, sessions): for turn_index, turn in enumerate(session): identifier = f"{session_id}#{turn_index}" documents.append(Document(identifier=identifier, text=format_turn(turn))) if bool(turn.get("has_answer")): relevant_ids.append(identifier) return documents, relevant_ids def format_session(session: list[dict[str, Any]]) -> str: return "\n".join(format_turn(turn) for turn in session) def format_turn(turn: dict[str, Any]) -> str: role = str(turn.get("role", "unknown")) content = str(turn.get("content", "")) return f"{role}: {content}" def is_abstention(instance: dict[str, Any]) -> bool: question_id = str(instance.get("question_id", "")) return question_id.endswith("_abs") def retrieve(*, query: str, documents: list[Document], embedder: Embedder, top_k: int) -> list[Document]: query_embedding = embedder.embed(query, input_type="query") document_embeddings = embedder.embed_many([document.text for document in documents], input_type="passage") scored = [] for document, document_embedding in zip(documents, document_embeddings): scored.append( Document( identifier=document.identifier, text=document.text, score=cosine_similarity(query_embedding, document_embedding), ) ) scored.sort(key=lambda document: document.score, reverse=True) return scored[:top_k] class HashingEmbedder: def embed(self, text: str, input_type: str = "passage") -> list[float]: del input_type vector = [0.0] * 768 for token in text.lower().split(): digest = hashlib.blake2b(token.encode("utf-8"), digest_size=4).digest() index = int.from_bytes(digest, byteorder="big") % len(vector) vector[index] += 1.0 return normalize_vector(vector) def embed_many(self, texts: list[str], input_type: str = "passage") -> list[list[float]]: return [self.embed(text, input_type=input_type) for text in texts] class HFE5Embedder: def __init__(self, model_name: str) -> None: import torch from transformers import AutoModel, AutoTokenizer self.torch = torch self.tokenizer = AutoTokenizer.from_pretrained(model_name) self.model = AutoModel.from_pretrained(model_name) self.model.eval() self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.model.to(self.device) self.cache: dict[tuple[str, str], list[float]] = {} self.batch_size = 64 if self.device.type == "cuda" else 32 print(f"[retrieval:hf-e5] device={self.device} batch_size={self.batch_size}", flush=True) def embed(self, text: str, input_type: str = "passage") -> list[float]: return self.embed_many([text], input_type=input_type)[0] def embed_many(self, texts: list[str], input_type: str = "passage") -> list[list[float]]: keys = [(input_type, text) for text in texts] missing_texts = [text for key, text in zip(keys, texts) if key not in self.cache] for start in range(0, len(missing_texts), self.batch_size): batch_texts = missing_texts[start : start + self.batch_size] if not batch_texts: continue prefix = "query: " if input_type == "query" else "passage: " inputs = self.tokenizer( [prefix + text.strip() for text in batch_texts], return_tensors="pt", truncation=True, padding=True, max_length=512, ) inputs = {key: value.to(self.device) for key, value in inputs.items()} with self.torch.no_grad(): output = self.model(**inputs) attention_mask = inputs["attention_mask"].unsqueeze(-1) embeddings = output.last_hidden_state * attention_mask pooled = embeddings.sum(dim=1) / attention_mask.sum(dim=1).clamp(min=1) pooled = self.torch.nn.functional.normalize(pooled, p=2, dim=1) for text, vector in zip(batch_texts, pooled.cpu()): self.cache[(input_type, text)] = [float(value) for value in vector] return [self.cache[key] for key in keys] def calculate_query_metrics(relevant: list[str], retrieved: list[str]) -> dict[str, Any]: relevant_set = set(relevant) metrics: dict[str, Any] = { "relevant_count": len(relevant_set), "reciprocal_rank_at_5": reciprocal_rank_at_k(relevant, retrieved, limit=5), } for limit in TOP_K_LIMITS: hit_count = len(relevant_set & set(retrieved[:limit])) precision = hit_count / limit recall = hit_count / max(len(relevant_set), 1) hit_rate = 1.0 if hit_count > 0 else 0.0 f1 = 2 * precision * recall / (precision + recall) if precision + recall > 0 else 0.0 metrics[f"hit_count_at_{limit}"] = hit_count metrics[f"precision_at_{limit}"] = precision metrics[f"recall_at_{limit}"] = recall metrics[f"hit_rate_at_{limit}"] = hit_rate metrics[f"f1_at_{limit}"] = f1 validate_query_invariants(metrics) return metrics def reciprocal_rank_at_k(relevant: list[str], retrieved: list[str], limit: int) -> float: relevant_set = set(relevant) for index, identifier in enumerate(retrieved[:limit], start=1): if identifier in relevant_set: return 1.0 / index return 0.0 def validate_query_invariants(metrics: dict[str, Any]) -> None: tolerance = 1e-12 if not ( metrics["recall_at_1"] <= metrics["recall_at_3"] + tolerance and metrics["recall_at_3"] <= metrics["recall_at_5"] + tolerance ): raise ValueError(f"Recall@k monotonicity invariant failed: {metrics}") if not ( metrics["hit_rate_at_1"] <= metrics["hit_rate_at_3"] + tolerance and metrics["hit_rate_at_3"] <= metrics["hit_rate_at_5"] + tolerance ): raise ValueError(f"HitRate@k monotonicity invariant failed: {metrics}") if metrics["reciprocal_rank_at_5"] > metrics["hit_rate_at_5"] + tolerance: raise ValueError(f"MRR@5 <= HitRate@5 invariant failed: {metrics}") for limit in TOP_K_LIMITS: if metrics[f"recall_at_{limit}"] > metrics[f"hit_rate_at_{limit}"] + tolerance: raise ValueError(f"Recall@{limit} <= HitRate@{limit} invariant failed: {metrics}") def validate_metric_invariants(report: dict[str, Any]) -> None: tolerance = 1e-12 if not ( report["recall_at_1"] <= report["recall_at_3"] + tolerance and report["recall_at_3"] <= report["recall_at_5"] + tolerance ): raise ValueError(f"Aggregate Recall@k monotonicity invariant failed: {report}") if not ( report["hit_rate_at_1"] <= report["hit_rate_at_3"] + tolerance and report["hit_rate_at_3"] <= report["hit_rate_at_5"] + tolerance ): raise ValueError(f"Aggregate HitRate@k monotonicity invariant failed: {report}") if report["mrr_at_5"] > report["hit_rate_at_5"] + tolerance: raise ValueError(f"Aggregate MRR@5 <= HitRate@5 invariant failed: {report}") for limit in TOP_K_LIMITS: if report[f"recall_at_{limit}"] > report[f"hit_rate_at_{limit}"] + tolerance: raise ValueError(f"Aggregate Recall@{limit} <= HitRate@{limit} invariant failed: {report}") def metric_definitions() -> dict[str, str]: return { "precision_at_k": "Mean share of retrieved top-k items that are relevant, using k as denominator.", "recall_at_k": "Mean share of all relevant items recovered in the retrieved top-k items.", "f1_at_k": "Mean harmonic average of per-query Precision@k and Recall@k.", "hit_rate_at_k": "Mean indicator that at least one relevant item appears in the retrieved top-k items.", "mrr_at_5": "Mean reciprocal rank of the first relevant item within the retrieved top-5 items.", } def default_audit_output(output_path: Path) -> Path: return output_path.with_name(f"{output_path.stem}_audit.jsonl") def write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None: path.parent.mkdir(parents=True, exist_ok=True) with path.open("w", encoding="utf-8") as file: for row in rows: file.write(json.dumps(row, ensure_ascii=False) + "\n") def cosine_similarity(left: list[float], right: list[float]) -> float: numerator = sum(left_value * right_value for left_value, right_value in zip(left, right)) left_norm = math.sqrt(sum(value * value for value in left)) right_norm = math.sqrt(sum(value * value for value in right)) if left_norm == 0.0 or right_norm == 0.0: return 0.0 return numerator / (left_norm * right_norm) def normalize_vector(vector: list[float]) -> list[float]: norm = math.sqrt(sum(value * value for value in vector)) if norm == 0.0: return vector return [value / norm for value in vector] if __name__ == "__main__": raise SystemExit(main())