"""记忆系统:短期记忆 + 长期记忆 + 遗忘曲线""" import time import json import heapq from dataclasses import dataclass, field from typing import List, Dict, Optional, Any from collections import deque import math @dataclass class MemoryItem: """记忆条目""" id: str content: str memory_type: str # "conversation" / "event" / "preference" / "feedback" timestamp: float = field(default_factory=time.time) importance: float = 0.5 # 0-1 access_count: int = 0 last_access: float = field(default_factory=time.time) tags: List[str] = field(default_factory=list) related_ids: List[str] = field(default_factory=list) @property def age(self) -> float: return time.time() - self.timestamp def strength(self, half_life: float = 86400.0) -> float: """基于遗忘曲线的记忆强度""" elapsed = self.age # 指数遗忘: S = importance × 2^(-elapsed/half_life) decay = math.pow(2, -elapsed / half_life) # 访问增强: boost = log2(access_count + 1) × 0.1 boost = math.log2(self.access_count + 1) * 0.1 return min(1.0, self.importance * decay + boost) class ForgettingCurve: """ 遗忘曲线:决定记忆的半衰期和遗忘速度 公式: S(t) = I · 2^(-t / τ) 其中: S(t) = t时刻的记忆强度 I = 初始重要性 τ = 半衰期(新记忆默认1天,高频访问的记忆半衰期延长) """ def __init__(self, default_half_life: float = 86400.0): self.default_half_life = default_half_life def compute_half_life(self, item: MemoryItem) -> float: """根据访问频率动态调整半衰期""" base = self.default_half_life if item.access_count > 10: base *= 7 # 高频访问延至1周 elif item.access_count > 5: base *= 3 # 中频延至3天 # 高重要性记忆更持久 base *= (0.5 + item.importance) return base def should_forget(self, item: MemoryItem, threshold: float = 0.05) -> bool: """判断是否应该遗忘(强度低于阈值)""" half_life = self.compute_half_life(item) return item.strength(half_life) < threshold class MemoryGraph: """ 异构图记忆网络 节点类型: - user: 用户实体 - conversation: 对话片段 - event: 事件 - preference: 偏好 - topic: 话题 边类型: - has_topic: 对话→话题 - mentioned: 对话→用户实体 - leads_to: 事件→事件(因果关系) - prefers: 用户→偏好 """ def __init__(self, config=None): self._items: Dict[str, MemoryItem] = {} self._graph: Dict[str, List[tuple]] = {} # node_id -> [(edge_type, target_id)] self._forgetting = ForgettingCurve() self._short_term: deque = deque(maxlen=100) self._long_term_capacity = 10000 def add(self, item: MemoryItem, relations: List[tuple] = None) -> str: """添加记忆""" self._items[item.id] = item self._short_term.append(item.id) if item.memory_type in ("conversation", "event"): self._graph[item.id] = relations or [] for edge_type, target in (relations or []): if target not in self._graph: self._graph[target] = [] self._graph[target].append((f"inverse_{edge_type}", item.id)) return item.id def get(self, item_id: str) -> Optional[MemoryItem]: """获取记忆并更新访问计数""" item = self._items.get(item_id) if item: item.access_count += 1 item.last_access = time.time() return item def recall_by_topic(self, topic: str, top_k: int = 5) -> List[MemoryItem]: """按话题召回""" results = [] for item in self._items.values(): if topic in item.tags or topic in item.content: # 跳过已被遗忘的 if self._forgetting.should_forget(item): continue results.append((item.strength(), item)) results.sort(key=lambda x: -x[0]) return [item for _, item in results[:top_k]] def recall_recent(self, n: int = 10) -> List[MemoryItem]: """召回最近的n条记忆""" ids = list(self._short_term)[-n:] return [self._items[i] for i in ids if i in self._items] def prune(self): """清理被遗忘的记忆""" to_delete = [] for item_id, item in self._items.items(): if self._forgetting.should_forget(item): to_delete.append(item_id) for item_id in to_delete: del self._items[item_id] self._graph.pop(item_id, None) def to_json(self) -> str: data = { "items": {k: v.__dict__ for k, v in self._items.items()}, "graph": {k: v for k, v in self._graph.items()}, } return json.dumps(data, ensure_ascii=False) @classmethod def from_json(cls, json_str: str) -> "MemoryGraph": mg = cls() data = json.loads(json_str) for k, v in data.get("items", {}).items(): item = MemoryItem(**v) mg._items[k] = item mg._short_term.append(k) mg._graph = data.get("graph", {}) return mg