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