Project-Aura / aura /memory.py
ljsysfurry's picture
Upload aura/memory.py with huggingface_hub
96cf104 verified
Raw
History Blame Contribute Delete
5.48 kB
"""记忆系统:短期记忆 + 长期记忆 + 遗忘曲线"""
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