| 91 | json.dump([item.to_dict() for item in self.contents], file, indent=2, ensure_ascii=False) |
| 92 | |
| 93 | def retrieve( |
| 94 | self, |
| 95 | agent_role: str, |
| 96 | query: MemoryContentSnapshot, |
| 97 | top_k: int, |
| 98 | similarity_threshold: float, |
| 99 | ) -> List[MemoryItem]: |
| 100 | if self.count_memories() == 0 or not self.embedding: |
| 101 | return [] |
| 102 | |
| 103 | # Build an optimized query for retrieval |
| 104 | query_text = self.retrieve_prompt.format(input=query.text) |
| 105 | query_text = self._extract_key_content(query_text) |
| 106 | |
| 107 | inputs_embedding = self.embedding.get_embedding(query_text) |
| 108 | if isinstance(inputs_embedding, list): |
| 109 | inputs_embedding = np.array(inputs_embedding, dtype=np.float32) |
| 110 | inputs_embedding = inputs_embedding.reshape(1, -1) |
| 111 | faiss.normalize_L2(inputs_embedding) |
| 112 | |
| 113 | expected_dim = inputs_embedding.shape[1] |
| 114 | |
| 115 | memory_embeddings = [] |
| 116 | valid_items = [] |
| 117 | for item in self.contents: |
| 118 | if item.embedding is not None: |
| 119 | if len(item.embedding) != expected_dim: |
| 120 | logger.warning( |
| 121 | "Skipping memory item %s: embedding dim %d != expected %d", |
| 122 | item.id, len(item.embedding), expected_dim, |
| 123 | ) |
| 124 | continue |
| 125 | memory_embeddings.append(item.embedding) |
| 126 | valid_items.append(item) |
| 127 | |
| 128 | if not memory_embeddings: |
| 129 | return [] |
| 130 | |
| 131 | memory_embeddings = np.array(memory_embeddings, dtype=np.float32) |
| 132 | |
| 133 | # Use an efficient inner-product index |
| 134 | index = faiss.IndexFlatIP(memory_embeddings.shape[1]) |
| 135 | index.add(memory_embeddings) |
| 136 | |
| 137 | # Retrieve extra candidates for reranking |
| 138 | retrieval_k = min(top_k * 3, len(valid_items)) |
| 139 | similarities, indices = index.search(inputs_embedding, retrieval_k) |
| 140 | |
| 141 | # Filter and rerank the candidates |
| 142 | candidates = [] |
| 143 | for i in range(len(indices[0])): |
| 144 | idx = indices[0][i] |
| 145 | similarity = similarities[0][i] |
| 146 | |
| 147 | if idx != -1 and similarity >= similarity_threshold: |
| 148 | item = valid_items[idx] |
| 149 | # Calculate an auxiliary semantic similarity score |
| 150 | semantic_score = self._calculate_semantic_similarity(query_text, item.content_summary) |