File size: 5,480 Bytes
96cf104 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 | """记忆系统:短期记忆 + 长期记忆 + 遗忘曲线"""
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
|