| """记忆系统:短期记忆 + 长期记忆 + 遗忘曲线""" |
|
|
| 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 |
| timestamp: float = field(default_factory=time.time) |
| importance: float = 0.5 |
| 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 |
| |
| decay = math.pow(2, -elapsed / half_life) |
| |
| 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 |
| elif item.access_count > 5: |
| base *= 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]] = {} |
| 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 |
|
|